package runner import ( "context" "io" "log/slog" "os" "strconv" "strings" "testing" "time" "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" "github.com/easyai/easyai-ai-gateway/apps/api/internal/config" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" ) func TestAsyncQueueClientEnqueuesWithoutExecutionWorker(t *testing.T) { databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL")) if databaseURL == "" { t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run the queue client integration test") } ctx, cancel := context.WithCancel(context.Background()) t.Cleanup(cancel) db, err := store.Connect(ctx, databaseURL) if err != nil { t.Fatalf("connect store: %v", err) } t.Cleanup(db.Close) suffix := strconv.FormatInt(time.Now().UnixNano(), 10) task, err := db.CreateTask(ctx, store.CreateTaskInput{ Kind: "images.edits", Model: "queue-client-" + suffix, Request: map[string]any{"prompt": "enqueue without execution worker"}, Async: true, RunMode: "simulation", }, &auth.User{ID: "queue-client-" + suffix, Source: "gateway"}) if err != nil { t.Fatalf("create task: %v", err) } t.Cleanup(func() { _, _ = db.Pool().Exec(context.Background(), `DELETE FROM gateway_tasks WHERE id = $1::uuid`, task.ID) }) service := New(config.Config{AppEnv: "test"}, db, slog.New(slog.NewTextHandler(io.Discard, nil))) service.StartAsyncQueueClient(ctx) if err := service.EnqueueAsyncTask(ctx, task); err != nil { t.Fatalf("enqueue task without execution worker: %v", err) } queued, err := db.GetTask(ctx, task.ID) if err != nil { t.Fatalf("load queued task: %v", err) } if queued.RiverJobID <= 0 { t.Fatal("queue client did not persist a River job ID") } t.Cleanup(func() { _, _ = db.Pool().Exec(context.Background(), `DELETE FROM river_job WHERE id = $1`, queued.RiverJobID) }) var attemptedByCount int if err := db.Pool().QueryRow(ctx, ` SELECT COALESCE(cardinality(attempted_by), 0) FROM river_job WHERE id = $1`, queued.RiverJobID).Scan(&attemptedByCount); err != nil { t.Fatalf("load River job: %v", err) } if attemptedByCount != 0 { t.Fatalf("control-only queue client executed the job: attempted_by=%d", attemptedByCount) } } func TestWaitingAsyncAdmissionDispatcherOwnsScheduledRiverJobs(t *testing.T) { databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL")) if databaseURL == "" { t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run the async admission dispatcher integration test") } ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() db, err := store.Connect(ctx, databaseURL) if err != nil { t.Fatalf("connect store: %v", err) } t.Cleanup(db.Close) var databaseName string if err := db.Pool().QueryRow(ctx, `SELECT current_database()`).Scan(&databaseName); err != nil { t.Fatalf("read test database name: %v", err) } if !strings.Contains(strings.ToLower(databaseName), "test") { t.Fatalf("refusing to use non-test database %q", databaseName) } suffix := strconv.FormatInt(time.Now().UnixNano(), 10) platform, err := db.CreatePlatform(ctx, store.CreatePlatformInput{ Provider: "async-dispatch-test", PlatformKey: "async-dispatch-test-" + suffix, Name: "Async Dispatch Test " + suffix, AuthType: "none", Status: "enabled", }) if err != nil { t.Fatalf("create platform: %v", err) } modelName := "async-dispatch-model-" + suffix var platformModelID string if err := db.Pool().QueryRow(ctx, ` INSERT INTO platform_models ( platform_id, model_name, provider_model_name, model_alias, model_type, display_name, pricing_mode, enabled ) VALUES ($1::uuid, $2, $2, $2, '["image_generate"]'::jsonb, $2, 'inherit_discount', true) RETURNING id::text`, platform.ID, modelName).Scan(&platformModelID); err != nil { t.Fatalf("create platform model: %v", err) } task, err := db.CreateTask(ctx, store.CreateTaskInput{ Kind: "images.edits", Model: modelName, Request: map[string]any{"prompt": "scheduled River job re-admission"}, Async: true, RunMode: "simulation", }, &auth.User{ID: "async-dispatch-user-" + suffix, Source: "gateway"}) if err != nil { t.Fatalf("create task: %v", err) } t.Cleanup(func() { cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), 10*time.Second) defer cleanupCancel() _ = db.DeleteTaskAdmission(cleanupCtx, task.ID) _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM river_job WHERE id = (SELECT river_job_id FROM gateway_tasks WHERE id = $1::uuid)`, task.ID) _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM gateway_tasks WHERE id = $1::uuid`, task.ID) _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM platform_models WHERE id = $1::uuid`, platformModelID) _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM integration_platforms WHERE id = $1::uuid`, platform.ID) }) service := New(config.Config{AppEnv: "test"}, db, slog.New(slog.NewTextHandler(io.Discard, nil))) service.StartAsyncQueueClient(ctx) if err := service.EnqueueAsyncTask(ctx, task); err != nil { t.Fatalf("enqueue task: %v", err) } queued, err := db.GetTask(ctx, task.ID) if err != nil { t.Fatalf("load queued task: %v", err) } if queued.RiverJobID <= 0 { t.Fatal("queued task did not persist a River job ID") } admissionInput := store.TaskAdmissionInput{ TaskID: task.ID, PlatformID: platform.ID, PlatformModelID: platformModelID, QueueKey: "async-dispatch-test:" + suffix, Mode: "async", Priority: 100, Scopes: []store.AdmissionScope{{ ScopeType: "worker_capacity", ScopeKey: "global", ScopeName: "Worker capacity", ConcurrentLimit: 1, Amount: 1, LeaseTTLSeconds: 120, }}, } if _, err := db.QueueTaskAdmissionWithHook(ctx, admissionInput, nil); err != nil { t.Fatalf("queue task admission: %v", err) } assertDispatchable := func(want bool) { t.Helper() taskIDs, listErr := db.ListWaitingAsyncAdmissionTaskIDs(ctx, 1000) if listErr != nil { t.Fatalf("list waiting async admissions: %v", listErr) } found := false for _, taskID := range taskIDs { if taskID == task.ID { found = true break } } if found != want { t.Fatalf("task dispatchable=%v, want %v", found, want) } } assertDispatchable(true) if _, err := db.Pool().Exec(ctx, `UPDATE river_job SET state = 'running' WHERE id = $1`, queued.RiverJobID); err != nil { t.Fatalf("mark River job running: %v", err) } assertDispatchable(false) if _, err := db.Pool().Exec(ctx, `UPDATE river_job SET state = 'scheduled' WHERE id = $1`, queued.RiverJobID); err != nil { t.Fatalf("snooze River job: %v", err) } assertDispatchable(true) admitted, err := db.TryTaskAdmission(ctx, admissionInput) if err != nil || !admitted.Admitted || len(admitted.Leases) != 1 { t.Fatalf("admit scheduled task: result=%+v err=%v", admitted, err) } assertDispatchable(false) if _, err := db.Pool().Exec(ctx, ` UPDATE gateway_concurrency_leases SET expires_at = now() - interval '1 second' WHERE task_id = $1::uuid AND released_at IS NULL`, task.ID); err != nil { t.Fatalf("expire admission lease: %v", err) } assertDispatchable(true) } func TestRecoverOrphanedAsyncRiverJobAndYieldStaleAdmission(t *testing.T) { databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL")) if databaseURL == "" { t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run the orphaned River job integration test") } ctx, cancel := context.WithCancel(context.Background()) t.Cleanup(cancel) db, err := store.Connect(ctx, databaseURL) if err != nil { t.Fatalf("connect store: %v", err) } t.Cleanup(db.Close) service := New(config.Config{AppEnv: "test"}, db, slog.New(slog.NewTextHandler(io.Discard, nil))) service.StartAsyncQueueClient(ctx) suffix := strconv.FormatInt(time.Now().UnixNano(), 10) createQueuedTask := func(label string) store.GatewayTask { t.Helper() task, createErr := db.CreateTask(ctx, store.CreateTaskInput{ Kind: "images.edits", Model: "orphan-recovery-" + label + "-" + suffix, Request: map[string]any{"prompt": "recover orphaned River job"}, Async: true, RunMode: "simulation", }, &auth.User{ID: "orphan-recovery-" + suffix, Source: "gateway"}) if createErr != nil { t.Fatalf("create %s task: %v", label, createErr) } if enqueueErr := service.EnqueueAsyncTask(ctx, task); enqueueErr != nil { t.Fatalf("enqueue %s task: %v", label, enqueueErr) } queued, getErr := db.GetTask(ctx, task.ID) if getErr != nil { t.Fatalf("load %s task: %v", label, getErr) } return queued } orphaned := createQueuedTask("orphaned") protected := createQueuedTask("protected") yielded := createQueuedTask("yielded") staleWorkerID := "orphan-recovery-stale-" + suffix activeWorkerID := "orphan-recovery-active-" + suffix t.Cleanup(func() { cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), 10*time.Second) defer cleanupCancel() _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM gateway_worker_instances WHERE instance_id = ANY($1::text[])`, []string{staleWorkerID, activeWorkerID}) _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM river_job WHERE id = ANY($1::bigint[])`, []int64{orphaned.RiverJobID, protected.RiverJobID, yielded.RiverJobID}) _, _ = db.Pool().Exec(cleanupCtx, `DELETE FROM gateway_tasks WHERE id = ANY($1::uuid[])`, []string{orphaned.ID, protected.ID, yielded.ID}) }) if _, err := db.Pool().Exec(ctx, ` INSERT INTO gateway_worker_instances ( instance_id, status, desired_capacity, allocated_capacity, heartbeat_at, updated_at ) VALUES ($1, 'active', 1, 0, now() - interval '2 minutes', now() - interval '2 minutes'), ($2, 'active', 1, 1, now(), now())`, staleWorkerID, activeWorkerID, ); err != nil { t.Fatalf("seed worker heartbeats: %v", err) } if _, err := db.Pool().Exec(ctx, ` UPDATE river_job SET state = 'running', attempt = 1, attempted_at = now() - interval '1 minute', attempted_by = ARRAY[ CASE id WHEN $1::bigint THEN $3 WHEN $2::bigint THEN $4 WHEN $5::bigint THEN $4 END ]::text[] WHERE id = ANY($6::bigint[])`, orphaned.RiverJobID, protected.RiverJobID, staleWorkerID+"-exec-1-test", activeWorkerID+"-exec-1-test", yielded.RiverJobID, []int64{orphaned.RiverJobID, protected.RiverJobID, yielded.RiverJobID}, ); err != nil { t.Fatalf("mark River jobs running: %v", err) } var platformID, platformModelID string if err := db.Pool().QueryRow(ctx, ` SELECT platform.id::text, model.id::text FROM integration_platforms platform JOIN platform_models model ON model.platform_id = platform.id ORDER BY platform.priority ASC, model.created_at ASC LIMIT 1`).Scan(&platformID, &platformModelID); err != nil { t.Fatalf("load admission binding: %v", err) } if _, err := db.Pool().Exec(ctx, ` UPDATE gateway_tasks SET execution_lease_expires_at = now() - interval '1 minute' WHERE id = $1::uuid`, yielded.ID, ); err != nil { t.Fatalf("expire yielded task execution lease: %v", err) } if _, err := db.Pool().Exec(ctx, ` UPDATE gateway_tasks SET execution_lease_expires_at = now() + interval '5 minutes' WHERE id = $1::uuid`, protected.ID, ); err != nil { t.Fatalf("renew protected task execution lease: %v", err) } if _, err := db.Pool().Exec(ctx, ` INSERT INTO gateway_task_admissions ( task_id, platform_id, platform_model_id, queue_key, mode, status, priority, enqueued_at, wait_deadline_at ) VALUES ( $1::uuid, $2::uuid, $3::uuid, 'integration-test', 'async', 'waiting', 100, now() - interval '1 minute', now() + interval '10 minutes' )`, yielded.ID, platformID, platformModelID, ); err != nil { t.Fatalf("seed stale async admission: %v", err) } if _, err := db.Pool().Exec(ctx, ` INSERT INTO gateway_task_admissions ( task_id, platform_id, platform_model_id, queue_key, mode, status, priority, enqueued_at, wait_deadline_at ) VALUES ( $1::uuid, $2::uuid, $3::uuid, 'integration-test', 'async', 'waiting', 100, now(), now() + interval '10 minutes' )`, protected.ID, platformID, platformModelID, ); err != nil { t.Fatalf("seed protected async admission: %v", err) } yieldedCount, err := db.YieldStaleAsyncTaskAdmissions(ctx, 30*time.Second, 10) if err != nil { t.Fatalf("yield stale async task admission: %v", err) } if yieldedCount != 1 { t.Fatalf("yielded admissions=%d, want 1", yieldedCount) } var yieldedAdmissions, protectedAdmissions int if err := db.Pool().QueryRow(ctx, ` SELECT count(*) FILTER (WHERE task_id = $1::uuid), count(*) FILTER (WHERE task_id = $2::uuid) FROM gateway_task_admissions WHERE task_id = ANY($3::uuid[])`, yielded.ID, protected.ID, []string{yielded.ID, protected.ID}, ).Scan(&yieldedAdmissions, &protectedAdmissions); err != nil { t.Fatalf("count yielded admissions: %v", err) } var yieldedState string if err := db.Pool().QueryRow(ctx, ` SELECT state::text FROM river_job WHERE id = $1`, yielded.RiverJobID).Scan(&yieldedState); err != nil { t.Fatalf("read yielded River job state: %v", err) } if yieldedAdmissions != 0 || protectedAdmissions != 1 || yieldedState != "running" { t.Fatalf( "yielded/protected admissions=%d/%d River state=%s, want 0/1/running", yieldedAdmissions, protectedAdmissions, yieldedState, ) } recovered, err := db.RecoverOrphanedAsyncRiverJobs(ctx, 30*time.Second, 10) if err != nil { t.Fatalf("recover orphaned River jobs: %v", err) } if recovered != 2 { t.Fatalf("recovered jobs=%d, want 2", recovered) } var orphanedState, protectedState, yieldedRecoveredState string if err := db.Pool().QueryRow(ctx, `SELECT state::text FROM river_job WHERE id = $1`, orphaned.RiverJobID).Scan(&orphanedState); err != nil { t.Fatalf("read orphaned River job state: %v", err) } if err := db.Pool().QueryRow(ctx, `SELECT state::text FROM river_job WHERE id = $1`, protected.RiverJobID).Scan(&protectedState); err != nil { t.Fatalf("read protected River job state: %v", err) } if err := db.Pool().QueryRow(ctx, `SELECT state::text FROM river_job WHERE id = $1`, yielded.RiverJobID).Scan(&yieldedRecoveredState); err != nil { t.Fatalf("read yielded River job state: %v", err) } if orphanedState != "retryable" || protectedState != "running" || yieldedRecoveredState != "retryable" { t.Fatalf( "River states orphaned=%s protected=%s yielded=%s, want retryable/running/retryable", orphanedState, protectedState, yieldedRecoveredState, ) } recovered, err = db.RecoverOrphanedAsyncRiverJobs(ctx, 30*time.Second, 10) if err != nil || recovered != 0 { t.Fatalf("second orphan recovery count=%d err=%v, want idempotent zero", recovered, err) } }