From 5bff2a6c9e14763ca8764cf75adba0efbf990e50 Mon Sep 17 00:00:00 2001 From: Stefan Majewsky Date: Fri, 31 Jul 2026 22:27:58 +0200 Subject: pgruntime: in ConnectForTest, only truncate tables instead of recreating the database DROP + CREATE DATABASE was measured at 100 ms per ConnectForTest(), which is prohibitively slow in large test suites. This should be more in the ballpark of 10 ms per test. --- CHANGELOG.md | 2 ++ pgruntime/behavior.go | 3 +- pgruntime/connector.go | 85 +++++++++++++++++++++++++++++++++++--------------- pgruntime/helpers.go | 35 +++++++++++++++++++++ pgruntime/pgruntime.go | 1 - 5 files changed, 97 insertions(+), 29 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 92f062d..34962eb 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,8 @@ SPDX-License-Identifier: Apache-2.0 Changes: - Add package pgruntime. +- TODO: add SelectOneValue, SelectSeveralValues, etc. to gsql +- TODO: add WithinTransaction to gsql.DB, gsql.Conn (and same on gg-pgx) # v1.12.0 (2026-07-31) diff --git a/pgruntime/behavior.go b/pgruntime/behavior.go index 9fb8662..3ffff1e 100644 --- a/pgruntime/behavior.go +++ b/pgruntime/behavior.go @@ -68,8 +68,7 @@ func applyMigrations(ctx context.Context, db gsql.ConnectionHandle, migrations m } // read schema_migrations table - var rowCount int64 - err = queryRow(ctx, db, `SELECT COUNT(*) FROM schema_migrations`, nil, []any{&rowCount}) + rowCount, err := selectOneValue[int64](ctx, db, `SELECT COUNT(*) FROM schema_migrations`) if err != nil { return fmt.Errorf("could not check row count for schema_migrations: %w", err) } diff --git a/pgruntime/connector.go b/pgruntime/connector.go index 60e0de2..92780f6 100644 --- a/pgruntime/connector.go +++ b/pgruntime/connector.go @@ -5,7 +5,9 @@ package pgruntime import ( "context" + "crypto/sha256" "database/sql" + "encoding/base32" "fmt" "regexp" "strings" @@ -19,13 +21,17 @@ import ( // This type acts as a dependency injection surface, abstracting the different ways in which different database drivers and libraries perform connection. // // [libpq-style connection URI]: https://www.postgresql.org/docs/current/libpq-connect.html#LIBPQ-CONNSTRING-URIS -type Connector[T gsql.ConnectionHandle] func(ctx context.Context, dbURL string) (T, error) +type Connector[T gsql.ConnectionHandle] func(context.Context, ConnectionTarget) (T, error) // StdConnector returns a [Connector] for database/sql drivers. // When used with [lib/pq], the driver name must be "postgres". func StdConnector(driverName string) Connector[*gsql.DB] { - return func(ctx context.Context, dbURL string) (*gsql.DB, error) { - db, err := sql.Open(driverName, dbURL) + return func(ctx context.Context, target ConnectionTarget) (*gsql.DB, error) { + u, err := target.IntoURL() + if err != nil { + return nil, err + } + db, err := sql.Open(driverName, u.String()) if err != nil { return nil, err } @@ -37,11 +43,7 @@ func StdConnector(driverName string) Connector[*gsql.DB] { func (c Connector[T]) Connect(ctx context.Context, target ConnectionTarget, behavior ConnectionBehavior) (T, error) { var none T // shorthand for error return paths - u, err := target.IntoURL() - if err != nil { - return none, err - } - db, err := c(ctx, u.String()) + db, err := c(ctx, target) if err != nil { return none, err } @@ -90,15 +92,16 @@ func (c Connector[T]) ConnectForTest(t assert.TestingTB, behavior ConnectionBeha // normalize t.Name() into an acceptable database name for PostgreSQL // - only alphanumerics and underscore -> replace all other symbols with _ - // - max 63 chars -> reject longer names + // - max 63 chars -> if overflown, truncate and append a short digest to hopefully make it unique dbName := strings.ToLower(params.DatabaseName) dbName = regexp.MustCompile(`[^a-z_]`).ReplaceAllString(dbName, "_") if len(dbName) > 63 { - t.Fatalf("cannot use t.Name() = %q (normalized to %q) as a database name because it is longer than 63 chars", params.DatabaseName, dbName) + digest := sha256.Sum256([]byte(params.DatabaseName)) + encoded := base32.HexEncoding.EncodeToString(digest[:]) + dbName = dbName[0:53] + "__" + strings.ToLower(encoded[0:8]) } - // connect to "postgres" database for the DROP/CREATE DATABASE queries - // TODO: DROP/CREATE DATABASE turns out to be very slow (in a real-world scenario: 100ms per test instead of ~10ms to wipe just the DB contents and reset sequences) + // connect to "postgres" database for the CREATE DATABASE query (if necessary) target := ConnectionTarget{ HostName: "127.0.0.1", Port: testdbPort, @@ -106,36 +109,66 @@ func (c Connector[T]) ConnectForTest(t assert.TestingTB, behavior ConnectionBeha DatabaseName: "postgres", ConnectionOptions: "sslmode=disable", } - err := c.prepareTestDatabase(ctx, target, dbName) + adminDB, err := c(ctx, target) + if err != nil { + t.Fatal(err.Error() + " (if this error is about the database server not running, check if your TestMain() calls pgruntime.WithTestDB())") + } + err = createDatabaseIfMissing(ctx, adminDB, dbName) + err = errext.WithCleanup(err, "db.Close", adminDB.GSQLClose(ctx)) if err != nil { t.Fatal(err.Error()) } // connect to actual test database target.DatabaseName = dbName - handle, err := c.Connect(ctx, target, behavior) + testDB, err := c.Connect(ctx, target, behavior) + if err != nil { + t.Fatal(err.Error()) + } + err = resetTestDatabase(ctx, testDB, params, behavior) if err != nil { + err = errext.WithCleanup(err, "db.Close", testDB.GSQLClose(ctx)) t.Fatal(err.Error()) } - return handle, target + return testDB, target } -func (c Connector[T]) prepareTestDatabase(ctx context.Context, target ConnectionTarget, dbName string) error { - u, err := target.IntoURL() +func createDatabaseIfMissing(ctx context.Context, db gsql.Handle, dbName string) error { + // check if database exists + exists, err := selectOneValue[bool](ctx, db, `SELECT COUNT(*) > 0 FROM pg_catalog.pg_database WHERE datname = $1`, dbName) if err != nil { - return err + return fmt.Errorf("while reading from pg_catalog.pg_database: %w", err) } - db, err := c(ctx, u.String()) - if err != nil { - return fmt.Errorf("%w (if this error is about the database server not running, check if your TestMain() calls pgruntime.WithTestDB())", err) + + // create database if necessary + if !exists { + _, err = execQuery(ctx, db, "CREATE DATABASE "+quoteIdentifier(dbName), nil) + if err != nil { + return fmt.Errorf("during CREATE DATABASE: %w", err) + } + } + + return nil +} + +func resetTestDatabase(ctx context.Context, db gsql.Handle, params testSetupParams, behavior ConnectionBehavior) error { + // enumerate all tables that need to be truncated (all tables that are not managed by pgruntime) + condition := `table_schema = 'public' AND table_type = 'BASE TABLE'` + if len(behavior.Migrations) > 0 { + condition += ` AND table_name != 'schema_migrations'` } - _, err = execQuery(ctx, db, "DROP DATABASE IF EXISTS "+quoteIdentifier(dbName), nil) + query := fmt.Sprintf(`SELECT quote_ident(table_name) FROM information_schema.tables WHERE %s ORDER BY table_name`, condition) + quotedTableNames, err := selectSeveralValues[string](ctx, db, query) if err != nil { - return errext.WithCleanup(fmt.Errorf("during DROP DATABASE: %w", err), "db.Close", db.GSQLClose(ctx)) + return fmt.Errorf("while listing tables to truncate: %w", err) } - _, err = execQuery(ctx, db, "CREATE DATABASE "+quoteIdentifier(dbName), nil) + + // truncate all tables at once + query = fmt.Sprintf(`TRUNCATE %s RESTART IDENTITY CASCADE`, strings.Join(quotedTableNames, ", ")) + _, err = execQuery(ctx, db, query, nil) if err != nil { - return errext.WithCleanup(fmt.Errorf("during CREATE DATABASE: %w", err), "db.Close", db.GSQLClose(ctx)) + return fmt.Errorf("during %s: %w", query, err) } - return errext.WithCleanup(nil, "db.Close", db.GSQLClose(ctx)) + + return nil } diff --git a/pgruntime/helpers.go b/pgruntime/helpers.go index 11bc84e..1019e6e 100644 --- a/pgruntime/helpers.go +++ b/pgruntime/helpers.go @@ -13,6 +13,7 @@ import ( ) // Convenience function for executing a one-off SQL query returning no rows. +// TODO: move to gsql func execQuery(ctx context.Context, db gsql.Handle, query string, args []any) (sql.Result, error) { stmt, err := db.GSQLPrepare(ctx, query, false) if err != nil { @@ -32,6 +33,40 @@ func queryRow(ctx context.Context, db gsql.Handle, query string, args, slots []a return errext.WithCleanup(err, "stmt.Close", stmt.Close()) } +// Convenience function for executing a one-off SQL query returning one value. +// TODO: move to gsql +func selectOneValue[T any](ctx context.Context, db gsql.Handle, query string, args ...any) (T, error) { + stmt, err := db.GSQLPrepare(ctx, query, false) + if err != nil { + var none T + return none, err + } + + var result T + err = stmt.QueryRow(ctx, args, []any{&result}) + return result, errext.WithCleanup(err, "stmt.Close", stmt.Close()) +} + +// Convenience function for executing a one-off SQL query returning several single-column rows. +// TODO: move to gsql (and also add ForeachValue with a callback instead of a slice return, maybe even ForeachPair and ForeachTriple) +func selectSeveralValues[T any](ctx context.Context, db gsql.Handle, query string, args ...any) ([]T, error) { + rows, err := db.GSQLQuery(ctx, query, args) + if err != nil { + return nil, err + } + var result []T + for rows.Next() { + // TODO: this should share growRecordSlice() from Oblast to optimize allocations + var value T + err := rows.Scan(&value) + if err != nil { + return nil, errext.WithCleanup(err, "rows.Close", rows.Close()) + } + result = append(result, value) + } + return result, errext.WithCleanup(nil, "rows.Err", rows.Err()) +} + // Convenience function for preparing an identifier that needs to be inserted into a query verbatim // (e.g. a database name for CREATE DATABASE). func quoteIdentifier(name string) string { diff --git a/pgruntime/pgruntime.go b/pgruntime/pgruntime.go index ed44076..17a6c94 100644 --- a/pgruntime/pgruntime.go +++ b/pgruntime/pgruntime.go @@ -28,5 +28,4 @@ // [gg-pgx]: https://git.xyrillian.de/go-gg-pgx/ package pgruntime -// TODO: before merging this branch, start work on go-gg-pgx to verify that we're not painting ourselves into a corner with the gsql.Handle interfaces // TODO: test coverage via separate module importing github.com/lib/pq -- cgit v1.3.1