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)
+		}
+	})
+}