aboutsummaryrefslogtreecommitdiff
path: root/pgruntime
diff options
context:
space:
mode:
authorStefan Majewsky <majewsky@gmx.net>2026-07-31 22:27:58 +0200
committerStefan Majewsky <majewsky@gmx.net>2026-07-31 22:27:59 +0200
commit5bff2a6c9e14763ca8764cf75adba0efbf990e50 (patch)
tree466030af2d3da2210e806e083567ee35893adb1f /pgruntime
parentc953ebcc54296ac36fd199fa66fb0763f1790334 (diff)
downloadgo-gg-5bff2a6c9e14763ca8764cf75adba0efbf990e50.tar.gz
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.
Diffstat (limited to 'pgruntime')
-rw-r--r--pgruntime/behavior.go3
-rw-r--r--pgruntime/connector.go85
-rw-r--r--pgruntime/helpers.go35
-rw-r--r--pgruntime/pgruntime.go1
4 files changed, 95 insertions, 29 deletions
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