package postgres import ( "context" "sync" "testing" "time" "github.com/barkstack/pawsql/internal/config" "github.com/barkstack/pawsql/internal/router" ) func TestResolverDefersManagedDatabaseStartup(t *testing.T) { routes, err := router.NewStaticResolver([]config.DatabaseConfig{ {Name: "external", Upstream: "192.0.2.1:5432"}, {Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD"}}, }) if err != nil { t.Fatal(err) } ensurer := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}} resolver := NewResolver(routes, []config.DatabaseConfig{ {Name: "external", Upstream: "192.0.2.1:5432"}, {Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD"}}, }, ensurer) backend, err := resolver.ResolveDatabase(context.Background(), "external") if err != nil { t.Fatal(err) } if backend.Address != "192.0.2.1:5432" || ensurer.calls != 0 { t.Errorf("external backend = %#v, calls = %d", backend, ensurer.calls) } backend, err = resolver.ResolveDatabase(context.Background(), "managed") if err != nil { t.Fatal(err) } if backend.Address != "127.0.0.1:55432" || ensurer.calls != 1 { t.Errorf("managed backend = %#v, calls = %d", backend, ensurer.calls) } } func TestResolverSerializesManagedStartup(t *testing.T) { routes, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD"}}}) if err != nil { t.Fatal(err) } ensurer := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}, gate: make(chan struct{}), started: make(chan struct{})} resolver := NewResolver(routes, []config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD"}}}, ensurer) var wg sync.WaitGroup for range 2 { wg.Add(1) go func() { defer wg.Done() _, _ = resolver.ResolveDatabase(context.Background(), "managed") }() } <-ensurer.started close(ensurer.gate) wg.Wait() if ensurer.calls != 2 { t.Errorf("EnsureDatabase calls = %d, want 2 route resolutions", ensurer.calls) } if ensurer.maxActive != 1 { t.Errorf("concurrent EnsureDatabase calls = %d, want 1", ensurer.maxActive) } } func TestResolverStopsDatabaseAfterIdleTimeout(t *testing.T) { routes, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD", IdleTimeout: 20 * time.Millisecond}}}) if err != nil { t.Fatal(err) } controller := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}, stopped: make(chan string, 1)} resolver := NewResolver(routes, []config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD", IdleTimeout: 20 * time.Millisecond}}}, controller) if _, err := resolver.ResolveDatabase(context.Background(), "managed"); err != nil { t.Fatal(err) } resolver.ReleaseConnection("managed") select { case database := <-controller.stopped: if database != "managed" { t.Errorf("stopped database = %q", database) } case <-time.After(time.Second): t.Fatal("managed database was not stopped after idle timeout") } } func TestResolverCancelsIdleStopForNewConnection(t *testing.T) { routes, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD", IdleTimeout: 40 * time.Millisecond}}}) if err != nil { t.Fatal(err) } controller := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}, stopped: make(chan string, 1)} resolver := NewResolver(routes, []config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD", IdleTimeout: 40 * time.Millisecond}}}, controller) if _, err := resolver.ResolveDatabase(context.Background(), "managed"); err != nil { t.Fatal(err) } resolver.ReleaseConnection("managed") if _, err := resolver.ResolveDatabase(context.Background(), "managed"); err != nil { t.Fatal(err) } select { case database := <-controller.stopped: t.Fatalf("stopped active database %q", database) case <-time.After(80 * time.Millisecond): } resolver.ReleaseConnection("managed") select { case <-controller.stopped: case <-time.After(time.Second): t.Fatal("managed database was not stopped after final lease release") } } func TestResolverMetersTrafficAndStopsTrafficIdleSession(t *testing.T) { postgres := &config.PostgresConfig{ Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD", TrafficIdleTimeout: 100 * time.Millisecond, } routes, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "managed", Postgres: postgres}}) if err != nil { t.Fatal(err) } controller := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}, stopped: make(chan string, 1)} resolver := NewResolver(routes, []config.DatabaseConfig{{Name: "managed", Postgres: postgres}}, controller) if _, err := resolver.ResolveDatabase(context.Background(), "managed"); err != nil { t.Fatal(err) } time.Sleep(70 * time.Millisecond) resolver.RecordTraffic("managed", true, 7) resolver.RecordTraffic("managed", false, 11) stats, ok := resolver.TrafficStats("managed") if !ok || stats.ActiveConnections != 1 || stats.ClientBytes != 7 || stats.BackendBytes != 11 || stats.LastActivity.IsZero() { t.Errorf("TrafficStats() = %#v, %t", stats, ok) } select { case database := <-controller.stopped: t.Fatalf("stopped database %q despite recent traffic", database) case <-time.After(50 * time.Millisecond): } select { case database := <-controller.stopped: if database != "managed" { t.Errorf("stopped database = %q", database) } case <-time.After(time.Second): t.Fatal("managed database was not stopped after traffic idle timeout") } } type fakeEnsurer struct { mu sync.Mutex addresses map[string]string calls int active int maxActive int gate chan struct{} started chan struct{} startedOnce sync.Once stopped chan string } func (f *fakeEnsurer) EnsureDatabase(_ context.Context, database string, _ config.PostgresConfig) (string, error) { f.mu.Lock() f.calls++ f.active++ if f.active > f.maxActive { f.maxActive = f.active } if f.started != nil { f.startedOnce.Do(func() { close(f.started) }) } f.mu.Unlock() if f.gate != nil { <-f.gate } f.mu.Lock() f.active-- address := f.addresses[database] f.mu.Unlock() return address, nil } func (f *fakeEnsurer) StopDatabase(_ context.Context, database string) error { if f.stopped != nil { f.stopped <- database } return nil }