Files
Crussell/backend/db/db_test.go
T
2026-06-25 01:01:32 +01:00

185 lines
4.1 KiB
Go

//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")
if os.Getenv("POSTGRES_HOST") == "" {
os.Setenv("POSTGRES_HOST", "localhost")
}
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
// =============================================================================