1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
|
// SPDX-FileCopyrightText: 2026 Stefan Majewsky <majewsky@gmx.net>
// SPDX-License-Identifier: Apache-2.0
package pgruntime
import (
"context"
"database/sql"
"fmt"
"regexp"
"strings"
"go.xyrillian.de/gg/assert"
"go.xyrillian.de/gg/errext"
"go.xyrillian.de/gg/gsql"
)
// Connector describes how to connect to a PostgreSQL database given a [libpq-style connection URI].
// 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)
// 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)
if err != nil {
return nil, err
}
return gsql.NewDB(db), nil
}
}
// Connect connects to a PostgreSQL database.
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())
if err != nil {
return none, err
}
err = behavior.applyTo(ctx, db)
if err != nil {
return none, errext.WithCleanup(err, "db.Close", db.GSQLClose(ctx))
}
return db, nil
}
// TestSetupOption is an optional behavior that can be given to [Connector.ConnectForTest].
type TestSetupOption func(*testSetupParams)
type testSetupParams struct {
DatabaseName string
}
// OverrideDatabaseName is a [TestSetupOption] that picks a different database name than the default of t.Name().
//
// This is only necessary if a single test needs to use multiple database connections at the same time,
// e.g. to simulate two separate deployments of the application next to each other.
func OverrideDatabaseName(dbName string) TestSetupOption {
return func(params *testSetupParams) {
params.DatabaseName = dbName
}
}
// ConnectForTest connects to the test database server managed by [WithTestDB].
//
// Each test will run in its own separate database (whose name is the same as t.Name()),
// so it is safe to mark tests as t.Parallel() to run multiple tests within the same package concurrently.
// ConnectForTest will always drop and recreate the database to ensure deterministic test behavior.
//
// Some tests require setting up multiple separate connections to the same database.
// The second return argument can be be used with [Connector.Connect] to obtain additional connections as needed.
func (c Connector[T]) ConnectForTest(t assert.TestingTB, behavior ConnectionBehavior, opts ...TestSetupOption) (T, ConnectionTarget) {
t.Helper()
ctx := t.Context()
params := testSetupParams{
DatabaseName: t.Name(),
}
for _, opt := range opts {
opt(¶ms)
}
// 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
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)
}
// 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)
target := ConnectionTarget{
HostName: "127.0.0.1",
Port: testdbPort,
UserName: testdbUserName,
DatabaseName: "postgres",
ConnectionOptions: "sslmode=disable",
}
err := c.prepareTestDatabase(ctx, target, dbName)
if err != nil {
t.Fatal(err.Error())
}
// connect to actual test database
target.DatabaseName = dbName
handle, err := c.Connect(ctx, target, behavior)
if err != nil {
t.Fatal(err.Error())
}
return handle, target
}
func (c Connector[T]) prepareTestDatabase(ctx context.Context, target ConnectionTarget, dbName string) error {
u, err := target.IntoURL()
if err != nil {
return 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)
}
_, err = execQuery(ctx, db, "DROP DATABASE IF EXISTS "+quoteIdentifier(dbName), nil)
if err != nil {
return errext.WithCleanup(fmt.Errorf("during DROP DATABASE: %w", err), "db.Close", db.GSQLClose(ctx))
}
_, err = execQuery(ctx, db, "CREATE DATABASE "+quoteIdentifier(dbName), nil)
if err != nil {
return errext.WithCleanup(fmt.Errorf("during CREATE DATABASE: %w", err), "db.Close", db.GSQLClose(ctx))
}
return errext.WithCleanup(nil, "db.Close", db.GSQLClose(ctx))
}
|