Merge pull request #1462 from mattn/fix-conn-use-after-close
Return an error instead of crashing when a connection is used after Close
diff --git a/sqlite3.go b/sqlite3.go
index 2dcc2d3..a092c09 100644
--- a/sqlite3.go
+++ b/sqlite3.go
@@ -752,6 +752,12 @@
return err
}
+// errConnClosed is returned by methods called on a connection whose
+// Close has already run. SQLite dereferences the connection handle
+// without checking it, so calling into it with a released handle
+// crashes the process instead of reporting an error.
+var errConnClosed = errors.New("sqlite connection is already closed")
+
// RegisterCollation makes a Go function available as a collation.
//
// cmp receives two UTF-8 strings, a and b. The result should be 0 if
@@ -765,6 +771,9 @@
// If cmp does not obey these constraints, sqlite3's behavior is
// undefined when the collation is used.
func (c *SQLiteConn) RegisterCollation(name string, cmp func(string, string) int) error {
+ if !c.dbConnOpen() {
+ return errConnClosed
+ }
handle := newHandle(c, cmp)
cname := C.CString(name)
defer C.free(unsafe.Pointer(cname))
@@ -782,7 +791,12 @@
// If there is an existing commit hook for this connection, it will be
// removed. If callback is nil the existing hook (if any) will be removed
// without creating a new one.
+//
+// The call is a no-op once the connection has been closed.
func (c *SQLiteConn) RegisterCommitHook(callback func() int) {
+ if !c.dbConnOpen() {
+ return
+ }
if callback == nil {
C.sqlite3_commit_hook(c.db, nil, nil)
} else {
@@ -795,7 +809,12 @@
// If there is an existing rollback hook for this connection, it will be
// removed. If callback is nil the existing hook (if any) will be removed
// without creating a new one.
+//
+// The call is a no-op once the connection has been closed.
func (c *SQLiteConn) RegisterRollbackHook(callback func()) {
+ if !c.dbConnOpen() {
+ return
+ }
if callback == nil {
C.sqlite3_rollback_hook(c.db, nil, nil)
} else {
@@ -812,7 +831,12 @@
// If there is an existing update hook for this connection, it will be
// removed. If callback is nil the existing hook (if any) will be removed
// without creating a new one.
+//
+// The call is a no-op once the connection has been closed.
func (c *SQLiteConn) RegisterUpdateHook(callback func(int, string, string, int64)) {
+ if !c.dbConnOpen() {
+ return
+ }
if callback == nil {
C.sqlite3_update_hook(c.db, nil, nil)
} else {
@@ -826,7 +850,12 @@
// SQLITE_INSERT, SQLITE_DELETE, or SQLITE_UPDATE), and 1 to 3 arguments,
// depending on operation. More details see:
// https://www.sqlite.org/c3ref/c_alter_table.html
+//
+// The call is a no-op once the connection has been closed.
func (c *SQLiteConn) RegisterAuthorizer(callback func(int, string, string, string) int) {
+ if !c.dbConnOpen() {
+ return
+ }
c.authorizerMu.Lock()
defer c.authorizerMu.Unlock()
@@ -861,6 +890,9 @@
//
// See _example/go_custom_funcs for a detailed example.
func (c *SQLiteConn) RegisterFunc(name string, impl any, pure bool) error {
+ if !c.dbConnOpen() {
+ return errConnClosed
+ }
var fi functionInfo
fi.f = reflect.ValueOf(impl)
t := fi.f.Type()
@@ -943,6 +975,9 @@
//
// See _example/go_custom_funcs for a detailed example.
func (c *SQLiteConn) RegisterAggregator(name string, impl any, pure bool) error {
+ if !c.dbConnOpen() {
+ return errConnClosed
+ }
var ai aggInfo
ai.constructor = reflect.ValueOf(impl)
t := ai.constructor.Type()
@@ -1051,9 +1086,13 @@
}
// AutoCommit return which currently auto commit or not.
+// It reports false once the connection has been closed.
func (c *SQLiteConn) AutoCommit() bool {
c.mu.Lock()
defer c.mu.Unlock()
+ if c.db == nil {
+ return false
+ }
return int(C.sqlite3_get_autocommit(c.db)) != 0
}
@@ -2225,8 +2264,12 @@
// GetFilename returns the absolute path to the file containing
// the requested schema. When passed an empty string, it will
// instead use the database's default schema: "main".
+// It returns an empty string once the connection has been closed.
// See: sqlite3_db_filename, https://www.sqlite.org/c3ref/db_filename.html
func (c *SQLiteConn) GetFilename(schemaName string) string {
+ if !c.dbConnOpen() {
+ return ""
+ }
if schemaName == "" {
schemaName = "main"
}
@@ -2236,15 +2279,23 @@
}
// GetLimit returns the current value of a run-time limit.
+// It returns -1 once the connection has been closed.
// See: sqlite3_limit, http://www.sqlite.org/c3ref/limit.html
func (c *SQLiteConn) GetLimit(id int) int {
+ if !c.dbConnOpen() {
+ return -1
+ }
return int(C._sqlite3_limit(c.db, C.int(id), C.int(-1)))
}
// SetLimit changes the value of a run-time limits.
// Then this method returns the prior value of the limit.
+// It returns -1 once the connection has been closed.
// See: sqlite3_limit, http://www.sqlite.org/c3ref/limit.html
func (c *SQLiteConn) SetLimit(id int, newVal int) int {
+ if !c.dbConnOpen() {
+ return -1
+ }
return int(C._sqlite3_limit(c.db, C.int(id), C.int(newVal)))
}
@@ -2261,6 +2312,9 @@
//
// See: sqlite3_file_control, https://www.sqlite.org/c3ref/file_control.html
func (c *SQLiteConn) SetFileControlInt(dbName string, op int, arg int) error {
+ if !c.dbConnOpen() {
+ return errConnClosed
+ }
if dbName == "" {
dbName = "main"
}
@@ -2289,6 +2343,9 @@
//
// See: sqlite3_file_control, https://www.sqlite.org/c3ref/file_control.html
func (c *SQLiteConn) SetFileControlInt64(dbName string, op int, arg int64) error {
+ if !c.dbConnOpen() {
+ return errConnClosed
+ }
if dbName == "" {
dbName = "main"
}
diff --git a/sqlite3_test.go b/sqlite3_test.go
index 506d409..4fec621 100644
--- a/sqlite3_test.go
+++ b/sqlite3_test.go
@@ -2980,3 +2980,88 @@
}
})
}
+
+// TestConnMethodsAfterClose exercises the *SQLiteConn methods that reach
+// into SQLite with the connection handle. SQLite dereferences that handle
+// without checking it, so before these guards every one of them crashed
+// the process when called on a connection captured from a ConnectHook and
+// used after Close.
+func TestConnMethodsAfterClose(t *testing.T) {
+ driverName := fmt.Sprintf("sqlite3_after_close_%d", time.Now().UnixNano())
+ var conn *SQLiteConn
+ sql.Register(driverName, &SQLiteDriver{
+ ConnectHook: func(c *SQLiteConn) error {
+ conn = c
+ return nil
+ },
+ })
+ db, err := sql.Open(driverName, ":memory:")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := db.Ping(); err != nil {
+ db.Close()
+ t.Fatal(err)
+ }
+ if err := db.Close(); err != nil {
+ t.Fatal(err)
+ }
+ if conn == nil {
+ t.Fatal("ConnectHook did not run")
+ }
+
+ t.Run("errors", func(t *testing.T) {
+ cases := []struct {
+ name string
+ call func() error
+ }{
+ {"RegisterCollation", func() error {
+ return conn.RegisterCollation("c", func(string, string) int { return 0 })
+ }},
+ {"RegisterFunc", func() error {
+ return conn.RegisterFunc("f", func() int64 { return 1 }, true)
+ }},
+ {"RegisterAggregator", func() error {
+ return conn.RegisterAggregator("a", func() *sumAggregator {
+ var ret sumAggregator
+ return &ret
+ }, true)
+ }},
+ {"SetFileControlInt", func() error {
+ return conn.SetFileControlInt("", SQLITE_FCNTL_CHUNK_SIZE, 4096)
+ }},
+ {"SetFileControlInt64", func() error {
+ return conn.SetFileControlInt64("", SQLITE_FCNTL_CHUNK_SIZE, 4096)
+ }},
+ }
+ for _, tc := range cases {
+ if err := tc.call(); !errors.Is(err, errConnClosed) {
+ t.Errorf("%s error = %v, want %v", tc.name, err, errConnClosed)
+ }
+ }
+ })
+
+ t.Run("noops", func(t *testing.T) {
+ // These have no error to report, so they must simply do nothing.
+ conn.RegisterCommitHook(func() int { return 0 })
+ conn.RegisterRollbackHook(func() {})
+ conn.RegisterUpdateHook(func(int, string, string, int64) {})
+ conn.RegisterAuthorizer(func(int, string, string, string) int { return SQLITE_OK })
+ conn.RegisterAuthorizer(nil)
+ })
+
+ t.Run("values", func(t *testing.T) {
+ if got := conn.AutoCommit(); got {
+ t.Errorf("AutoCommit() = %v, want false", got)
+ }
+ if got := conn.GetFilename("main"); got != "" {
+ t.Errorf("GetFilename() = %q, want empty", got)
+ }
+ if got := conn.GetLimit(SQLITE_LIMIT_LENGTH); got != -1 {
+ t.Errorf("GetLimit() = %d, want -1", got)
+ }
+ if got := conn.SetLimit(SQLITE_LIMIT_LENGTH, 1000); got != -1 {
+ t.Errorf("SetLimit() = %d, want -1", got)
+ }
+ })
+}