diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 4af4d9bccd..ed56fb894b 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -213,12 +213,17 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { soraMediaStorage := service.ProvideSoraMediaStorage(configConfig) soraGatewayService := service.NewSoraGatewayService(soraSDKClient, rateLimitService, httpUpstream, configConfig) soraClientHandler := handler.NewSoraClientHandler(soraGenerationService, soraQuotaService, soraStorageRouter, soraGatewayService, gatewayService, soraMediaStorage, apiKeyService) + soraTaskRepository := repository.NewSoraTaskRepository(db) + soraTaskService := service.NewSoraTaskService(soraTaskRepository, soraSDKClient, httpUpstream) + soraTaskWorker := service.NewSoraTaskWorker(soraTaskService, accountRepository, soraStorageRouter, soraMediaStorage, 60*time.Second) + soraTaskWorker.Start() + soraVideosHandler := handler.NewSoraVideosHandler(soraTaskService, gatewayService, soraStorageRouter, soraMediaStorage, soraGatewayService) soraGatewayHandler := handler.NewSoraGatewayHandler(gatewayService, soraGatewayService, concurrencyService, billingCacheService, usageRecordWorkerPool, configConfig) handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo) totpHandler := handler.NewTotpHandler(totpService) idempotencyCoordinator := service.ProvideIdempotencyCoordinator(idempotencyRepository, configConfig) idempotencyCleanupService := service.ProvideIdempotencyCleanupService(idempotencyRepository, configConfig) - handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, soraGatewayHandler, soraClientHandler, handlerSettingHandler, totpHandler, idempotencyCoordinator, idempotencyCleanupService) + handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, soraGatewayHandler, soraClientHandler, soraVideosHandler, handlerSettingHandler, totpHandler, idempotencyCoordinator, idempotencyCleanupService) jwtAuthMiddleware := middleware.NewJWTAuthMiddleware(authService, userService) adminAuthMiddleware := middleware.NewAdminAuthMiddleware(authService, userService, settingService) apiKeyAuthMiddleware := middleware.NewAPIKeyAuthMiddleware(apiKeyService, subscriptionService, configConfig) @@ -234,7 +239,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { accountExpiryService := service.ProvideAccountExpiryService(accountRepository) subscriptionExpiryService := service.ProvideSubscriptionExpiryService(userSubscriptionRepository) scheduledTestRunnerService := service.ProvideScheduledTestRunnerService(scheduledTestPlanRepository, scheduledTestService, accountTestService, configConfig) - v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, soraMediaCleanupService, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService) + v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, soraMediaCleanupService, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, soraTaskWorker) application := &Application{ Server: httpServer, Cleanup: v, @@ -283,6 +288,7 @@ func provideCleanup( antigravityOAuth *service.AntigravityOAuthService, openAIGateway *service.OpenAIGatewayService, scheduledTestRunner *service.ScheduledTestRunnerService, + soraTaskWorker *service.SoraTaskWorker, ) func() { return func() { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -418,6 +424,12 @@ func provideCleanup( } return nil }}, + {"SoraTaskWorker", func() error { + if soraTaskWorker != nil { + soraTaskWorker.Stop() + } + return nil + }}, } infraSteps := []cleanupStep{ diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go index 8e203cb81c..cd50c4020e 100644 --- a/backend/cmd/server/wire_gen_test.go +++ b/backend/cmd/server/wire_gen_test.go @@ -75,6 +75,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) { antigravityOAuthSvc, nil, // openAIGateway nil, // scheduledTestRunner + nil, // soraTaskWorker ) require.NotPanics(t, func() { diff --git a/backend/internal/handler/handler.go b/backend/internal/handler/handler.go index f908800c67..37b3ab474b 100644 --- a/backend/internal/handler/handler.go +++ b/backend/internal/handler/handler.go @@ -45,6 +45,7 @@ type Handlers struct { OpenAIGateway *OpenAIGatewayHandler SoraGateway *SoraGatewayHandler SoraClient *SoraClientHandler + SoraVideos *SoraVideosHandler Setting *SettingHandler Totp *TotpHandler } diff --git a/backend/internal/handler/sora_videos_handler.go b/backend/internal/handler/sora_videos_handler.go new file mode 100644 index 0000000000..c75b1ce91f --- /dev/null +++ b/backend/internal/handler/sora_videos_handler.go @@ -0,0 +1,379 @@ +package handler + +import ( + "encoding/json" + "net/http" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + + "github.com/gin-gonic/gin" +) + +// SoraVideosHandler handles Sora video/image async task API. +type SoraVideosHandler struct { + taskService *service.SoraTaskService + gatewayService *service.GatewayService + objectStorage *service.SoraS3Storage + mediaStorage *service.SoraMediaStorage + soraGatewayService *service.SoraGatewayService +} + +func NewSoraVideosHandler( + taskService *service.SoraTaskService, + gatewayService *service.GatewayService, + objectStorage *service.SoraS3Storage, + mediaStorage *service.SoraMediaStorage, + soraGatewayService *service.SoraGatewayService, +) *SoraVideosHandler { + if taskService == nil { + return nil + } + return &SoraVideosHandler{ + taskService: taskService, + gatewayService: gatewayService, + objectStorage: objectStorage, + mediaStorage: mediaStorage, + soraGatewayService: soraGatewayService, + } +} + +func (h *SoraVideosHandler) CreateVideo(c *gin.Context) { + apiKey, account, ok := h.selectAccount(c, "") + if !ok { + return + } + + body, err := readBody(c) + if err != nil { + return + } + + var req service.CreateVideoRequest + if err := json.Unmarshal(body, &req); err != nil { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body") + return + } + if req.Model == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "model is required") + return + } + + task, err := h.taskService.CreateVideoTask(c.Request.Context(), apiKey.ID, account, &req, body) + if err != nil { + logger.LegacyPrintf("handler.sora_videos", "[CreateVideo] error: %v", err) + soraErrorResponse(c, http.StatusInternalServerError, "server_error", "Failed to create video task") + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +func (h *SoraVideosHandler) GetVideo(c *gin.Context) { + apiKey, ok := h.getAPIKey(c) + if !ok { + return + } + + taskID := c.Param("id") + if taskID == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "video_id is required") + return + } + + task, err := h.taskService.GetTask(c.Request.Context(), taskID, apiKey.ID) + if err != nil { + soraErrorResponse(c, http.StatusNotFound, "not_found", "Task not found") + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +func (h *SoraVideosHandler) RemixVideo(c *gin.Context) { + apiKey, ok := h.getAPIKey(c) + if !ok { + return + } + + taskID := c.Param("id") + if taskID == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "video_id is required") + return + } + + body, err := readBody(c) + if err != nil { + return + } + + var req service.RemixRequest + if err := json.Unmarshal(body, &req); err != nil || req.Prompt == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "prompt is required") + return + } + + originalTask, err := h.taskService.GetTask(c.Request.Context(), taskID, apiKey.ID) + if err != nil { + soraErrorResponse(c, http.StatusNotFound, "not_found", "Original video task not found") + return + } + if originalTask.Status != service.SoraTaskCompleted { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Original video must be completed before remix") + return + } + + remixTargetID := originalTask.ShareID + if remixTargetID == "" { + remixTargetID = originalTask.UpstreamTaskID + } + + account, err := h.selectAccountByID(c, originalTask.AccountID) + if err != nil { + soraErrorResponse(c, http.StatusServiceUnavailable, "server_error", "Failed to get account") + return + } + + videoReq := &service.CreateVideoRequest{ + Model: originalTask.Model, + Prompt: req.Prompt, + RemixTargetID: remixTargetID, + } + reqBody, _ := json.Marshal(videoReq) + + task, err := h.taskService.CreateVideoTask(c.Request.Context(), apiKey.ID, account, videoReq, reqBody) + if err != nil { + logger.LegacyPrintf("handler.sora_videos", "[RemixVideo] error: %v", err) + soraErrorResponse(c, http.StatusInternalServerError, "server_error", "Failed to create remix task") + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +// GetVideoContent returns video content based on storage configuration. +func (h *SoraVideosHandler) GetVideoContent(c *gin.Context) { + apiKey, ok := h.getAPIKey(c) + if !ok { + return + } + + taskID := c.Param("id") + if taskID == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "video_id is required") + return + } + + task, err := h.taskService.GetTask(c.Request.Context(), taskID, apiKey.ID) + if err != nil { + soraErrorResponse(c, http.StatusNotFound, "not_found", "Task not found") + return + } + + switch task.Status { + case service.SoraTaskCompleted: + contentURL := h.resolveContentURL(c, task) + if contentURL == "" { + soraErrorResponse(c, http.StatusNotFound, "not_found", "Video URL not available") + return + } + c.Redirect(http.StatusFound, contentURL) + + case service.SoraTaskFailed: + c.JSON(http.StatusGone, gin.H{ + "id": task.ID, + "object": task.ObjectType, + "status": task.Status, + "error": gin.H{ + "message": task.ErrorMessage, + "type": task.ErrorType, + }, + }) + + default: + c.JSON(http.StatusAccepted, service.TaskToResponse(task)) + } +} + +func (h *SoraVideosHandler) resolveContentURL(c *gin.Context, task *service.SoraTask) string { + if task.StoredKey == "" { + return task.VideoURL + } + + switch task.StorageType { + case "s3", "gdrive": + if h.objectStorage != nil { + accessURL, err := h.objectStorage.GetAccessURL(c.Request.Context(), task.StoredKey) + if err != nil { + logger.LegacyPrintf("handler.sora_videos", + "[GetVideoContent] task=%s get access URL error: %v, fallback to upstream", task.ID, err) + return task.VideoURL + } + return accessURL + } + return task.VideoURL + + case "local": + return "/sora/media" + task.StoredKey + + default: + return task.VideoURL + } +} + +func (h *SoraVideosHandler) CreateImage(c *gin.Context) { + apiKey, account, ok := h.selectAccount(c, "") + if !ok { + return + } + + body, err := readBody(c) + if err != nil { + return + } + + var req service.CreateImageRequest + if err := json.Unmarshal(body, &req); err != nil { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body") + return + } + if req.Prompt == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "prompt is required") + return + } + if req.Model == "" { + req.Model = inferImageModel(req.Size) + } + + task, err := h.taskService.CreateImageGeneration(c.Request.Context(), apiKey.ID, account, &req, body) + if err != nil { + logger.LegacyPrintf("handler.sora_videos", "[CreateImage] error: %v", err) + soraErrorResponse(c, http.StatusInternalServerError, "server_error", "Failed to create image task") + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +func (h *SoraVideosHandler) EditImage(c *gin.Context) { + apiKey, account, ok := h.selectAccount(c, "") + if !ok { + return + } + + body, err := readBody(c) + if err != nil { + return + } + + var req service.EditImageRequest + if err := json.Unmarshal(body, &req); err != nil { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Invalid request body") + return + } + if req.Image == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "image is required") + return + } + if req.Prompt == "" { + soraErrorResponse(c, http.StatusBadRequest, "parameter_missing", "prompt is required") + return + } + if req.Model == "" { + req.Model = inferImageModel(req.Size) + } + + imageReq := &service.CreateImageRequest{ + Model: req.Model, + Prompt: req.Prompt, + Size: req.Size, + ResponseFormat: req.ResponseFormat, + N: 1, + } + + task, err := h.taskService.CreateImageGeneration(c.Request.Context(), apiKey.ID, account, imageReq, body) + if err != nil { + logger.LegacyPrintf("handler.sora_videos", "[EditImage] error: %v", err) + soraErrorResponse(c, http.StatusInternalServerError, "server_error", "Failed to create image edit task") + return + } + + c.JSON(http.StatusOK, service.TaskToResponse(task)) +} + +// ── Internal helpers ── + +func (h *SoraVideosHandler) getAPIKey(c *gin.Context) (*service.APIKey, bool) { + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok { + soraErrorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key") + return nil, false + } + return apiKey, true +} + +func (h *SoraVideosHandler) selectAccount(c *gin.Context, model string) (*service.APIKey, *service.Account, bool) { + apiKey, ok := h.getAPIKey(c) + if !ok { + return nil, nil, false + } + + selection, err := h.gatewayService.SelectAccountWithLoadAwareness( + c.Request.Context(), apiKey.GroupID, "", model, nil, "", + ) + if err != nil { + soraErrorResponse(c, http.StatusServiceUnavailable, "server_error", "No available accounts") + return nil, nil, false + } + if selection.ReleaseFunc != nil { + defer selection.ReleaseFunc() + } + return apiKey, selection.Account, true +} + +func (h *SoraVideosHandler) selectAccountByID(c *gin.Context, _ int64) (*service.Account, error) { + selection, err := h.gatewayService.SelectAccountWithLoadAwareness( + c.Request.Context(), nil, "", "", nil, "", + ) + if err != nil { + return nil, err + } + if selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + return selection.Account, nil +} + +func readBody(c *gin.Context) ([]byte, error) { + body, err := c.GetRawData() + if err != nil { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body") + return nil, err + } + if len(body) == 0 { + soraErrorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty") + return nil, err + } + return body, nil +} + +func soraErrorResponse(c *gin.Context, status int, errType, message string) { + c.JSON(status, gin.H{ + "error": gin.H{ + "message": message, + "type": errType, + }, + }) +} + +func inferImageModel(size string) string { + switch size { + case "540x360": + return "gpt-image-landscape" + case "360x540": + return "gpt-image-portrait" + default: + return "gpt-image" + } +} diff --git a/backend/internal/handler/wire.go b/backend/internal/handler/wire.go index 7916af6e43..88c211887e 100644 --- a/backend/internal/handler/wire.go +++ b/backend/internal/handler/wire.go @@ -84,6 +84,7 @@ func ProvideHandlers( openaiGatewayHandler *OpenAIGatewayHandler, soraGatewayHandler *SoraGatewayHandler, soraClientHandler *SoraClientHandler, + soraVideosHandler *SoraVideosHandler, settingHandler *SettingHandler, totpHandler *TotpHandler, _ *service.IdempotencyCoordinator, @@ -102,6 +103,7 @@ func ProvideHandlers( OpenAIGateway: openaiGatewayHandler, SoraGateway: soraGatewayHandler, SoraClient: soraClientHandler, + SoraVideos: soraVideosHandler, Setting: settingHandler, Totp: totpHandler, } diff --git a/backend/internal/repository/sora_task_repo.go b/backend/internal/repository/sora_task_repo.go new file mode 100644 index 0000000000..49a101b139 --- /dev/null +++ b/backend/internal/repository/sora_task_repo.go @@ -0,0 +1,131 @@ +package repository + +import ( + "context" + "database/sql" + "encoding/json" + + "github.com/Wei-Shaw/sub2api/internal/service" +) + +type SoraTaskRepository struct { + db *sql.DB +} + +func NewSoraTaskRepository(sqlDB *sql.DB) service.SoraTaskRepository { + return &SoraTaskRepository{db: sqlDB} +} + +func (r *SoraTaskRepository) Create(ctx context.Context, task *service.SoraTask) error { + charJSON, _ := json.Marshal(task.CharacterInfo) + if task.CharacterInfo == nil { + charJSON = nil + } + var reqBody []byte + if len(task.RequestBody) > 0 { + reqBody = task.RequestBody + } + + _, err := r.db.ExecContext(ctx, ` + INSERT INTO sora_tasks ( + id, account_id, api_key_id, upstream_task_id, object_type, + model, prompt, status, progress, video_url, stored_key, storage_type, + share_id, character_info, error_message, error_type, + request_body, seconds, size, created_at, completed_at + ) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21)`, + task.ID, task.AccountID, task.APIKeyID, task.UpstreamTaskID, task.ObjectType, + task.Model, task.Prompt, task.Status, task.Progress, task.VideoURL, + task.StoredKey, task.StorageType, + task.ShareID, charJSON, task.ErrorMessage, task.ErrorType, + reqBody, task.Seconds, task.Size, task.CreatedAt, task.CompletedAt, + ) + return err +} + +func (r *SoraTaskRepository) GetByID(ctx context.Context, id string) (*service.SoraTask, error) { + row := r.db.QueryRowContext(ctx, `SELECT `+soraTaskColumns+` FROM sora_tasks WHERE id = $1`, id) + return scanTask(row) +} + +func (r *SoraTaskRepository) GetByIDAndAPIKey(ctx context.Context, id string, apiKeyID int64) (*service.SoraTask, error) { + row := r.db.QueryRowContext(ctx, + `SELECT `+soraTaskColumns+` FROM sora_tasks WHERE id = $1 AND api_key_id = $2`, + id, apiKeyID, + ) + return scanTask(row) +} + +func (r *SoraTaskRepository) Update(ctx context.Context, task *service.SoraTask) error { + charJSON, _ := json.Marshal(task.CharacterInfo) + if task.CharacterInfo == nil { + charJSON = nil + } + + _, err := r.db.ExecContext(ctx, ` + UPDATE sora_tasks SET + upstream_task_id = $2, status = $3, progress = $4, + video_url = $5, stored_key = $6, storage_type = $7, + share_id = $8, character_info = $9, + error_message = $10, error_type = $11, + completed_at = $12, seconds = $13, size = $14, object_type = $15 + WHERE id = $1`, + task.ID, task.UpstreamTaskID, task.Status, task.Progress, + task.VideoURL, task.StoredKey, task.StorageType, + task.ShareID, charJSON, + task.ErrorMessage, task.ErrorType, + task.CompletedAt, task.Seconds, task.Size, task.ObjectType, + ) + return err +} + +func (r *SoraTaskRepository) ListPending(ctx context.Context) ([]*service.SoraTask, error) { + rows, err := r.db.QueryContext(ctx, + `SELECT `+soraTaskColumns+` FROM sora_tasks WHERE status IN ('queued', 'in_progress') ORDER BY created_at ASC`, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + var tasks []*service.SoraTask + for rows.Next() { + t, err := scanTaskFromRow(rows) + if err != nil { + return nil, err + } + tasks = append(tasks, t) + } + return tasks, rows.Err() +} + +const soraTaskColumns = `id, account_id, api_key_id, upstream_task_id, object_type, + model, prompt, status, progress, video_url, stored_key, storage_type, + share_id, character_info, error_message, error_type, + request_body, seconds, size, created_at, completed_at` + +func scanTask(s scannable) (*service.SoraTask, error) { + var t service.SoraTask + var charJSON, reqBody []byte + err := s.Scan( + &t.ID, &t.AccountID, &t.APIKeyID, &t.UpstreamTaskID, &t.ObjectType, + &t.Model, &t.Prompt, &t.Status, &t.Progress, &t.VideoURL, + &t.StoredKey, &t.StorageType, + &t.ShareID, &charJSON, &t.ErrorMessage, &t.ErrorType, + &reqBody, &t.Seconds, &t.Size, &t.CreatedAt, &t.CompletedAt, + ) + if err != nil { + return nil, err + } + if len(charJSON) > 0 { + var ch service.SoraCharacter + if json.Unmarshal(charJSON, &ch) == nil && ch.Username != "" { + t.CharacterInfo = &ch + } + } + t.RequestBody = reqBody + return &t, nil +} + +func scanTaskFromRow(rows *sql.Rows) (*service.SoraTask, error) { + return scanTask(rows) +} diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index ade2ed8337..86ef38e0fb 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -143,6 +143,16 @@ func RegisterGatewayRoutes( { soraV1.POST("/chat/completions", h.SoraGateway.ChatCompletions) soraV1.GET("/models", h.Gateway.Models) + + // Sora Videos/Images async task API + if h.SoraVideos != nil { + soraV1.POST("/videos", h.SoraVideos.CreateVideo) + soraV1.GET("/videos/:id", h.SoraVideos.GetVideo) + soraV1.POST("/videos/:id/remix", h.SoraVideos.RemixVideo) + soraV1.GET("/videos/:id/content", h.SoraVideos.GetVideoContent) + soraV1.POST("/images/generations", h.SoraVideos.CreateImage) + soraV1.POST("/images/edits", h.SoraVideos.EditImage) + } } // Sora 媒体代理(可选 API Key 验证) diff --git a/backend/internal/service/sora_task.go b/backend/internal/service/sora_task.go new file mode 100644 index 0000000000..6b909b9b2b --- /dev/null +++ b/backend/internal/service/sora_task.go @@ -0,0 +1,63 @@ +package service + +import ( + "context" + "time" +) + +// ── 任务状态常量 ── + +const ( + SoraTaskQueued = "queued" + SoraTaskInProgress = "in_progress" + SoraTaskCompleted = "completed" + SoraTaskFailed = "failed" +) + +// ── 对象类型常量 ── + +const ( + SoraObjectVideo = "video" + SoraObjectCharacter = "character" + SoraObjectImage = "image" +) + +// SoraTask 表示一个 Sora 异步任务记录。 +type SoraTask struct { + ID string + AccountID int64 + APIKeyID *int64 + UpstreamTaskID string + ObjectType string + Model string + Prompt string + Status string + Progress int + VideoURL string // 原始上游 URL + StoredKey string // 存储后的 key(本地路径或 S3 key) + StorageType string // local / s3 / gdrive / 空 + ShareID string + CharacterInfo *SoraCharacter + ErrorMessage string + ErrorType string + RequestBody []byte + Seconds string + Size string + CreatedAt time.Time + CompletedAt *time.Time +} + +// SoraCharacter 角色信息。 +type SoraCharacter struct { + Username string `json:"username"` + DisplayName string `json:"display_name"` +} + +// SoraTaskRepository 持久层接口。 +type SoraTaskRepository interface { + Create(ctx context.Context, task *SoraTask) error + GetByID(ctx context.Context, id string) (*SoraTask, error) + GetByIDAndAPIKey(ctx context.Context, id string, apiKeyID int64) (*SoraTask, error) + Update(ctx context.Context, task *SoraTask) error + ListPending(ctx context.Context) ([]*SoraTask, error) +} diff --git a/backend/internal/service/sora_task_service.go b/backend/internal/service/sora_task_service.go new file mode 100644 index 0000000000..7983345722 --- /dev/null +++ b/backend/internal/service/sora_task_service.go @@ -0,0 +1,528 @@ +package service + +import ( + "bytes" + "context" + "crypto/rand" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" +) + +// SoraTaskService manages Sora async task creation and status queries. +type SoraTaskService struct { + repo SoraTaskRepository + soraClient SoraClient + httpUpstream HTTPUpstream +} + +func NewSoraTaskService( + repo SoraTaskRepository, + soraClient SoraClient, + httpUpstream HTTPUpstream, +) *SoraTaskService { + return &SoraTaskService{ + repo: repo, + soraClient: soraClient, + httpUpstream: httpUpstream, + } +} + +// ── Request structs ── + +type CreateVideoRequest struct { + Model string `json:"model"` + Prompt string `json:"prompt"` + StyleID string `json:"style_id,omitempty"` + Orientation string `json:"orientation,omitempty"` + Image string `json:"image,omitempty"` + Video string `json:"video,omitempty"` + RemixTargetID string `json:"remix_target_id,omitempty"` + Size string `json:"size,omitempty"` +} + +type CreateImageRequest struct { + Model string `json:"model,omitempty"` + Prompt string `json:"prompt"` + Size string `json:"size,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` + N int `json:"n,omitempty"` +} + +type EditImageRequest struct { + Image string `json:"image"` + Prompt string `json:"prompt"` + Model string `json:"model,omitempty"` + Size string `json:"size,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` +} + +type RemixRequest struct { + Prompt string `json:"prompt"` +} + +// ── Response structs ── + +type SoraTaskResponse struct { + ID string `json:"id"` + Object string `json:"object"` + CreatedAt int64 `json:"created_at"` + Status string `json:"status"` + Model string `json:"model,omitempty"` + Prompt string `json:"prompt,omitempty"` + Progress int `json:"progress"` + VideoURL string `json:"video_url,omitempty"` + ShareID string `json:"share_id,omitempty"` + Seconds string `json:"seconds,omitempty"` + Size string `json:"size,omitempty"` + Character *SoraCharacter `json:"character,omitempty"` + URL string `json:"url,omitempty"` + CompletedAt *int64 `json:"completed_at,omitempty"` + Error *SoraTaskError `json:"error,omitempty"` +} + +type SoraTaskError struct { + Message string `json:"message"` + Type string `json:"type"` +} + +func TaskToResponse(t *SoraTask) *SoraTaskResponse { + resp := &SoraTaskResponse{ + ID: t.ID, + Object: t.ObjectType, + CreatedAt: t.CreatedAt.Unix(), + Status: t.Status, + Model: t.Model, + Prompt: t.Prompt, + Progress: t.Progress, + Seconds: t.Seconds, + Size: t.Size, + } + if t.VideoURL != "" { + resp.VideoURL = t.VideoURL + } + if t.ShareID != "" { + resp.ShareID = t.ShareID + } + if t.CharacterInfo != nil { + resp.Character = t.CharacterInfo + } + if t.CompletedAt != nil { + ts := t.CompletedAt.Unix() + resp.CompletedAt = &ts + } + if t.Status == SoraTaskFailed && t.ErrorMessage != "" { + resp.Error = &SoraTaskError{ + Message: t.ErrorMessage, + Type: t.ErrorType, + } + } + return resp +} + +// ── Task creation ── + +func (s *SoraTaskService) CreateVideoTask( + ctx context.Context, + apiKeyID int64, + account *Account, + req *CreateVideoRequest, + body []byte, +) (*SoraTask, error) { + modelCfg, ok := GetSoraModelConfig(req.Model) + + objectType := SoraObjectVideo + if req.Video != "" && req.Prompt == "" { + objectType = SoraObjectCharacter + } + + taskID := generateTaskID(objectType) + + seconds := "" + size := "" + if ok && modelCfg.Type == "video" { + seconds = fmt.Sprintf("%d", modelCfg.Frames/30) + if modelCfg.Orientation == "landscape" { + size = "1920x1080" + } else { + size = "1080x1920" + } + } + + task := &SoraTask{ + ID: taskID, + AccountID: account.ID, + APIKeyID: &apiKeyID, + ObjectType: objectType, + Model: req.Model, + Prompt: req.Prompt, + Status: SoraTaskQueued, + Progress: 0, + Seconds: seconds, + Size: size, + CreatedAt: time.Now(), + } + + if account.Type == AccountTypeAPIKey { + task.RequestBody = body + upstreamResp, err := s.forwardCreateToUpstream(ctx, account, "/v1/videos", body) + if err != nil { + return nil, fmt.Errorf("forward to upstream: %w", err) + } + s.applyUpstreamResponse(task, upstreamResp) + } else { + if s.soraClient == nil || !s.soraClient.Enabled() { + return nil, fmt.Errorf("sora SDK client not configured") + } + upstreamID, err := s.createViaSdk(ctx, account, req, modelCfg, objectType) + if err != nil { + return nil, fmt.Errorf("sdk create task: %w", err) + } + task.UpstreamTaskID = upstreamID + } + + if err := s.repo.Create(ctx, task); err != nil { + return nil, fmt.Errorf("save task: %w", err) + } + return task, nil +} + +func (s *SoraTaskService) CreateImageGeneration( + ctx context.Context, + apiKeyID int64, + account *Account, + req *CreateImageRequest, + body []byte, +) (*SoraTask, error) { + taskID := generateTaskID(SoraObjectImage) + + task := &SoraTask{ + ID: taskID, + AccountID: account.ID, + APIKeyID: &apiKeyID, + ObjectType: SoraObjectImage, + Model: req.Model, + Prompt: req.Prompt, + Status: SoraTaskQueued, + Progress: 0, + Size: req.Size, + CreatedAt: time.Now(), + } + + if account.Type == AccountTypeAPIKey { + task.RequestBody = body + upstreamResp, err := s.forwardCreateToUpstream(ctx, account, "/v1/images/generations", body) + if err != nil { + return nil, fmt.Errorf("forward to upstream: %w", err) + } + s.applyUpstreamResponse(task, upstreamResp) + } else { + if s.soraClient == nil || !s.soraClient.Enabled() { + return nil, fmt.Errorf("sora SDK client not configured") + } + modelCfg, _ := GetSoraModelConfig(req.Model) + width, height := modelCfg.Width, modelCfg.Height + if width == 0 { + width = 360 + } + if height == 0 { + height = 360 + } + upstreamID, err := s.soraClient.CreateImageTask(ctx, account, SoraImageRequest{ + Prompt: req.Prompt, + Width: width, + Height: height, + }) + if err != nil { + return nil, fmt.Errorf("sdk create image task: %w", err) + } + task.UpstreamTaskID = upstreamID + } + + if err := s.repo.Create(ctx, task); err != nil { + return nil, fmt.Errorf("save task: %w", err) + } + return task, nil +} + +func (s *SoraTaskService) GetTask(ctx context.Context, taskID string, apiKeyID int64) (*SoraTask, error) { + return s.repo.GetByIDAndAPIKey(ctx, taskID, apiKeyID) +} + +func (s *SoraTaskService) GetTaskByID(ctx context.Context, taskID string) (*SoraTask, error) { + return s.repo.GetByID(ctx, taskID) +} + +// ── Internal methods ── + +func (s *SoraTaskService) createViaSdk( + ctx context.Context, + account *Account, + req *CreateVideoRequest, + modelCfg SoraModelConfig, + objectType string, +) (string, error) { + if objectType == SoraObjectCharacter { + return "", fmt.Errorf("character creation via /v1/videos is not yet supported for OAuth accounts") + } + + orientation := modelCfg.Orientation + if req.Orientation != "" { + orientation = req.Orientation + } + + videoReq := SoraVideoRequest{ + Prompt: req.Prompt, + Orientation: orientation, + Frames: modelCfg.Frames, + Model: modelCfg.Model, + Size: modelCfg.Size, + RemixTargetID: req.RemixTargetID, + } + return s.soraClient.CreateVideoTask(ctx, account, videoReq) +} + +func (s *SoraTaskService) forwardCreateToUpstream( + ctx context.Context, + account *Account, + path string, + body []byte, +) (map[string]any, error) { + apiKey := account.GetCredential("api_key") + baseURL := account.GetBaseURL() + if apiKey == "" || baseURL == "" { + return nil, fmt.Errorf("account %d missing api_key or base_url", account.ID) + } + + upstreamURL := strings.TrimRight(baseURL, "/") + path + req, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamURL, bytes.NewReader(body)) + if err != nil { + return nil, err + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+apiKey) + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, fmt.Errorf("upstream request: %w", err) + } + defer resp.Body.Close() + + respBody, err := io.ReadAll(io.LimitReader(resp.Body, 256*1024)) + if err != nil { + return nil, fmt.Errorf("read upstream response: %w", err) + } + if resp.StatusCode >= 400 { + return nil, fmt.Errorf("upstream error %d: %s", resp.StatusCode, string(respBody)) + } + + var result map[string]any + if err := json.Unmarshal(respBody, &result); err != nil { + return nil, fmt.Errorf("parse upstream response: %w", err) + } + return result, nil +} + +func (s *SoraTaskService) forwardGetToUpstream( + ctx context.Context, + account *Account, + path string, +) (map[string]any, error) { + apiKey := account.GetCredential("api_key") + baseURL := account.GetBaseURL() + if apiKey == "" || baseURL == "" { + return nil, fmt.Errorf("account %d missing api_key or base_url", account.ID) + } + + upstreamURL := strings.TrimRight(baseURL, "/") + path + req, err := http.NewRequestWithContext(ctx, http.MethodGet, upstreamURL, nil) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+apiKey) + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, fmt.Errorf("upstream request: %w", err) + } + defer resp.Body.Close() + + respBody, err := io.ReadAll(io.LimitReader(resp.Body, 256*1024)) + if err != nil { + return nil, fmt.Errorf("read upstream response: %w", err) + } + if resp.StatusCode >= 400 { + return nil, fmt.Errorf("upstream error %d: %s", resp.StatusCode, string(respBody)) + } + + var result map[string]any + if err := json.Unmarshal(respBody, &result); err != nil { + return nil, fmt.Errorf("parse upstream response: %w", err) + } + return result, nil +} + +func (s *SoraTaskService) applyUpstreamResponse(task *SoraTask, resp map[string]any) { + if id, ok := resp["id"].(string); ok && id != "" { + task.UpstreamTaskID = id + } + if status, ok := resp["status"].(string); ok && status != "" { + task.Status = status + } + if progress, ok := resp["progress"].(float64); ok { + task.Progress = int(progress) + } + if videoURL, ok := resp["video_url"].(string); ok { + task.VideoURL = videoURL + } + if shareID, ok := resp["share_id"].(string); ok { + task.ShareID = shareID + } + if seconds, ok := resp["seconds"].(string); ok { + task.Seconds = seconds + } + if size, ok := resp["size"].(string); ok { + task.Size = size + } + if obj, ok := resp["object"].(string); ok && obj != "" { + task.ObjectType = obj + } +} + +// PollTask polls a single task's latest status (called by worker). +func (s *SoraTaskService) PollTask(ctx context.Context, task *SoraTask, account *Account) error { + if account.Type == AccountTypeAPIKey { + return s.pollUpstreamTask(ctx, task, account) + } + return s.pollSdkTask(ctx, task, account) +} + +func (s *SoraTaskService) pollUpstreamTask(ctx context.Context, task *SoraTask, account *Account) error { + upstreamID := task.UpstreamTaskID + if upstreamID == "" { + upstreamID = task.ID + } + + path := fmt.Sprintf("/v1/videos/%s", upstreamID) + resp, err := s.forwardGetToUpstream(ctx, account, path) + if err != nil { + logger.LegacyPrintf("service.sora_task", "[PollUpstream] task=%s error=%v", task.ID, err) + return err + } + + s.applyUpstreamResponse(task, resp) + now := time.Now() + if task.Status == SoraTaskCompleted || task.Status == SoraTaskFailed { + task.CompletedAt = &now + } + + if errObj, ok := resp["error"].(map[string]any); ok { + if msg, ok := errObj["message"].(string); ok { + task.ErrorMessage = msg + } + if typ, ok := errObj["type"].(string); ok { + task.ErrorType = typ + } + } + + return s.repo.Update(ctx, task) +} + +func (s *SoraTaskService) pollSdkTask(ctx context.Context, task *SoraTask, account *Account) error { + if s.soraClient == nil { + return fmt.Errorf("sora SDK client not configured") + } + upstreamID := task.UpstreamTaskID + if upstreamID == "" { + return fmt.Errorf("task %s has no upstream_task_id", task.ID) + } + + now := time.Now() + + switch task.ObjectType { + case SoraObjectImage: + status, err := s.soraClient.GetImageTask(ctx, account, upstreamID) + if err != nil { + return err + } + task.Progress = int(status.ProgressPct) + switch status.Status { + case "complete", "completed": + task.Status = SoraTaskCompleted + task.Progress = 100 + task.CompletedAt = &now + if len(status.URLs) > 0 { + task.VideoURL = status.URLs[0] + } + case "failed": + task.Status = SoraTaskFailed + task.CompletedAt = &now + task.ErrorMessage = status.ErrorMsg + task.ErrorType = "server_error" + default: + task.Status = SoraTaskInProgress + } + + case SoraObjectVideo, SoraObjectCharacter: + status, err := s.soraClient.GetVideoTask(ctx, account, upstreamID) + if err != nil { + return err + } + task.Progress = status.ProgressPct + switch status.Status { + case "complete", "completed": + task.Status = SoraTaskCompleted + task.Progress = 100 + task.CompletedAt = &now + if len(status.URLs) > 0 { + task.VideoURL = status.URLs[0] + } + if status.GenerationID != "" { + task.ShareID = status.GenerationID + } + case "failed": + task.Status = SoraTaskFailed + task.CompletedAt = &now + task.ErrorMessage = status.ErrorMsg + task.ErrorType = "server_error" + default: + task.Status = SoraTaskInProgress + } + + default: + return fmt.Errorf("unknown object type: %s", task.ObjectType) + } + + return s.repo.Update(ctx, task) +} + +func generateTaskID(objectType string) string { + b := make([]byte, 8) + if _, err := rand.Read(b); err != nil { + b = []byte(fmt.Sprintf("%x", time.Now().UnixNano())) + } + switch objectType { + case SoraObjectCharacter: + return "char_" + hex.EncodeToString(b) + case SoraObjectImage: + return "img_" + hex.EncodeToString(b) + default: + return "video_" + hex.EncodeToString(b) + } +} diff --git a/backend/internal/service/sora_task_worker.go b/backend/internal/service/sora_task_worker.go new file mode 100644 index 0000000000..eff12f149a --- /dev/null +++ b/backend/internal/service/sora_task_worker.go @@ -0,0 +1,183 @@ +package service + +import ( + "context" + "sync" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" +) + +// SoraTaskWorker polls unfinished Sora tasks in the background. +// When a task completes, it downloads media to configured storage if available. +type SoraTaskWorker struct { + taskService *SoraTaskService + accountRepo AccountRepository + s3Storage *SoraS3Storage + mediaStorage *SoraMediaStorage + interval time.Duration + pollTimeout time.Duration + stopCh chan struct{} + stopOnce sync.Once + wg sync.WaitGroup +} + +func NewSoraTaskWorker( + taskService *SoraTaskService, + accountRepo AccountRepository, + s3Storage *SoraS3Storage, + mediaStorage *SoraMediaStorage, + interval time.Duration, +) *SoraTaskWorker { + if interval <= 0 { + interval = 60 * time.Second + } + return &SoraTaskWorker{ + taskService: taskService, + accountRepo: accountRepo, + s3Storage: s3Storage, + mediaStorage: mediaStorage, + interval: interval, + pollTimeout: 30 * time.Second, + stopCh: make(chan struct{}), + } +} + +func (w *SoraTaskWorker) Start() { + if w == nil || w.taskService == nil { + return + } + w.wg.Add(1) + go func() { + defer w.wg.Done() + ticker := time.NewTicker(w.interval) + defer ticker.Stop() + + logger.LegacyPrintf("service.sora_task_worker", "[Start] polling interval=%s", w.interval) + w.pollAll() + + for { + select { + case <-ticker.C: + w.pollAll() + case <-w.stopCh: + return + } + } + }() +} + +func (w *SoraTaskWorker) Stop() { + if w == nil { + return + } + w.stopOnce.Do(func() { + close(w.stopCh) + }) + w.wg.Wait() +} + +func (w *SoraTaskWorker) pollAll() { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute) + defer cancel() + + tasks, err := w.taskService.repo.ListPending(ctx) + if err != nil { + logger.LegacyPrintf("service.sora_task_worker", "[PollAll] list pending tasks error: %v", err) + return + } + if len(tasks) == 0 { + return + } + + logger.LegacyPrintf("service.sora_task_worker", "[PollAll] found %d pending tasks", len(tasks)) + + for _, task := range tasks { + select { + case <-w.stopCh: + return + default: + } + w.pollOne(ctx, task) + } +} + +func (w *SoraTaskWorker) pollOne(ctx context.Context, task *SoraTask) { + pollCtx, cancel := context.WithTimeout(ctx, w.pollTimeout) + defer cancel() + + account, err := w.accountRepo.GetByID(pollCtx, task.AccountID) + if err != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[PollOne] task=%s get account=%d error: %v", task.ID, task.AccountID, err) + w.markTaskFailed(pollCtx, task, "account not found") + return + } + + if err := w.taskService.PollTask(pollCtx, task, account); err != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[PollOne] task=%s poll error: %v", task.ID, err) + return + } + + logger.LegacyPrintf("service.sora_task_worker", + "[PollOne] task=%s status=%s progress=%d", task.ID, task.Status, task.Progress) + + if task.Status == SoraTaskCompleted && task.VideoURL != "" && task.StoredKey == "" { + w.tryStoreMedia(pollCtx, task) + } +} + +// tryStoreMedia downloads media to configured storage. +// Priority: S3 > local disk > keep upstream URL. +func (w *SoraTaskWorker) tryStoreMedia(ctx context.Context, task *SoraTask) { + mediaType := "video" + if task.ObjectType == SoraObjectImage { + mediaType = "image" + } + + if w.s3Storage != nil && w.s3Storage.Enabled(ctx) { + key, _, err := w.s3Storage.UploadFromURL(ctx, 0, task.VideoURL) + if err != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[StoreMedia] task=%s S3 upload error: %v", task.ID, err) + } else { + task.StoredKey = key + task.StorageType = "s3" + if updateErr := w.taskService.repo.Update(ctx, task); updateErr != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[StoreMedia] task=%s update stored_key error: %v", task.ID, updateErr) + } + return + } + } + + if w.mediaStorage != nil && w.mediaStorage.Enabled() { + stored, err := w.mediaStorage.StoreFromURLs(ctx, mediaType, []string{task.VideoURL}) + if err != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[StoreMedia] task=%s local storage error: %v", task.ID, err) + return + } + if len(stored) > 0 && stored[0] != task.VideoURL { + task.StoredKey = stored[0] + task.StorageType = "local" + if updateErr := w.taskService.repo.Update(ctx, task); updateErr != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[StoreMedia] task=%s update stored_key error: %v", task.ID, updateErr) + } + } + } +} + +func (w *SoraTaskWorker) markTaskFailed(ctx context.Context, task *SoraTask, message string) { + now := time.Now() + task.Status = SoraTaskFailed + task.ErrorMessage = message + task.ErrorType = "server_error" + task.CompletedAt = &now + if updateErr := w.taskService.repo.Update(ctx, task); updateErr != nil { + logger.LegacyPrintf("service.sora_task_worker", + "[PollOne] task=%s update failed: %v", task.ID, updateErr) + } +} diff --git a/backend/migrations/070_add_sora_tasks.sql b/backend/migrations/070_add_sora_tasks.sql new file mode 100644 index 0000000000..3a484976d9 --- /dev/null +++ b/backend/migrations/070_add_sora_tasks.sql @@ -0,0 +1,36 @@ +-- +migrate Up + +-- Sora 异步任务表:记录视频/角色/图片生成任务状态 +CREATE TABLE IF NOT EXISTS sora_tasks ( + id VARCHAR(64) PRIMARY KEY, -- 对外 ID (video_xxx / char_xxx / img_xxx) + account_id BIGINT NOT NULL, -- 使用的账号 ID + api_key_id BIGINT, -- 调用方 API Key ID + upstream_task_id VARCHAR(128) NOT NULL DEFAULT '', -- 上游任务 ID + object_type VARCHAR(16) NOT NULL DEFAULT 'video', -- video / character / image + model VARCHAR(64) NOT NULL, -- 请求模型名 + prompt TEXT NOT NULL DEFAULT '', -- 生成提示词 + status VARCHAR(16) NOT NULL DEFAULT 'queued', -- queued / in_progress / completed / failed + progress INT NOT NULL DEFAULT 0, -- 进度 0-100 + video_url TEXT NOT NULL DEFAULT '', -- 原始上游下载 URL + stored_key TEXT NOT NULL DEFAULT '', -- 存储后的 key(本地相对路径或 S3 object key) + storage_type VARCHAR(16) NOT NULL DEFAULT '', -- 存储类型:local / s3 / gdrive / 空=未存储 + share_id VARCHAR(128) NOT NULL DEFAULT '', -- 可分享 ID + character_info JSONB, -- 角色信息 {username, display_name} + error_message TEXT NOT NULL DEFAULT '', -- 错误信息 + error_type VARCHAR(64) NOT NULL DEFAULT '', -- 错误类型 + request_body JSONB, -- 原始请求体(用于上游重放) + seconds VARCHAR(8) NOT NULL DEFAULT '', -- 视频时长 + size VARCHAR(16) NOT NULL DEFAULT '', -- 分辨率 (如 1920x1080) + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + completed_at TIMESTAMPTZ +); + +-- 仅索引未完成任务,供后台 worker 轮询 +CREATE INDEX IF NOT EXISTS idx_sora_tasks_pending ON sora_tasks(status) WHERE status IN ('queued', 'in_progress'); + +-- 按 API Key 查询 +CREATE INDEX IF NOT EXISTS idx_sora_tasks_api_key ON sora_tasks(api_key_id) WHERE api_key_id IS NOT NULL; + +-- +migrate Down + +DROP TABLE IF EXISTS sora_tasks;