//go:build test // +build test package db import ( "context" "fmt" "os" "sync" "testing" ) func resetEnv() { host := "localhost" if h := os.Getenv("POSTGRES_HOST"); h != "" { host = h } else if h := os.Getenv("TEST_DB_HOST"); h != "" { host = h } os.Setenv("POSTGRES_USER", "myuser") os.Setenv("POSTGRES_PASSWORD", "mypassword") os.Setenv("POSTGRES_HOST", host) os.Setenv("POSTGRES_DB", "crussell_test_db") } func closePool() { if Conn != nil { Conn.Pool().Close() Conn = nil } } var testDBName = "crussell_test_db" // ============================================================================= // 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 Conn == nil { t.Fatal("Conn 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() poolConn, err := Conn.Acquire(context.Background()) if err != nil { t.Fatalf("Acquire failed: %v", err) } defer poolConn.Release() var result int err = poolConn.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() poolConn, err := Conn.Acquire(context.Background()) if err != nil { errs <- fmt.Errorf("goroutine %d: acquire: %w", id, err) return } defer poolConn.Release() var result int err = poolConn.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 // =============================================================================