mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
feat: add Sora async task API (videos/images)
Add async task management for Sora video/image generation with two forwarding modes: OAuth (SDK) and API Key (HTTP upstream). Includes background worker for status polling and 3-layer storage degradation (S3 > local disk > upstream URL passthrough). Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -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{
|
||||
|
||||
@@ -75,6 +75,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
|
||||
antigravityOAuthSvc,
|
||||
nil, // openAIGateway
|
||||
nil, // scheduledTestRunner
|
||||
nil, // soraTaskWorker
|
||||
)
|
||||
|
||||
require.NotPanics(t, func() {
|
||||
|
||||
@@ -45,6 +45,7 @@ type Handlers struct {
|
||||
OpenAIGateway *OpenAIGatewayHandler
|
||||
SoraGateway *SoraGatewayHandler
|
||||
SoraClient *SoraClientHandler
|
||||
SoraVideos *SoraVideosHandler
|
||||
Setting *SettingHandler
|
||||
Totp *TotpHandler
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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 验证)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
Reference in New Issue
Block a user