package postgres import ( "context" "encoding/json" "errors" "reflect" "strings" "testing" ) func TestReadinessSQLFreezesPrivilegeMatrixAndFunctionSignatures(t *testing.T) { wantTables := []string{ "assets", "generation_jobs", "usage_events", "projects", "image_templates", "platform_organizations", "platform_users", "platform_account_migrations", "platform_runtime_settings", "billing_price_rules", "billing_wallets", "billing_ledger", } for _, table := range wantTables { if !strings.Contains(ReadinessSQL, "('"+table+"',") { t.Errorf("ReadinessSQL missing table %q", table) } } if got := strings.Count(ReadinessSQL, "has_function_privilege("); got != 2 { t.Fatalf("has_function_privilege count = %d, want 2", got) } for _, signature := range []string{ "public.claim_generation_jobs(text,integer,integer)", "public.billing_post_wallet_entry(text,text,text,text,text,bigint,text,text,text,jsonb)", } { if !strings.Contains(ReadinessSQL, signature) { t.Errorf("ReadinessSQL missing function signature %q", signature) } } } func TestReadinessSQLRequiresPublicGenerationJobLifecycleColumns(t *testing.T) { for _, fragment := range []string{ "required_generation_job_columns(column_name)", "information_schema.columns", "table_schema = 'public'", "table_name = 'generation_jobs'", } { if !strings.Contains(ReadinessSQL, fragment) { t.Errorf("ReadinessSQL missing public generation_jobs column check %q", fragment) } } for _, column := range []string{ "provider_dispatch_started_at", "dispatch_ready_at", "finalized_at", } { if !strings.Contains(ReadinessSQL, "('"+column+"')") { t.Errorf("ReadinessSQL missing required lifecycle column %q", column) } } } func TestReadinessSQLRequiresBillingProviderMigrations(t *testing.T) { for _, fragment := range []string{ "billing_price_rules_provider_check", "pg_get_constraintdef", "position('seedream'", "position('minimax'", } { if !strings.Contains(ReadinessSQL, fragment) { t.Errorf("ReadinessSQL missing billing migration check %q", fragment) } } } func TestReadinessLocalSucceedsWithoutQuery(t *testing.T) { db := NewDatabase(Config{Backend: BackendLocal}, &fakeQuerier{err: errors.New("must not query")}) if err := db.Readiness(context.Background()); err != nil { t.Fatalf("Readiness() error = %v", err) } } func TestReadinessPostgresFailsWhenMatrixIsNotReady(t *testing.T) { q := &fakeQuerier{rows: [][]any{{false}}} db := NewDatabase(Config{Backend: BackendPostgres}, q) if err := db.Readiness(context.Background()); err == nil { t.Fatal("Readiness() error = nil, want not-ready error") } if q.sql != ReadinessSQL { t.Fatalf("query = %q, want exact ReadinessSQL", q.sql) } } func TestClaimGenerationJobsCallsExactFunction(t *testing.T) { const wantSQL = `SELECT id FROM public.claim_generation_jobs($1::text, $2::integer, $3::integer)` if ClaimGenerationJobsSQL != wantSQL { t.Fatalf("ClaimGenerationJobsSQL = %q, want %q", ClaimGenerationJobsSQL, wantSQL) } q := &fakeQuerier{rows: [][]any{{"job-1"}}} db := NewDatabase(Config{Backend: BackendPostgres}, q) jobs, err := db.ClaimGenerationJobs(context.Background(), "worker-1", 2, 300) if err != nil { t.Fatalf("ClaimGenerationJobs() error = %v", err) } if q.sql != ClaimGenerationJobsSQL || !reflect.DeepEqual(q.args, []any{"worker-1", 2, 300}) { t.Fatalf("query = %q args = %#v", q.sql, q.args) } if !reflect.DeepEqual(jobs, []GenerationJob{{ID: "job-1"}}) { t.Fatalf("jobs = %#v", jobs) } } func TestClaimGenerationJobsBoundsBatchLikeCurrentBackend(t *testing.T) { for _, test := range []struct { name string requested int want int }{ {name: "minimum", requested: 0, want: 1}, {name: "maximum", requested: 25, want: 20}, } { t.Run(test.name, func(t *testing.T) { q := &fakeQuerier{} db := NewDatabase(Config{Backend: BackendPostgres}, q) if _, err := db.ClaimGenerationJobs(context.Background(), "worker-1", test.requested, 300); err != nil { t.Fatalf("ClaimGenerationJobs() error = %v", err) } if !reflect.DeepEqual(q.args, []any{"worker-1", test.want, 300}) { t.Fatalf("args = %#v, want bounded limit %d", q.args, test.want) } }) } } func TestPostWalletEntryCallsExactFunction(t *testing.T) { const wantSQL = `SELECT ledger_id, balance_after_fen, balance_fen, total_recharged_fen, total_charged_fen, created_at, updated_at, delta_fen FROM public.billing_post_wallet_entry($1::text, $2::text, $3::text, $4::text, $5::text, $6::bigint, $7::text, $8::text, $9::text, $10::jsonb)` if PostWalletEntrySQL != wantSQL { t.Fatalf("PostWalletEntrySQL = %q, want %q", PostWalletEntrySQL, wantSQL) } q := &fakeQuerier{rows: [][]any{{"ledger-1", int64(120), int64(120), int64(200), int64(80), nil, nil, int64(-80)}}} db := NewDatabase(Config{Backend: BackendPostgres}, q) metadata := json.RawMessage(`{"source":"test"}`) entry, err := db.PostWalletEntry(context.Background(), WalletEntryParams{ LedgerID: "ledger-1", OrganizationID: "org-1", AccountID: "acct-1", JobID: "job-1", Kind: "charge", DeltaFen: -80, Currency: "CNY", IdempotencyKey: "idem-1", Description: "generation", Metadata: metadata, }) if err != nil { t.Fatalf("PostWalletEntry() error = %v", err) } if q.sql != PostWalletEntrySQL { t.Fatalf("query = %q, want exact PostWalletEntrySQL", q.sql) } wantArgs := []any{"ledger-1", "org-1", "acct-1", "job-1", "charge", int64(-80), "CNY", "idem-1", "generation", metadata} if !reflect.DeepEqual(q.args, wantArgs) { t.Fatalf("args = %#v, want %#v", q.args, wantArgs) } if entry.LedgerID != "ledger-1" || entry.BalanceFen != 120 { t.Fatalf("entry = %#v", entry) } } func TestPostWalletEntryNormalizesOptionalValuesLikeCurrentBackend(t *testing.T) { q := &fakeQuerier{rows: [][]any{{"ledger-1", int64(200), int64(200), int64(200), int64(0), nil, nil, int64(200)}}} db := NewDatabase(Config{Backend: BackendPostgres}, q) _, err := db.PostWalletEntry(context.Background(), WalletEntryParams{ LedgerID: "ledger-1", OrganizationID: "org-1", Kind: "recharge", DeltaFen: 200, IdempotencyKey: "idem-1", Description: "recharge", }) if err != nil { t.Fatalf("PostWalletEntry() error = %v", err) } wantArgs := []any{"ledger-1", "org-1", nil, nil, "recharge", int64(200), "CNY", "idem-1", "recharge", json.RawMessage(nil)} if !reflect.DeepEqual(q.args, wantArgs) { t.Fatalf("args = %#v, want %#v", q.args, wantArgs) } } type fakeQuerier struct { rows [][]any err error sql string args []any } func (q *fakeQuerier) Query(_ context.Context, sql string, args ...any) (Rows, error) { q.sql = sql q.args = args return &fakeRows{rows: q.rows, err: q.err}, nil } type fakeRows struct { rows [][]any idx int err error } func (r *fakeRows) Close() {} func (r *fakeRows) Err() error { return r.err } func (r *fakeRows) Next() bool { return r.idx < len(r.rows) } func (r *fakeRows) Scan(dest ...any) error { if r.idx >= len(r.rows) { return errors.New("scan past end") } row := r.rows[r.idx] r.idx++ if len(dest) != len(row) { return errors.New("scan arity mismatch") } for i := range dest { switch target := dest[i].(type) { case *bool: *target = row[i].(bool) case *string: *target = row[i].(string) case *int64: *target = row[i].(int64) default: // nil timestamp fixtures intentionally leave zero values. } } return nil }