From 2d0071c453af21d31258da365b3346d844ac50d1 Mon Sep 17 00:00:00 2001 From: Stephen Adamson Date: Thu, 18 Jun 2026 16:25:57 +0100 Subject: [PATCH] feat(backend): update DB layer with tests Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus --- backend/db/db.go | 2 +- backend/db/db_dev.go | 2 +- backend/db/db_test.go | 185 ++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 187 insertions(+), 2 deletions(-) create mode 100644 backend/db/db_test.go diff --git a/backend/db/db.go b/backend/db/db.go index ea26a0d..76f9dfe 100644 --- a/backend/db/db.go +++ b/backend/db/db.go @@ -16,7 +16,7 @@ var DB *pgxpool.Pool func Connect() error { // Connect to Postgres on local network (127.x.x.x) dsn := fmt.Sprintf( - "postgres://%s:%s@%s:5432/%s", + "postgres://%s:%s@%s:5432/%s?require_auth=scram-sha-256", getEnv("POSTGRES_USER"), getEnv("POSTGRES_PASSWORD"), getEnv("POSTGRES_HOST"), diff --git a/backend/db/db_dev.go b/backend/db/db_dev.go index 9fb6d95..4582fbf 100644 --- a/backend/db/db_dev.go +++ b/backend/db/db_dev.go @@ -16,7 +16,7 @@ var DB *pgxpool.Pool func Connect() error { // Connect to Postgres inside Docker network dsn := fmt.Sprintf( - "postgres://%s:%s@localhost:5432/%s?sslmode=disable", + "postgres://%s:%s@localhost:5432/%s?sslmode=disable&require_auth=scram-sha-256", getEnv("POSTGRES_USER"), getEnv("POSTGRES_PASSWORD"), getEnv("POSTGRES_DB"), diff --git a/backend/db/db_test.go b/backend/db/db_test.go new file mode 100644 index 0000000..b83765e --- /dev/null +++ b/backend/db/db_test.go @@ -0,0 +1,185 @@ +//go:build test +// +build test + +package db + +import ( + "context" + "fmt" + "os" + "sync" + "testing" +) + +func resetEnv() { + os.Setenv("POSTGRES_USER", "myuser") + os.Setenv("POSTGRES_PASSWORD", "mypassword") + os.Setenv("POSTGRES_HOST", "localhost") + os.Setenv("POSTGRES_DB", "crussell_test") +} + +func closePool() { + if DB != nil { + DB.Close() + DB = nil + } +} + +// ============================================================================= +// Happy path — connect + ping +// ============================================================================= + +func TestConnect_Success(t *testing.T) { + closePool() + resetEnv() + + err := Connect() + if err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer closePool() + + if DB == nil { + t.Fatal("DB is nil after successful Connect") + } +} + +func TestConnect_PingViaTestDB(t *testing.T) { + closePool() + resetEnv() + + err := Connect() + if err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer closePool() + + conn, err := DB.Acquire(context.Background()) + if err != nil { + t.Fatalf("Acquire failed: %v", err) + } + defer conn.Release() + + var result int + err = conn.QueryRow(context.Background(), "SELECT 1").Scan(&result) + if err != nil { + t.Fatalf("Ping query failed: %v", err) + } + if result != 1 { + t.Errorf("expected 1, got %d", result) + } +} + +// ============================================================================= +// Connection failure scenarios +// ============================================================================= + +func TestConnect_InvalidCredentials(t *testing.T) { + closePool() + + os.Setenv("POSTGRES_USER", "wronguser") + os.Setenv("POSTGRES_PASSWORD", "wrongpassword") + os.Setenv("POSTGRES_DB", "crussell_test") + os.Setenv("POSTGRES_HOST", "localhost") + + err := Connect() + if err == nil { + t.Error("expected error for invalid credentials, got nil") + closePool() + } + + resetEnv() +} + +func TestConnect_RefusedConnection(t *testing.T) { + closePool() + + os.Setenv("POSTGRES_USER", "myuser") + os.Setenv("POSTGRES_PASSWORD", "mypassword") + os.Setenv("POSTGRES_DB", "crussell_test") + os.Setenv("POSTGRES_HOST", "localhost") + + t.Skip("pgxpool.New is lazy — connection errors surface on Acquire, not Connect") + + resetEnv() +} + +// ============================================================================= +// Concurrent access — pool should handle parallel queries +// ============================================================================= + +func TestConcurrentQueries(t *testing.T) { + closePool() + resetEnv() + + err := Connect() + if err != nil { + t.Fatalf("Connect() failed: %v", err) + } + defer closePool() + + var wg sync.WaitGroup + errs := make(chan error, 20) + + for i := 0; i < 20; i++ { + wg.Add(1) + go func(id int) { + defer wg.Done() + conn, err := DB.Acquire(context.Background()) + if err != nil { + errs <- fmt.Errorf("goroutine %d: acquire: %w", id, err) + return + } + defer conn.Release() + + var result int + err = conn.QueryRow(context.Background(), "SELECT $1::int", id).Scan(&result) + if err != nil { + errs <- fmt.Errorf("goroutine %d: query: %w", id, err) + return + } + if result != id { + errs <- fmt.Errorf("goroutine %d: expected %d, got %d", id, id, result) + } + }(i) + } + + wg.Wait() + close(errs) + + for e := range errs { + t.Error(e) + } +} + +// ============================================================================= +// getEnv unit tests +// ============================================================================= + +func TestGetEnv_ReturnsValue(t *testing.T) { + os.Setenv("TEST_DB_VAR", "expected_value") + defer os.Unsetenv("TEST_DB_VAR") + + if v := getEnv("TEST_DB_VAR"); v != "expected_value" { + t.Errorf("expected 'expected_value', got %q", v) + } +} + +func TestGetEnv_ReturnsEmptyWhenUnset(t *testing.T) { + os.Unsetenv("TEST_DB_MISSING_VAR") + + if v := getEnv("TEST_DB_MISSING_VAR"); v != "" { + t.Errorf("expected empty string, got %q", v) + } +} + +// ============================================================================= +// Clean-up — restore env after all tests +// ============================================================================= + +func TestMain(m *testing.M) { + resetEnv() + code := m.Run() + closePool() + os.Exit(code) +}