mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:28:03 +08:00
Merge pull request #4406 from StarryKira/agent/fix-4326-async-image-object-storage
feat: 异步生图任务与结果轮询(重新引入 #4381)+ 结果落对象存储
This commit is contained in:
@@ -135,6 +135,7 @@ docs/*
|
||||
!docs/PAYMENT.md
|
||||
!docs/PAYMENT_CN.md
|
||||
!docs/ADMIN_PAYMENT_INTEGRATION_API.md
|
||||
!docs/ASYNC_IMAGE_TASKS.md
|
||||
!docs/legal/
|
||||
!docs/legal/*.md
|
||||
.serena/
|
||||
|
||||
@@ -688,6 +688,12 @@ Simple Mode is designed for individual developers or internal teams who want qui
|
||||
|
||||
---
|
||||
|
||||
## Asynchronous Image Tasks
|
||||
|
||||
Long-running OpenAI/Grok image generation and editing can be submitted through `/v1/images/generations/async` or `/v1/images/edits/async`, then polled at `/v1/images/tasks/{task_id}` without holding a CDN connection open. See [Asynchronous Image Tasks](docs/ASYNC_IMAGE_TASKS.md) for request and response examples.
|
||||
|
||||
---
|
||||
|
||||
## Grok / xAI Support
|
||||
|
||||
Sub2API supports both Grok subscription accounts through xAI OAuth and standard xAI API-key accounts. Both account types forward OpenAI-compatible Responses traffic to xAI.
|
||||
|
||||
@@ -70,7 +70,6 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig)
|
||||
schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig)
|
||||
accountRepository := repository.NewAccountRepository(client, db, schedulerCache)
|
||||
adminAccountRepository := repository.NewAdminAccountRepository(client, db, schedulerCache)
|
||||
concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig)
|
||||
apiKeyService := service.ProvideAPIKeyService(apiKeyRepository, userRepository, groupRepository, userSubscriptionRepository, userGroupRateRepository, apiKeyCache, configConfig, billingCacheService, concurrencyService)
|
||||
apiKeyAuthCacheInvalidator := service.ProvideAPIKeyAuthCacheInvalidator(apiKeyService)
|
||||
@@ -174,6 +173,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
leaderLockCache := repository.NewLeaderLockCache(redisClient)
|
||||
dashboardAggregationService := service.ProvideDashboardAggregationService(dashboardAggregationRepository, timingWheelService, leaderLockCache, db, configConfig)
|
||||
dashboardHandler := admin.NewDashboardHandler(dashboardService, dashboardAggregationService)
|
||||
adminAccountRepository := repository.NewAdminAccountRepository(client, db, schedulerCache)
|
||||
proxyExitInfoProber := repository.NewProxyExitInfoProber(configConfig)
|
||||
proxyLatencyCache := repository.NewProxyLatencyCache(redisClient)
|
||||
adminService := service.NewAdminService(userRepository, groupRepository, adminAccountRepository, proxyRepository, apiKeyRepository, redeemCodeRepository, userGroupRateRepository, userRPMCache, billingCacheService, proxyExitInfoProber, proxyLatencyCache, apiKeyAuthCacheInvalidator, client, settingService, subscriptionService, userSubscriptionRepository, privacyClientFactory, openAIGatewayService, affiliateService)
|
||||
@@ -250,10 +250,10 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
paymentHandler := admin.NewPaymentHandler(paymentService, paymentConfigService)
|
||||
affiliateHandler := admin.NewAffiliateHandler(affiliateService, adminService)
|
||||
complianceHandler := admin.NewComplianceHandler(settingService)
|
||||
upstreamBillingProbeService := service.ProvideUpstreamBillingProbeService(accountRepository, accountTestService, settingService, leaderLockCache, db)
|
||||
auditLogRepository := repository.NewAuditLogRepository(db)
|
||||
auditLogService := service.ProvideAuditLogService(auditLogRepository, settingService)
|
||||
auditLogHandler := admin.NewAuditLogHandler(auditLogService, totpService)
|
||||
upstreamBillingProbeService := service.ProvideUpstreamBillingProbeService(accountRepository, accountTestService, settingService, leaderLockCache, db)
|
||||
adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, paymentHandler, affiliateHandler, complianceHandler, auditLogHandler, upstreamBillingProbeService)
|
||||
usageRecordWorkerPool := service.NewUsageRecordWorkerPool(configConfig)
|
||||
userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient)
|
||||
@@ -265,6 +265,13 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService)
|
||||
paymentWebhookHandler := handler.NewPaymentWebhookHandler(paymentService, registry)
|
||||
availableChannelHandler := handler.NewAvailableChannelHandler(channelService, apiKeyService, settingService)
|
||||
imageTaskStore := repository.NewImageTaskStore(redisClient)
|
||||
imageStorage, err := repository.ProvideImageStorage(configConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
imageTaskService := service.ProvideImageTaskService(imageTaskStore, imageStorage, configConfig)
|
||||
asyncImageHandler := handler.NewAsyncImageHandler(imageTaskService, openAIGatewayHandler)
|
||||
batchImageRepository := repository.NewBatchImageRepository(db)
|
||||
batchImageQueue := repository.NewBatchImageQueue(redisClient, configConfig)
|
||||
batchImageModelPricingResolver := service.ProvideBatchImageModelPricingResolver(modelPricingResolver)
|
||||
@@ -275,7 +282,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
batchImageHandler := handler.NewBatchImageHandler(batchImagePublicService, batchImageDownloadService, batchImageCleanupService)
|
||||
idempotencyCoordinator := service.ProvideIdempotencyCoordinator(idempotencyRepository, configConfig)
|
||||
idempotencyCleanupService := service.ProvideIdempotencyCleanupService(idempotencyRepository, configConfig)
|
||||
handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, channelMonitorUserHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, handlerSettingHandler, totpHandler, handlerPaymentHandler, paymentWebhookHandler, availableChannelHandler, batchImageHandler, idempotencyCoordinator, idempotencyCleanupService)
|
||||
handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, channelMonitorUserHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, handlerSettingHandler, totpHandler, handlerPaymentHandler, paymentWebhookHandler, availableChannelHandler, asyncImageHandler, batchImageHandler, idempotencyCoordinator, idempotencyCleanupService)
|
||||
jwtAuthMiddleware := middleware.NewJWTAuthMiddleware(authService, userService, settingService, auditLogService)
|
||||
adminAuthMiddleware := middleware.NewAdminAuthMiddleware(authService, userService, settingService, auditLogService)
|
||||
apiKeyAuthMiddleware := middleware.NewAPIKeyAuthMiddleware(apiKeyService, subscriptionService, configConfig)
|
||||
@@ -369,12 +376,6 @@ func provideCleanup(
|
||||
}
|
||||
|
||||
parallelSteps := []cleanupStep{
|
||||
{"AuditLogService", func() error {
|
||||
if auditLog != nil {
|
||||
auditLog.Stop()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"OpsScheduledReportService", func() error {
|
||||
if opsScheduledReport != nil {
|
||||
opsScheduledReport.Stop()
|
||||
@@ -393,6 +394,12 @@ func provideCleanup(
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"AuditLogService", func() error {
|
||||
if auditLog != nil {
|
||||
auditLog.Stop()
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"OpsAlertEvaluatorService", func() error {
|
||||
if opsAlertEvaluator != nil {
|
||||
opsAlertEvaluator.Stop()
|
||||
|
||||
@@ -166,6 +166,8 @@ github.com/google/go-querystring v1.1.0/go.mod h1:Kcdr2DB4koayq7X8pmAG4sNG59So17
|
||||
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
|
||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs=
|
||||
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
|
||||
github.com/google/subcommands v1.2.0 h1:vWQspBTo2nEqTUFita5/KeEWlUL8kQObDFbub/EN9oE=
|
||||
github.com/google/subcommands v1.2.0/go.mod h1:ZjhPrFU+Olkh9WazFPsl27BQ4UPiG37m3yTrtFlrHVk=
|
||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||
github.com/google/wire v0.7.0 h1:JxUKI6+CVBgCO2WToKy/nQk0sS+amI9z9EjVmdaocj4=
|
||||
|
||||
@@ -95,6 +95,7 @@ type Config struct {
|
||||
Update UpdateConfig `mapstructure:"update"`
|
||||
Idempotency IdempotencyConfig `mapstructure:"idempotency"`
|
||||
BatchImage BatchImageConfig `mapstructure:"batch_image"`
|
||||
ImageStorage ImageStorageConfig `mapstructure:"image_storage"`
|
||||
}
|
||||
|
||||
type LogConfig struct {
|
||||
@@ -227,6 +228,33 @@ type BatchImageConfig struct {
|
||||
VertexGCSBaseURL string `mapstructure:"vertex_gcs_base_url"`
|
||||
}
|
||||
|
||||
// ImageStorageConfig 配置异步图片任务结果上传的 S3 兼容对象存储。
|
||||
// Enabled 同时作为异步图片任务功能的总开关:未启用或未配置完整凭证时,
|
||||
// 异步生图接口整体禁用,避免把上游返回的大 base64 结果塞进 Redis。
|
||||
type ImageStorageConfig struct {
|
||||
Enabled bool `mapstructure:"enabled"`
|
||||
Endpoint string `mapstructure:"endpoint"` // e.g. https://<account_id>.r2.cloudflarestorage.com
|
||||
Region string `mapstructure:"region"` // R2 用 "auto"
|
||||
Bucket string `mapstructure:"bucket"`
|
||||
AccessKeyID string `mapstructure:"access_key_id"`
|
||||
SecretAccessKey string `mapstructure:"secret_access_key"`
|
||||
Prefix string `mapstructure:"prefix"` // S3 key 前缀,如 "images/"
|
||||
ForcePathStyle bool `mapstructure:"force_path_style"` // MinIO/路径风格桶
|
||||
PublicBaseURL string `mapstructure:"public_base_url"` // 配了则返回 public_base_url/key 直链;否则 presigned
|
||||
PresignExpiry int `mapstructure:"presign_expiry_hours"` // public_base_url 为空时的 presigned 过期时长(小时)
|
||||
MaxDownloadByte int64 `mapstructure:"max_download_bytes"` // 下载上游 url 图片的字节上限
|
||||
}
|
||||
|
||||
// IsConfigured 检查对象存储必要字段是否已配置
|
||||
func (c *ImageStorageConfig) IsConfigured() bool {
|
||||
return c.Bucket != "" && c.AccessKeyID != "" && c.SecretAccessKey != ""
|
||||
}
|
||||
|
||||
// Active 返回异步图片任务是否可用:开关打开且凭证齐全
|
||||
func (c *ImageStorageConfig) Active() bool {
|
||||
return c.Enabled && c.IsConfigured()
|
||||
}
|
||||
|
||||
type LinuxDoConnectConfig struct {
|
||||
Enabled bool `mapstructure:"enabled"`
|
||||
ClientID string `mapstructure:"client_id"`
|
||||
@@ -1891,6 +1919,14 @@ func setDefaults() {
|
||||
viper.SetDefault("batch_image.vertex_batch_prediction_base_url", "")
|
||||
viper.SetDefault("batch_image.vertex_gcs_base_url", "")
|
||||
|
||||
// Image storage (async image task result offload to S3-compatible object storage)
|
||||
viper.SetDefault("image_storage.enabled", false)
|
||||
viper.SetDefault("image_storage.region", "auto")
|
||||
viper.SetDefault("image_storage.prefix", "images/")
|
||||
viper.SetDefault("image_storage.force_path_style", false)
|
||||
viper.SetDefault("image_storage.presign_expiry_hours", 24)
|
||||
viper.SetDefault("image_storage.max_download_bytes", 33554432)
|
||||
|
||||
// Ops (vNext)
|
||||
viper.SetDefault("ops.enabled", true)
|
||||
viper.SetDefault("ops.use_preaggregated_tables", true)
|
||||
|
||||
@@ -23,6 +23,7 @@ const (
|
||||
EndpointResponsesCompact = "/v1/responses/compact"
|
||||
EndpointImagesGenerations = "/v1/images/generations"
|
||||
EndpointImagesEdits = "/v1/images/edits"
|
||||
EndpointImageTasks = "/v1/images/tasks"
|
||||
EndpointVideosGenerations = "/v1/videos/generations"
|
||||
EndpointVideosEdits = "/v1/videos/edits"
|
||||
EndpointVideosExtensions = "/v1/videos/extensions"
|
||||
@@ -88,6 +89,8 @@ func NormalizeInboundEndpoint(path string) string {
|
||||
return EndpointImagesGenerations
|
||||
case strings.Contains(path, EndpointImagesEdits) || strings.Contains(path, "/images/edits"):
|
||||
return EndpointImagesEdits
|
||||
case strings.Contains(path, EndpointImageTasks) || strings.Contains(path, "/images/tasks/"):
|
||||
return EndpointImageTasks
|
||||
case strings.Contains(path, EndpointVideosGenerations) || strings.Contains(path, "/videos/generations"):
|
||||
return EndpointVideosGenerations
|
||||
case strings.Contains(path, EndpointVideosEdits) || strings.Contains(path, "/videos/edits"):
|
||||
|
||||
@@ -31,6 +31,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
|
||||
{"/v1/responses/compact/detail", EndpointResponsesCompact},
|
||||
{"/v1/images/generations", EndpointImagesGenerations},
|
||||
{"/v1/images/edits", EndpointImagesEdits},
|
||||
{"/v1/images/tasks/imgtask_123", EndpointImageTasks},
|
||||
{"/v1/videos/generations", EndpointVideosGenerations},
|
||||
{"/v1/videos/req_123", EndpointVideos},
|
||||
{"/v1beta/models", EndpointGeminiModels},
|
||||
@@ -52,6 +53,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
|
||||
{"/responses/compact", EndpointResponsesCompact},
|
||||
{"/responses/compact/detail", EndpointResponsesCompact},
|
||||
{"/alpha/search", EndpointAlphaSearch},
|
||||
{"/images/tasks/imgtask_123", EndpointImageTasks},
|
||||
|
||||
// Bare Codex direct alias route — root vs. compact.
|
||||
{"/backend-api/codex/responses", EndpointResponses},
|
||||
|
||||
@@ -59,6 +59,7 @@ type Handlers struct {
|
||||
Payment *PaymentHandler
|
||||
PaymentWebhook *PaymentWebhookHandler
|
||||
AvailableChannel *AvailableChannelHandler
|
||||
AsyncImage *AsyncImageHandler
|
||||
BatchImage *BatchImageHandler
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,281 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
|
||||
"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"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
type AsyncImageHandler struct {
|
||||
tasks *service.ImageTaskService
|
||||
openAI *OpenAIGatewayHandler
|
||||
execute func(platform string, c *gin.Context)
|
||||
}
|
||||
|
||||
func NewAsyncImageHandler(tasks *service.ImageTaskService, openAI *OpenAIGatewayHandler) *AsyncImageHandler {
|
||||
h := &AsyncImageHandler{tasks: tasks, openAI: openAI}
|
||||
h.execute = h.executeWithGateway
|
||||
return h
|
||||
}
|
||||
|
||||
// enabled reports whether the async image task feature is available. Object
|
||||
// storage is the enablement gate: without it the endpoints are fully disabled
|
||||
// so that large base64 results never land in Redis.
|
||||
func (h *AsyncImageHandler) enabled() bool {
|
||||
return h != nil && h.tasks != nil && h.tasks.Enabled()
|
||||
}
|
||||
|
||||
// Submit accepts the same payload as the synchronous Images endpoint and
|
||||
// returns before the upstream image generation begins.
|
||||
func (h *AsyncImageHandler) Submit(c *gin.Context) {
|
||||
if !h.enabled() {
|
||||
imageTaskJSONError(c, http.StatusNotFound, "not_found_error", "async image tasks are not enabled")
|
||||
return
|
||||
}
|
||||
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
|
||||
if !ok || apiKey == nil || apiKey.UserID <= 0 || apiKey.ID <= 0 {
|
||||
imageTaskError(c, service.ErrImageTaskForbidden)
|
||||
return
|
||||
}
|
||||
platform := ""
|
||||
if apiKey.Group != nil {
|
||||
platform = apiKey.Group.Platform
|
||||
}
|
||||
if platform != service.PlatformOpenAI && platform != service.PlatformGrok {
|
||||
imageTaskJSONError(c, http.StatusNotFound, "not_found_error", "Images API is not supported for this platform")
|
||||
return
|
||||
}
|
||||
if !service.GroupAllowsImageGeneration(apiKey.Group) {
|
||||
imageTaskJSONError(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
|
||||
return
|
||||
}
|
||||
if h == nil || h.tasks == nil || h.execute == nil {
|
||||
imageTaskError(c, service.ErrImageTaskUnavailable)
|
||||
return
|
||||
}
|
||||
|
||||
body, err := pkghttputil.ReadRequestBodyWithPrealloc(c.Request)
|
||||
if err != nil {
|
||||
if maxErr, ok := extractMaxBytesError(err); ok {
|
||||
imageTaskJSONError(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit))
|
||||
return
|
||||
}
|
||||
imageTaskJSONError(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
|
||||
return
|
||||
}
|
||||
if len(body) == 0 {
|
||||
imageTaskJSONError(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty")
|
||||
return
|
||||
}
|
||||
if asyncImageRequestStreams(c.GetHeader("Content-Type"), body) {
|
||||
imageTaskJSONError(c, http.StatusBadRequest, "invalid_request_error", "streaming image requests cannot be submitted as asynchronous tasks")
|
||||
return
|
||||
}
|
||||
if err := h.validateRequest(c, platform, body); err != nil {
|
||||
imageTaskJSONError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
taskCtx, recorder, cancel := newAsyncImageContext(c, body, h.tasks.ExecutionTimeout())
|
||||
task, err := h.tasks.Create(c.Request.Context(), service.ImageTaskOwner{UserID: apiKey.UserID, APIKeyID: apiKey.ID})
|
||||
if err != nil {
|
||||
cancel()
|
||||
imageTaskError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
pollURL := imageTaskPollURL(c.Request.URL.Path, task.ID)
|
||||
c.Header("Cache-Control", "no-store")
|
||||
c.Header("Location", pollURL)
|
||||
c.Header("Retry-After", "3")
|
||||
c.JSON(http.StatusAccepted, gin.H{
|
||||
"id": task.ID,
|
||||
"task_id": task.TaskID,
|
||||
"object": task.Object,
|
||||
"status": task.Status,
|
||||
"created_at": task.CreatedAt,
|
||||
"expires_at": task.ExpiresAt,
|
||||
"poll_url": pollURL,
|
||||
})
|
||||
|
||||
go h.run(task.ID, platform, taskCtx, recorder, cancel)
|
||||
}
|
||||
|
||||
func (h *AsyncImageHandler) Get(c *gin.Context) {
|
||||
if !h.enabled() {
|
||||
imageTaskJSONError(c, http.StatusNotFound, "not_found_error", "async image tasks are not enabled")
|
||||
return
|
||||
}
|
||||
apiKey, ok := middleware2.GetAPIKeyFromContext(c)
|
||||
if !ok || apiKey == nil || apiKey.UserID <= 0 || apiKey.ID <= 0 {
|
||||
imageTaskError(c, service.ErrImageTaskForbidden)
|
||||
return
|
||||
}
|
||||
task, err := h.tasks.Get(c.Request.Context(), service.ImageTaskOwner{UserID: apiKey.UserID, APIKeyID: apiKey.ID}, c.Param("task_id"))
|
||||
if err != nil {
|
||||
imageTaskError(c, err)
|
||||
return
|
||||
}
|
||||
c.Header("Cache-Control", "no-store")
|
||||
if task.Status == service.ImageTaskStatusProcessing {
|
||||
c.Header("Retry-After", "3")
|
||||
}
|
||||
c.JSON(http.StatusOK, task)
|
||||
}
|
||||
|
||||
func (h *AsyncImageHandler) validateRequest(c *gin.Context, platform string, body []byte) error {
|
||||
if h.openAI == nil || h.openAI.gatewayService == nil {
|
||||
return nil
|
||||
}
|
||||
if platform == service.PlatformGrok {
|
||||
parsed := service.ParseGrokMediaRequest(c.GetHeader("Content-Type"), body)
|
||||
if strings.TrimSpace(parsed.Model) == "" {
|
||||
return errors.New("model is required")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
parsed, err := h.openAI.gatewayService.ParseOpenAIImagesRequest(c, body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if parsed.Stream {
|
||||
return errors.New("streaming image requests cannot be submitted as asynchronous tasks")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *AsyncImageHandler) executeWithGateway(platform string, c *gin.Context) {
|
||||
if h.openAI == nil {
|
||||
imageTaskJSONError(c, http.StatusServiceUnavailable, "api_error", "image gateway is unavailable")
|
||||
return
|
||||
}
|
||||
if platform == service.PlatformGrok {
|
||||
h.openAI.GrokImages(c)
|
||||
return
|
||||
}
|
||||
h.openAI.Images(c)
|
||||
}
|
||||
|
||||
func (h *AsyncImageHandler) run(taskID, platform string, taskCtx *gin.Context, recorder *httptest.ResponseRecorder, cancel context.CancelFunc) {
|
||||
defer cancel()
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
logger.L().Error("image_task.execution_panicked", zap.String("task_id", taskID), zap.Any("panic", recovered))
|
||||
h.failTask(taskID, http.StatusInternalServerError, imageTaskErrorPayload("api_error", "image generation task panicked"))
|
||||
}
|
||||
}()
|
||||
|
||||
h.execute(platform, taskCtx)
|
||||
body := bytes.TrimSpace(recorder.Body.Bytes())
|
||||
if err := taskCtx.Request.Context().Err(); err != nil && len(body) == 0 {
|
||||
h.failTask(taskID, http.StatusGatewayTimeout, imageTaskErrorPayload("timeout_error", "image generation task timed out"))
|
||||
return
|
||||
}
|
||||
statusCode := recorder.Code
|
||||
if statusCode == 0 {
|
||||
statusCode = http.StatusOK
|
||||
}
|
||||
if statusCode >= http.StatusOK && statusCode < http.StatusMultipleChoices {
|
||||
if len(body) == 0 || !json.Valid(body) {
|
||||
h.failTask(taskID, http.StatusBadGateway, imageTaskErrorPayload("api_error", "upstream returned an invalid image response"))
|
||||
return
|
||||
}
|
||||
if err := h.tasks.Complete(context.Background(), taskID, statusCode, json.RawMessage(body)); err != nil {
|
||||
logger.L().Error("image_task.complete_store_failed", zap.String("task_id", taskID), zap.Error(err))
|
||||
}
|
||||
return
|
||||
}
|
||||
h.failTask(taskID, statusCode, extractImageTaskError(body))
|
||||
}
|
||||
|
||||
func (h *AsyncImageHandler) failTask(taskID string, statusCode int, taskErr json.RawMessage) {
|
||||
if err := h.tasks.Fail(context.Background(), taskID, statusCode, taskErr); err != nil {
|
||||
logger.L().Error("image_task.failure_store_failed", zap.String("task_id", taskID), zap.Error(err))
|
||||
}
|
||||
}
|
||||
|
||||
func newAsyncImageContext(c *gin.Context, body []byte, timeoutDuration time.Duration) (*gin.Context, *httptest.ResponseRecorder, context.CancelFunc) {
|
||||
base := context.WithoutCancel(c.Request.Context())
|
||||
executionCtx, cancel := context.WithTimeout(base, timeoutDuration)
|
||||
request := c.Request.Clone(executionCtx)
|
||||
request.Body = io.NopCloser(bytes.NewReader(body))
|
||||
request.GetBody = func() (io.ReadCloser, error) {
|
||||
return io.NopCloser(bytes.NewReader(body)), nil
|
||||
}
|
||||
request.ContentLength = int64(len(body))
|
||||
request.URL.Path = strings.TrimSuffix(request.URL.Path, "/async")
|
||||
|
||||
taskCtx := c.Copy()
|
||||
recorder := httptest.NewRecorder()
|
||||
recorderCtx, _ := gin.CreateTestContext(recorder)
|
||||
taskCtx.Writer = recorderCtx.Writer
|
||||
taskCtx.Request = request
|
||||
return taskCtx, recorder, cancel
|
||||
}
|
||||
|
||||
func asyncImageRequestStreams(contentType string, body []byte) bool {
|
||||
if isMultipartImagesContentType(contentType) {
|
||||
return false
|
||||
}
|
||||
var envelope struct {
|
||||
Stream bool `json:"stream"`
|
||||
}
|
||||
return json.Unmarshal(body, &envelope) == nil && envelope.Stream
|
||||
}
|
||||
|
||||
func imageTaskPollURL(submitPath, taskID string) string {
|
||||
if strings.HasPrefix(submitPath, "/v1/") {
|
||||
return "/v1/images/tasks/" + taskID
|
||||
}
|
||||
return "/images/tasks/" + taskID
|
||||
}
|
||||
|
||||
func extractImageTaskError(body []byte) json.RawMessage {
|
||||
if json.Valid(body) {
|
||||
var envelope struct {
|
||||
Error json.RawMessage `json:"error"`
|
||||
}
|
||||
if json.Unmarshal(body, &envelope) == nil && len(envelope.Error) > 0 && json.Valid(envelope.Error) {
|
||||
return envelope.Error
|
||||
}
|
||||
return json.RawMessage(body)
|
||||
}
|
||||
return imageTaskErrorPayload("api_error", "image generation failed")
|
||||
}
|
||||
|
||||
func imageTaskErrorPayload(errorType, message string) json.RawMessage {
|
||||
data, _ := json.Marshal(gin.H{"type": errorType, "message": message})
|
||||
return data
|
||||
}
|
||||
|
||||
func imageTaskError(c *gin.Context, err error) {
|
||||
status := infraerrors.Code(err)
|
||||
code := infraerrors.Reason(err)
|
||||
message := infraerrors.Message(err)
|
||||
if status <= 0 {
|
||||
status = http.StatusInternalServerError
|
||||
}
|
||||
if strings.TrimSpace(code) == "" {
|
||||
code = "IMAGE_TASK_ERROR"
|
||||
}
|
||||
imageTaskJSONError(c, status, code, message)
|
||||
}
|
||||
|
||||
func imageTaskJSONError(c *gin.Context, status int, code, message string) {
|
||||
c.Header("Cache-Control", "no-store")
|
||||
c.JSON(status, gin.H{"error": gin.H{"type": code, "code": code, "message": message}})
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type asyncImageMemoryStore struct {
|
||||
mu sync.RWMutex
|
||||
tasks map[string]*service.ImageTaskRecord
|
||||
}
|
||||
|
||||
func (s *asyncImageMemoryStore) Save(_ context.Context, task *service.ImageTaskRecord, _ time.Duration) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
copy := *task
|
||||
copy.Result = append(json.RawMessage(nil), task.Result...)
|
||||
copy.Error = append(json.RawMessage(nil), task.Error...)
|
||||
s.tasks[task.ID] = ©
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *asyncImageMemoryStore) Get(_ context.Context, id string) (*service.ImageTaskRecord, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
task := s.tasks[id]
|
||||
if task == nil {
|
||||
return nil, service.ErrImageTaskNotFound
|
||||
}
|
||||
copy := *task
|
||||
copy.Result = append(json.RawMessage(nil), task.Result...)
|
||||
copy.Error = append(json.RawMessage(nil), task.Error...)
|
||||
return ©, nil
|
||||
}
|
||||
|
||||
func TestAsyncImageHandlerSubmitAndPoll(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
store := &asyncImageMemoryStore{tasks: make(map[string]*service.ImageTaskRecord)}
|
||||
tasks := service.NewImageTaskServiceWithUploader(store, nil, time.Hour, time.Minute)
|
||||
release := make(chan struct{})
|
||||
h := &AsyncImageHandler{tasks: tasks}
|
||||
h.execute = func(_ string, c *gin.Context) {
|
||||
<-release
|
||||
c.JSON(http.StatusOK, gin.H{"created": 123, "data": []gin.H{{"url": "https://example.test/image.png"}}})
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.Use(func(c *gin.Context) {
|
||||
groupID := int64(3)
|
||||
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
|
||||
ID: 9,
|
||||
UserID: 7,
|
||||
GroupID: &groupID,
|
||||
Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI, AllowImageGeneration: true},
|
||||
})
|
||||
c.Next()
|
||||
})
|
||||
router.POST("/v1/images/generations/async", h.Submit)
|
||||
router.GET("/v1/images/tasks/:task_id", h.Get)
|
||||
|
||||
requestCtx, cancelRequest := context.WithCancel(context.Background())
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/images/generations/async", strings.NewReader(`{"model":"gpt-image-1","prompt":"cat"}`)).WithContext(requestCtx)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
require.Equal(t, http.StatusAccepted, w.Code)
|
||||
require.Equal(t, "no-store", w.Header().Get("Cache-Control"))
|
||||
require.Equal(t, "3", w.Header().Get("Retry-After"))
|
||||
|
||||
var accepted struct {
|
||||
TaskID string `json:"task_id"`
|
||||
Status string `json:"status"`
|
||||
PollURL string `json:"poll_url"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &accepted))
|
||||
require.Equal(t, service.ImageTaskStatusProcessing, accepted.Status)
|
||||
require.Equal(t, "/v1/images/tasks/"+accepted.TaskID, accepted.PollURL)
|
||||
require.Equal(t, accepted.PollURL, w.Header().Get("Location"))
|
||||
|
||||
// The detached background request must survive completion of/cancellation
|
||||
// from the short submission request.
|
||||
cancelRequest()
|
||||
close(release)
|
||||
require.Eventually(t, func() bool {
|
||||
got, err := tasks.Get(context.Background(), service.ImageTaskOwner{UserID: 7, APIKeyID: 9}, accepted.TaskID)
|
||||
return err == nil && got.Status == service.ImageTaskStatusCompleted
|
||||
}, time.Second, 10*time.Millisecond)
|
||||
|
||||
pollReq := httptest.NewRequest(http.MethodGet, accepted.PollURL, nil)
|
||||
pollWriter := httptest.NewRecorder()
|
||||
router.ServeHTTP(pollWriter, pollReq)
|
||||
require.Equal(t, http.StatusOK, pollWriter.Code)
|
||||
require.Equal(t, "no-store", pollWriter.Header().Get("Cache-Control"))
|
||||
require.Empty(t, pollWriter.Header().Get("Retry-After"))
|
||||
require.Contains(t, pollWriter.Body.String(), "https://example.test/image.png")
|
||||
}
|
||||
|
||||
// When object storage is not configured the feature is fully disabled: the
|
||||
// endpoints must return 404 without creating a task or writing to Redis.
|
||||
func TestAsyncImageHandlerDisabledReturns404(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
store := &asyncImageMemoryStore{tasks: make(map[string]*service.ImageTaskRecord)}
|
||||
tasks := service.NewImageTaskServiceWithOptions(store, time.Hour, time.Minute) // enabled == false
|
||||
h := &AsyncImageHandler{tasks: tasks}
|
||||
|
||||
router := gin.New()
|
||||
router.Use(func(c *gin.Context) {
|
||||
groupID := int64(3)
|
||||
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
|
||||
ID: 9,
|
||||
UserID: 7,
|
||||
GroupID: &groupID,
|
||||
Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI, AllowImageGeneration: true},
|
||||
})
|
||||
c.Next()
|
||||
})
|
||||
router.POST("/v1/images/generations/async", h.Submit)
|
||||
router.GET("/v1/images/tasks/:task_id", h.Get)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/v1/images/generations/async", strings.NewReader(`{"model":"gpt-image-1","prompt":"cat"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
router.ServeHTTP(w, req)
|
||||
require.Equal(t, http.StatusNotFound, w.Code)
|
||||
require.Contains(t, w.Body.String(), "not enabled")
|
||||
|
||||
pollReq := httptest.NewRequest(http.MethodGet, "/v1/images/tasks/imgtask_missing", nil)
|
||||
pollWriter := httptest.NewRecorder()
|
||||
router.ServeHTTP(pollWriter, pollReq)
|
||||
require.Equal(t, http.StatusNotFound, pollWriter.Code)
|
||||
|
||||
// No task was created / persisted.
|
||||
require.Empty(t, store.tasks)
|
||||
}
|
||||
@@ -119,6 +119,7 @@ func ProvideHandlers(
|
||||
paymentHandler *PaymentHandler,
|
||||
paymentWebhookHandler *PaymentWebhookHandler,
|
||||
availableChannelHandler *AvailableChannelHandler,
|
||||
asyncImageHandler *AsyncImageHandler,
|
||||
batchImageHandler *BatchImageHandler,
|
||||
_ *service.IdempotencyCoordinator,
|
||||
_ *service.IdempotencyCleanupService,
|
||||
@@ -140,6 +141,7 @@ func ProvideHandlers(
|
||||
Payment: paymentHandler,
|
||||
PaymentWebhook: paymentWebhookHandler,
|
||||
AvailableChannel: availableChannelHandler,
|
||||
AsyncImage: asyncImageHandler,
|
||||
BatchImage: batchImageHandler,
|
||||
}
|
||||
}
|
||||
@@ -162,6 +164,7 @@ var ProviderSet = wire.NewSet(
|
||||
NewPaymentHandler,
|
||||
NewPaymentWebhookHandler,
|
||||
NewAvailableChannelHandler,
|
||||
NewAsyncImageHandler,
|
||||
NewBatchImageHandler,
|
||||
|
||||
// Admin handlers
|
||||
|
||||
@@ -8,10 +8,6 @@ import (
|
||||
"path"
|
||||
"time"
|
||||
|
||||
"github.com/aws/aws-sdk-go-v2/aws"
|
||||
v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4"
|
||||
awsconfig "github.com/aws/aws-sdk-go-v2/config"
|
||||
"github.com/aws/aws-sdk-go-v2/credentials"
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
|
||||
@@ -27,32 +23,16 @@ type S3BackupStore struct {
|
||||
// NewS3BackupStoreFactory returns a BackupObjectStoreFactory that creates S3-backed stores
|
||||
func NewS3BackupStoreFactory() service.BackupObjectStoreFactory {
|
||||
return func(ctx context.Context, cfg *service.BackupS3Config) (service.BackupObjectStore, error) {
|
||||
region := cfg.Region
|
||||
if region == "" {
|
||||
region = "auto" // Cloudflare R2 默认 region
|
||||
}
|
||||
|
||||
awsCfg, err := awsconfig.LoadDefaultConfig(ctx,
|
||||
awsconfig.WithRegion(region),
|
||||
awsconfig.WithCredentialsProvider(
|
||||
credentials.NewStaticCredentialsProvider(cfg.AccessKeyID, cfg.SecretAccessKey, ""),
|
||||
),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load aws config: %w", err)
|
||||
}
|
||||
|
||||
client := s3.NewFromConfig(awsCfg, func(o *s3.Options) {
|
||||
if cfg.Endpoint != "" {
|
||||
o.BaseEndpoint = &cfg.Endpoint
|
||||
}
|
||||
if cfg.ForcePathStyle {
|
||||
o.UsePathStyle = true
|
||||
}
|
||||
o.APIOptions = append(o.APIOptions, v4.SwapComputePayloadSHA256ForUnsignedPayloadMiddleware)
|
||||
o.RequestChecksumCalculation = aws.RequestChecksumCalculationWhenRequired
|
||||
client, err := newS3Client(ctx, s3ClientParams{
|
||||
Endpoint: cfg.Endpoint,
|
||||
Region: cfg.Region,
|
||||
AccessKeyID: cfg.AccessKeyID,
|
||||
SecretAccessKey: cfg.SecretAccessKey,
|
||||
ForcePathStyle: cfg.ForcePathStyle,
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &S3BackupStore{client: client, bucket: cfg.Bucket}, nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/servertiming"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
)
|
||||
|
||||
// S3ImageStorage 用 S3 兼容对象存储实现 service.ImageStorage。
|
||||
type S3ImageStorage struct {
|
||||
client *s3.Client
|
||||
bucket string
|
||||
publicBaseURL string
|
||||
presignExpiry time.Duration
|
||||
}
|
||||
|
||||
var _ service.ImageStorage = (*S3ImageStorage)(nil)
|
||||
|
||||
// NewS3ImageStorage 依据配置构造 S3 图片存储(调用方应先确认 cfg.Active())。
|
||||
func NewS3ImageStorage(ctx context.Context, cfg *config.ImageStorageConfig) (*S3ImageStorage, error) {
|
||||
client, err := newS3Client(ctx, s3ClientParams{
|
||||
Endpoint: cfg.Endpoint,
|
||||
Region: cfg.Region,
|
||||
AccessKeyID: cfg.AccessKeyID,
|
||||
SecretAccessKey: cfg.SecretAccessKey,
|
||||
ForcePathStyle: cfg.ForcePathStyle,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
expiry := time.Duration(cfg.PresignExpiry) * time.Hour
|
||||
if expiry <= 0 {
|
||||
expiry = 24 * time.Hour
|
||||
}
|
||||
|
||||
return &S3ImageStorage{
|
||||
client: client,
|
||||
bucket: cfg.Bucket,
|
||||
publicBaseURL: strings.TrimRight(cfg.PublicBaseURL, "/"),
|
||||
presignExpiry: expiry,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Save 上传图片字节,返回可访问 URL:配了 public_base_url 则返回公开直链,否则返回 presigned 临时链接。
|
||||
func (s *S3ImageStorage) Save(ctx context.Context, key, contentType string, data []byte) (string, error) {
|
||||
finish := servertiming.ObserveDependency(ctx, "s3")
|
||||
_, err := s.client.PutObject(ctx, &s3.PutObjectInput{
|
||||
Bucket: &s.bucket,
|
||||
Key: &key,
|
||||
Body: bytes.NewReader(data),
|
||||
ContentType: &contentType,
|
||||
})
|
||||
finish()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("S3 PutObject: %w", err)
|
||||
}
|
||||
|
||||
if s.publicBaseURL != "" {
|
||||
return s.publicBaseURL + "/" + strings.TrimLeft(key, "/"), nil
|
||||
}
|
||||
|
||||
presignClient := s3.NewPresignClient(s.client)
|
||||
result, err := presignClient.PresignGetObject(ctx, &s3.GetObjectInput{
|
||||
Bucket: &s.bucket,
|
||||
Key: &key,
|
||||
}, s3.WithPresignExpires(s.presignExpiry))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("presign url: %w", err)
|
||||
}
|
||||
return result.URL, nil
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
const imageTaskKeyPrefix = "image_task:"
|
||||
|
||||
type imageTaskStore struct {
|
||||
rdb *redis.Client
|
||||
}
|
||||
|
||||
func NewImageTaskStore(rdb *redis.Client) service.ImageTaskStore {
|
||||
return &imageTaskStore{rdb: rdb}
|
||||
}
|
||||
|
||||
func (s *imageTaskStore) Save(ctx context.Context, task *service.ImageTaskRecord, ttl time.Duration) error {
|
||||
data, err := json.Marshal(task)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.rdb.Set(ctx, imageTaskKey(task.ID), data, ttl).Err()
|
||||
}
|
||||
|
||||
func (s *imageTaskStore) Get(ctx context.Context, id string) (*service.ImageTaskRecord, error) {
|
||||
data, err := s.rdb.Get(ctx, imageTaskKey(id)).Bytes()
|
||||
if err != nil {
|
||||
if err == redis.Nil {
|
||||
return nil, service.ErrImageTaskNotFound
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
var task service.ImageTaskRecord
|
||||
if err := json.Unmarshal(data, &task); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &task, nil
|
||||
}
|
||||
|
||||
func imageTaskKey(id string) string {
|
||||
return imageTaskKeyPrefix + strings.TrimSpace(id)
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestImageTaskStoreRoundTripAndTTL(t *testing.T) {
|
||||
mr := miniredis.RunT(t)
|
||||
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
|
||||
t.Cleanup(func() { _ = rdb.Close() })
|
||||
store := NewImageTaskStore(rdb)
|
||||
task := &service.ImageTaskRecord{
|
||||
ID: "imgtask_123",
|
||||
UserID: 7,
|
||||
APIKeyID: 9,
|
||||
Status: service.ImageTaskStatusProcessing,
|
||||
CreatedAt: 100,
|
||||
ExpiresAt: 200,
|
||||
}
|
||||
|
||||
require.NoError(t, store.Save(context.Background(), task, 24*time.Hour))
|
||||
got, err := store.Get(context.Background(), task.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, task, got)
|
||||
require.Equal(t, 24*time.Hour, mr.TTL(imageTaskKey(task.ID)))
|
||||
}
|
||||
|
||||
func TestImageTaskStoreMissing(t *testing.T) {
|
||||
mr := miniredis.RunT(t)
|
||||
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
|
||||
t.Cleanup(func() { _ = rdb.Close() })
|
||||
store := NewImageTaskStore(rdb)
|
||||
|
||||
_, err := store.Get(context.Background(), "imgtask_missing")
|
||||
require.ErrorIs(t, err, service.ErrImageTaskNotFound)
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/aws/aws-sdk-go-v2/aws"
|
||||
v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4"
|
||||
awsconfig "github.com/aws/aws-sdk-go-v2/config"
|
||||
"github.com/aws/aws-sdk-go-v2/credentials"
|
||||
"github.com/aws/aws-sdk-go-v2/service/s3"
|
||||
)
|
||||
|
||||
// s3ClientParams 描述构造 S3 兼容客户端所需的参数。
|
||||
type s3ClientParams struct {
|
||||
Endpoint string
|
||||
Region string
|
||||
AccessKeyID string
|
||||
SecretAccessKey string
|
||||
ForcePathStyle bool
|
||||
}
|
||||
|
||||
// newS3Client 构造一个 S3 兼容客户端,兼容 AWS S3 / Cloudflare R2 / 阿里云 OSS / MinIO。
|
||||
//
|
||||
// 通过 SwapComputePayloadSHA256ForUnsignedPayloadMiddleware + RequestChecksumCalculationWhenRequired
|
||||
// 规避阿里云 OSS 不兼容 s3manager 分片签名的问题(backup 与 image storage 共用此构造)。
|
||||
func newS3Client(ctx context.Context, p s3ClientParams) (*s3.Client, error) {
|
||||
region := p.Region
|
||||
if region == "" {
|
||||
region = "auto" // Cloudflare R2 默认 region
|
||||
}
|
||||
|
||||
awsCfg, err := awsconfig.LoadDefaultConfig(ctx,
|
||||
awsconfig.WithRegion(region),
|
||||
awsconfig.WithCredentialsProvider(
|
||||
credentials.NewStaticCredentialsProvider(p.AccessKeyID, p.SecretAccessKey, ""),
|
||||
),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load aws config: %w", err)
|
||||
}
|
||||
|
||||
return s3.NewFromConfig(awsCfg, func(o *s3.Options) {
|
||||
if p.Endpoint != "" {
|
||||
o.BaseEndpoint = &p.Endpoint
|
||||
}
|
||||
if p.ForcePathStyle {
|
||||
o.UsePathStyle = true
|
||||
}
|
||||
o.APIOptions = append(o.APIOptions, v4.SwapComputePayloadSHA256ForUnsignedPayloadMiddleware)
|
||||
o.RequestChecksumCalculation = aws.RequestChecksumCalculationWhenRequired
|
||||
}), nil
|
||||
}
|
||||
@@ -1,6 +1,7 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
|
||||
@@ -118,6 +119,7 @@ var ProviderSet = wire.NewSet(
|
||||
NewRedeemCache,
|
||||
NewUpdateCache,
|
||||
NewGeminiTokenCache,
|
||||
NewImageTaskStore,
|
||||
NewBatchImageQueue,
|
||||
NewBatchImageDownloadLimiter,
|
||||
NewLeaderLockCache,
|
||||
@@ -137,6 +139,9 @@ var ProviderSet = wire.NewSet(
|
||||
NewPgDumper,
|
||||
NewS3BackupStoreFactory,
|
||||
|
||||
// Image storage (async image task result offload)
|
||||
ProvideImageStorage,
|
||||
|
||||
// HTTP service ports (DI Strategy A: return interface directly)
|
||||
NewTurnstileVerifier,
|
||||
ProvidePricingRemoteClient,
|
||||
@@ -168,6 +173,19 @@ func ProvideEnt(cfg *config.Config) (*ent.Client, error) {
|
||||
return client, err
|
||||
}
|
||||
|
||||
// ProvideImageStorage 提供异步图片任务结果转存所用的对象存储实现。
|
||||
// 仅当开关打开且 S3 凭证齐全时返回具体实现,否则返回 nil(功能整体禁用)。
|
||||
func ProvideImageStorage(cfg *config.Config) (service.ImageStorage, error) {
|
||||
if !cfg.ImageStorage.Active() {
|
||||
return nil, nil
|
||||
}
|
||||
store, err := NewS3ImageStorage(context.Background(), &cfg.ImageStorage)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return store, nil
|
||||
}
|
||||
|
||||
// ProvideSQLDB 从 Ent 客户端提取底层的 *sql.DB 连接。
|
||||
//
|
||||
// 某些 Repository 需要直接执行原生 SQL(如复杂的批量更新、聚合查询),
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
@@ -25,8 +26,9 @@ func NewAPIKeyAuthMiddleware(apiKeyService *service.APIKeyService, subscriptionS
|
||||
// - 鉴权(Authentication):验证 Key 有效性、用户状态、IP 限制 —— 始终执行
|
||||
// - 计费执行(Billing Enforcement):过期/配额/订阅/余额检查 —— skipBilling 时整块跳过
|
||||
//
|
||||
// /v1/usage 和 /v1/sub2api/billing 端点只需鉴权,不需要计费执行。
|
||||
// 前者允许过期/配额耗尽的 Key 查询自身用量,后者用于读取当前 Key 的倍率配置。
|
||||
// /v1/usage、/v1/sub2api/billing 端点与异步生图任务查询只需鉴权,不需要计费执行。
|
||||
// usage 允许过期/配额耗尽的 Key 查询自身用量,billing 用于读取当前 Key 的倍率配置,
|
||||
// 异步生图查询允许已耗尽额度的 Key 拉取自身任务结果。
|
||||
func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscriptionService *service.SubscriptionService, cfg *config.Config) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// ── 1. 提取 API Key ──────────────────────────────────────────
|
||||
@@ -130,7 +132,10 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
|
||||
ctx := context.WithValue(c.Request.Context(), ctxkey.UserID, apiKey.User.ID)
|
||||
c.Request = c.Request.WithContext(ctx)
|
||||
billingInfoRequest := c.Request.URL.Path == "/v1/sub2api/billing"
|
||||
skipBilling := c.Request.URL.Path == "/v1/usage" || billingInfoRequest
|
||||
// Async image task polling only reads data that already belongs to the
|
||||
// authenticated key and must remain available after the completed
|
||||
// generation consumes the key's remaining balance.
|
||||
skipBilling := c.Request.URL.Path == "/v1/usage" || billingInfoRequest || isAsyncImageTaskRead(c.Request.Method, c.Request.URL.Path)
|
||||
|
||||
// ── 4. SimpleMode → early return ─────────────────────────────
|
||||
|
||||
@@ -248,6 +253,13 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
|
||||
}
|
||||
}
|
||||
|
||||
func isAsyncImageTaskRead(method, path string) bool {
|
||||
if method != http.MethodGet {
|
||||
return false
|
||||
}
|
||||
return strings.HasPrefix(path, "/v1/images/tasks/") || strings.HasPrefix(path, "/images/tasks/")
|
||||
}
|
||||
|
||||
// GetAPIKeyFromContext 从上下文中获取API key
|
||||
func GetAPIKeyFromContext(c *gin.Context) (*service.APIKey, bool) {
|
||||
value, exists := c.Get(string(ContextKeyAPIKey))
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestIsAsyncImageTaskRead(t *testing.T) {
|
||||
require.True(t, isAsyncImageTaskRead(http.MethodGet, "/v1/images/tasks/imgtask_123"))
|
||||
require.True(t, isAsyncImageTaskRead(http.MethodGet, "/images/tasks/imgtask_123"))
|
||||
require.False(t, isAsyncImageTaskRead(http.MethodPost, "/v1/images/tasks/imgtask_123"))
|
||||
require.False(t, isAsyncImageTaskRead(http.MethodGet, "/v1/images/generations"))
|
||||
}
|
||||
@@ -192,6 +192,9 @@ func RegisterGatewayRoutes(
|
||||
})
|
||||
gateway.POST("/images/generations", imagesHandler)
|
||||
gateway.POST("/images/edits", imagesHandler)
|
||||
gateway.POST("/images/generations/async", h.AsyncImage.Submit)
|
||||
gateway.POST("/images/edits/async", h.AsyncImage.Submit)
|
||||
gateway.GET("/images/tasks/:task_id", h.AsyncImage.Get)
|
||||
gateway.POST("/images/batches", h.BatchImage.Submit)
|
||||
gateway.GET("/images/batches", h.BatchImage.List)
|
||||
gateway.GET("/images/batches/models", h.BatchImage.Models)
|
||||
@@ -272,6 +275,9 @@ func RegisterGatewayRoutes(
|
||||
})
|
||||
r.POST("/images/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, imagesHandler)
|
||||
r.POST("/images/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, imagesHandler)
|
||||
r.POST("/images/generations/async", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.AsyncImage.Submit)
|
||||
r.POST("/images/edits/async", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.AsyncImage.Submit)
|
||||
r.GET("/images/tasks/:task_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.AsyncImage.Get)
|
||||
r.POST("/videos/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoGenerationHandler)
|
||||
r.POST("/videos/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoEditHandler)
|
||||
r.POST("/videos/extensions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, videoExtensionHandler)
|
||||
|
||||
@@ -28,6 +28,7 @@ func newGatewayRoutesTestRouter(platform ...string) *gin.Engine {
|
||||
&handler.Handlers{
|
||||
Gateway: &handler.GatewayHandler{},
|
||||
OpenAIGateway: &handler.OpenAIGatewayHandler{},
|
||||
AsyncImage: handler.NewAsyncImageHandler(nil, nil),
|
||||
},
|
||||
servermiddleware.APIKeyAuthMiddleware(func(c *gin.Context) {
|
||||
groupID := int64(1)
|
||||
@@ -113,6 +114,25 @@ func TestGatewayRoutesOpenAIImagesPathsAreRegistered(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayRoutesAsyncImagesPathsAreRegistered(t *testing.T) {
|
||||
router := newGatewayRoutesTestRouter()
|
||||
registered := make(map[string]bool)
|
||||
for _, route := range router.Routes() {
|
||||
registered[route.Method+" "+route.Path] = true
|
||||
}
|
||||
|
||||
for _, route := range []string{
|
||||
"POST /v1/images/generations/async",
|
||||
"POST /v1/images/edits/async",
|
||||
"GET /v1/images/tasks/:task_id",
|
||||
"POST /images/generations/async",
|
||||
"POST /images/edits/async",
|
||||
"GET /images/tasks/:task_id",
|
||||
} {
|
||||
require.True(t, registered[route], "%s should be registered", route)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGatewayRoutesGrokImagesAndVideosPathsAreRegistered(t *testing.T) {
|
||||
router := newGatewayRoutesTestRouter(service.PlatformGrok)
|
||||
|
||||
|
||||
@@ -0,0 +1,192 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const defaultImageMaxDownloadBytes int64 = 32 << 20 // 32 MiB
|
||||
|
||||
// ImageStorage 把图片字节写入对象存储并返回可访问 URL。
|
||||
//
|
||||
// 这是对象存储的可插拔抽象:适配一个新的对象存储厂商,只需实现本接口
|
||||
// (例如包一个厂商 SDK),无需改动任务/网关逻辑。仓库内自带一个 S3 兼容实现
|
||||
// (repository.S3ImageStorage),适用于 AWS S3 / Cloudflare R2 / 阿里云 OSS / MinIO 等。
|
||||
type ImageStorage interface {
|
||||
// Save 把 data 以 key 存入对象存储,返回可下载的 URL(公开直链或 presigned 临时链接)。
|
||||
// contentType 为图片 MIME 类型,如 "image/png"。
|
||||
Save(ctx context.Context, key, contentType string, data []byte) (url string, err error)
|
||||
}
|
||||
|
||||
// ImageResultUploader 是 ImageStorage 的上层编排器(与具体厂商无关):
|
||||
// 把上游生图响应里的每张图片(b64_json 解码 / url 下载)转存到对象存储,
|
||||
// 并把响应结果改写为只含短链接的紧凑 JSON,从而避免大 base64 落 Redis。
|
||||
type ImageResultUploader struct {
|
||||
storage ImageStorage
|
||||
httpClient *http.Client
|
||||
prefix string
|
||||
maxDownloadBytes int64
|
||||
}
|
||||
|
||||
// NewImageResultUploader 构造一个 uploader;storage 为 nil 时 Rewrite 直接透传。
|
||||
func NewImageResultUploader(storage ImageStorage, prefix string, maxDownloadBytes int64, httpClient *http.Client) *ImageResultUploader {
|
||||
if httpClient == nil {
|
||||
httpClient = defaultImageDownloadHTTPClient()
|
||||
}
|
||||
if maxDownloadBytes <= 0 {
|
||||
maxDownloadBytes = defaultImageMaxDownloadBytes
|
||||
}
|
||||
return &ImageResultUploader{
|
||||
storage: storage,
|
||||
httpClient: httpClient,
|
||||
prefix: prefix,
|
||||
maxDownloadBytes: maxDownloadBytes,
|
||||
}
|
||||
}
|
||||
|
||||
func defaultImageDownloadHTTPClient() *http.Client {
|
||||
return &http.Client{Timeout: 60 * time.Second}
|
||||
}
|
||||
|
||||
// Rewrite 将 result(上游生图响应 JSON)里的每张图片转存到对象存储,
|
||||
// 返回改写后的紧凑结果(data[i].url 指向对象存储,b64_json 被移除)。
|
||||
// 任一图片转存失败即返回 error(调用方据此将任务标记为失败,绝不把大 blob 落 Redis)。
|
||||
func (u *ImageResultUploader) Rewrite(ctx context.Context, taskID string, result json.RawMessage) (json.RawMessage, error) {
|
||||
if u == nil || u.storage == nil {
|
||||
return result, nil
|
||||
}
|
||||
var top map[string]json.RawMessage
|
||||
if err := json.Unmarshal(result, &top); err != nil {
|
||||
return nil, fmt.Errorf("parse image response: %w", err)
|
||||
}
|
||||
rawData, ok := top["data"]
|
||||
if !ok {
|
||||
// 没有 data 数组(结构不符合预期),保持原样返回,交由上层决定。
|
||||
return result, nil
|
||||
}
|
||||
var items []map[string]json.RawMessage
|
||||
if err := json.Unmarshal(rawData, &items); err != nil {
|
||||
return nil, fmt.Errorf("parse image response data: %w", err)
|
||||
}
|
||||
if len(items) == 0 {
|
||||
return result, nil
|
||||
}
|
||||
for i, item := range items {
|
||||
data, contentType, err := u.fetchImageBytes(ctx, item)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("image %d: %w", i, err)
|
||||
}
|
||||
key := u.buildKey(taskID, i, contentType)
|
||||
url, err := u.storage.Save(ctx, key, contentType, data)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("image %d: upload to object storage: %w", i, err)
|
||||
}
|
||||
urlRaw, err := json.Marshal(url)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("image %d: encode url: %w", i, err)
|
||||
}
|
||||
item["url"] = urlRaw
|
||||
delete(item, "b64_json")
|
||||
items[i] = item
|
||||
}
|
||||
newData, err := json.Marshal(items)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encode image response data: %w", err)
|
||||
}
|
||||
top["data"] = newData
|
||||
out, err := json.Marshal(top)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encode image response: %w", err)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (u *ImageResultUploader) fetchImageBytes(ctx context.Context, item map[string]json.RawMessage) ([]byte, string, error) {
|
||||
if raw, ok := item["b64_json"]; ok {
|
||||
var b64 string
|
||||
if err := json.Unmarshal(raw, &b64); err == nil {
|
||||
if b64 = strings.TrimSpace(b64); b64 != "" {
|
||||
data, err := base64.StdEncoding.DecodeString(b64)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("decode b64_json: %w", err)
|
||||
}
|
||||
return data, detectImageContentType(data), nil
|
||||
}
|
||||
}
|
||||
}
|
||||
if raw, ok := item["url"]; ok {
|
||||
var rawURL string
|
||||
if err := json.Unmarshal(raw, &rawURL); err == nil {
|
||||
if rawURL = strings.TrimSpace(rawURL); rawURL != "" {
|
||||
return u.download(ctx, rawURL)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil, "", errors.New("image item has neither b64_json nor url")
|
||||
}
|
||||
|
||||
func (u *ImageResultUploader) download(ctx context.Context, rawURL string) ([]byte, string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("build download request: %w", err)
|
||||
}
|
||||
resp, err := u.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("download image: %w", err)
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
|
||||
return nil, "", fmt.Errorf("download image: unexpected status %d", resp.StatusCode)
|
||||
}
|
||||
limit := u.maxDownloadBytes
|
||||
if limit <= 0 {
|
||||
limit = defaultImageMaxDownloadBytes
|
||||
}
|
||||
data, err := io.ReadAll(io.LimitReader(resp.Body, limit+1))
|
||||
if err != nil {
|
||||
return nil, "", fmt.Errorf("read image body: %w", err)
|
||||
}
|
||||
if int64(len(data)) > limit {
|
||||
return nil, "", fmt.Errorf("downloaded image exceeds %d bytes", limit)
|
||||
}
|
||||
contentType := strings.TrimSpace(strings.Split(resp.Header.Get("Content-Type"), ";")[0])
|
||||
if !strings.HasPrefix(contentType, "image/") {
|
||||
contentType = detectImageContentType(data)
|
||||
}
|
||||
return data, contentType, nil
|
||||
}
|
||||
|
||||
func (u *ImageResultUploader) buildKey(taskID string, index int, contentType string) string {
|
||||
return u.prefix + taskID + "-" + strconv.Itoa(index) + extensionForContentType(contentType)
|
||||
}
|
||||
|
||||
func detectImageContentType(data []byte) string {
|
||||
ct := strings.TrimSpace(strings.Split(http.DetectContentType(data), ";")[0])
|
||||
if strings.HasPrefix(ct, "image/") {
|
||||
return ct
|
||||
}
|
||||
return "image/png"
|
||||
}
|
||||
|
||||
func extensionForContentType(ct string) string {
|
||||
switch {
|
||||
case strings.Contains(ct, "png"):
|
||||
return ".png"
|
||||
case strings.Contains(ct, "jpeg"), strings.Contains(ct, "jpg"):
|
||||
return ".jpg"
|
||||
case strings.Contains(ct, "webp"):
|
||||
return ".webp"
|
||||
case strings.Contains(ct, "gif"):
|
||||
return ".gif"
|
||||
default:
|
||||
return ".png"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// pngBytes is a minimal payload whose signature makes http.DetectContentType
|
||||
// report image/png.
|
||||
var pngBytes = []byte("\x89PNG\r\n\x1a\nfake-png-payload")
|
||||
|
||||
type savedImage struct {
|
||||
key string
|
||||
contentType string
|
||||
data []byte
|
||||
}
|
||||
|
||||
type fakeImageStorage struct {
|
||||
saved []savedImage
|
||||
url string
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *fakeImageStorage) Save(_ context.Context, key, contentType string, data []byte) (string, error) {
|
||||
if f.err != nil {
|
||||
return "", f.err
|
||||
}
|
||||
f.saved = append(f.saved, savedImage{key: key, contentType: contentType, data: append([]byte(nil), data...)})
|
||||
if f.url != "" {
|
||||
return f.url, nil
|
||||
}
|
||||
return "https://cdn.test/" + key, nil
|
||||
}
|
||||
|
||||
func TestImageResultUploaderRewritesB64JSON(t *testing.T) {
|
||||
storage := &fakeImageStorage{}
|
||||
uploader := NewImageResultUploader(storage, "images/", 0, nil)
|
||||
|
||||
b64 := base64.StdEncoding.EncodeToString(pngBytes)
|
||||
result := json.RawMessage(`{"created":1,"data":[{"b64_json":"` + b64 + `","revised_prompt":"a cat"}]}`)
|
||||
|
||||
out, err := uploader.Rewrite(context.Background(), "imgtask_abc", result)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Len(t, storage.saved, 1)
|
||||
require.Equal(t, "images/imgtask_abc-0.png", storage.saved[0].key)
|
||||
require.Equal(t, "image/png", storage.saved[0].contentType)
|
||||
require.Equal(t, pngBytes, storage.saved[0].data)
|
||||
|
||||
var parsed struct {
|
||||
Data []map[string]json.RawMessage `json:"data"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(out, &parsed))
|
||||
require.Len(t, parsed.Data, 1)
|
||||
require.JSONEq(t, `"https://cdn.test/images/imgtask_abc-0.png"`, string(parsed.Data[0]["url"]))
|
||||
_, hasB64 := parsed.Data[0]["b64_json"]
|
||||
require.False(t, hasB64, "b64_json must be stripped after offload")
|
||||
require.JSONEq(t, `"a cat"`, string(parsed.Data[0]["revised_prompt"]), "unrelated fields preserved")
|
||||
}
|
||||
|
||||
func TestImageResultUploaderRewritesURL(t *testing.T) {
|
||||
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Type", "image/png")
|
||||
_, _ = w.Write(pngBytes)
|
||||
}))
|
||||
defer upstream.Close()
|
||||
|
||||
storage := &fakeImageStorage{}
|
||||
uploader := NewImageResultUploader(storage, "images/", 0, nil)
|
||||
|
||||
result := json.RawMessage(`{"created":1,"data":[{"url":"` + upstream.URL + `/pic.png"}]}`)
|
||||
out, err := uploader.Rewrite(context.Background(), "imgtask_xyz", result)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Len(t, storage.saved, 1)
|
||||
require.Equal(t, pngBytes, storage.saved[0].data)
|
||||
require.Equal(t, "image/png", storage.saved[0].contentType)
|
||||
|
||||
var parsed struct {
|
||||
Data []map[string]json.RawMessage `json:"data"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(out, &parsed))
|
||||
require.JSONEq(t, `"https://cdn.test/images/imgtask_xyz-0.png"`, string(parsed.Data[0]["url"]))
|
||||
}
|
||||
|
||||
func TestImageResultUploaderPropagatesStorageError(t *testing.T) {
|
||||
storage := &fakeImageStorage{err: errors.New("bucket unreachable")}
|
||||
uploader := NewImageResultUploader(storage, "images/", 0, nil)
|
||||
|
||||
b64 := base64.StdEncoding.EncodeToString(pngBytes)
|
||||
result := json.RawMessage(`{"data":[{"b64_json":"` + b64 + `"}]}`)
|
||||
|
||||
_, err := uploader.Rewrite(context.Background(), "imgtask_err", result)
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "bucket unreachable")
|
||||
}
|
||||
|
||||
func TestImageResultUploaderNilStoragePassthrough(t *testing.T) {
|
||||
var uploader *ImageResultUploader
|
||||
result := json.RawMessage(`{"data":[{"url":"https://example.test/x.png"}]}`)
|
||||
out, err := uploader.Rewrite(context.Background(), "imgtask_nil", result)
|
||||
require.NoError(t, err)
|
||||
require.JSONEq(t, string(result), string(out))
|
||||
}
|
||||
|
||||
func TestImageTaskServiceCompleteOffloadsToStorage(t *testing.T) {
|
||||
store := &imageTaskMemoryStore{}
|
||||
storage := &fakeImageStorage{}
|
||||
uploader := NewImageResultUploader(storage, "images/", 0, nil)
|
||||
svc := NewImageTaskServiceWithUploader(store, uploader, time.Hour, time.Minute)
|
||||
require.True(t, svc.Enabled())
|
||||
|
||||
owner := ImageTaskOwner{UserID: 1, APIKeyID: 2}
|
||||
created, err := svc.Create(context.Background(), owner)
|
||||
require.NoError(t, err)
|
||||
|
||||
b64 := base64.StdEncoding.EncodeToString(pngBytes)
|
||||
result := json.RawMessage(`{"created":1,"data":[{"b64_json":"` + b64 + `"}]}`)
|
||||
require.NoError(t, svc.Complete(context.Background(), created.ID, http.StatusOK, result))
|
||||
|
||||
got, err := svc.Get(context.Background(), owner, created.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ImageTaskStatusCompleted, got.Status)
|
||||
require.Equal(t, "https://cdn.test/images/"+created.ID+"-0.png", got.ImageURL)
|
||||
require.NotContains(t, string(got.Result), "b64_json", "large base64 must not be persisted to Redis")
|
||||
require.Len(t, storage.saved, 1)
|
||||
}
|
||||
|
||||
func TestImageTaskServiceCompleteOffloadFailureMarksFailed(t *testing.T) {
|
||||
store := &imageTaskMemoryStore{}
|
||||
storage := &fakeImageStorage{err: errors.New("bucket unreachable")}
|
||||
uploader := NewImageResultUploader(storage, "images/", 0, nil)
|
||||
svc := NewImageTaskServiceWithUploader(store, uploader, time.Hour, time.Minute)
|
||||
|
||||
owner := ImageTaskOwner{UserID: 1, APIKeyID: 2}
|
||||
created, err := svc.Create(context.Background(), owner)
|
||||
require.NoError(t, err)
|
||||
|
||||
b64 := base64.StdEncoding.EncodeToString(pngBytes)
|
||||
result := json.RawMessage(`{"data":[{"b64_json":"` + b64 + `"}]}`)
|
||||
require.NoError(t, svc.Complete(context.Background(), created.ID, http.StatusOK, result))
|
||||
|
||||
got, err := svc.Get(context.Background(), owner, created.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ImageTaskStatusFailed, got.Status)
|
||||
require.Equal(t, http.StatusBadGateway, got.HTTPStatus)
|
||||
require.Contains(t, string(got.Error), "object storage")
|
||||
require.NotContains(t, string(got.Result), "b64_json", "failed offload must not persist base64 to Redis")
|
||||
}
|
||||
@@ -0,0 +1,238 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
|
||||
"github.com/google/uuid"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
const (
|
||||
ImageTaskStatusProcessing = "processing"
|
||||
ImageTaskStatusCompleted = "completed"
|
||||
ImageTaskStatusFailed = "failed"
|
||||
|
||||
defaultImageTaskTTL = 24 * time.Hour
|
||||
defaultImageTaskExecutionTimeout = 30 * time.Minute
|
||||
)
|
||||
|
||||
var (
|
||||
ErrImageTaskNotFound = infraerrors.New(http.StatusNotFound, "IMAGE_TASK_NOT_FOUND", "image task not found")
|
||||
ErrImageTaskForbidden = infraerrors.New(http.StatusForbidden, "IMAGE_TASK_FORBIDDEN", "image task does not belong to this API key")
|
||||
ErrImageTaskUnavailable = infraerrors.New(http.StatusServiceUnavailable, "IMAGE_TASK_UNAVAILABLE", "image task storage is unavailable")
|
||||
)
|
||||
|
||||
// ImageTaskRecord is the private Redis representation of an asynchronous image
|
||||
// request. Ownership fields are intentionally omitted from the public view.
|
||||
type ImageTaskRecord struct {
|
||||
ID string `json:"id"`
|
||||
UserID int64 `json:"user_id"`
|
||||
APIKeyID int64 `json:"api_key_id"`
|
||||
Status string `json:"status"`
|
||||
HTTPStatus int `json:"http_status,omitempty"`
|
||||
Result json.RawMessage `json:"result,omitempty"`
|
||||
Error json.RawMessage `json:"error,omitempty"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
CompletedAt *int64 `json:"completed_at,omitempty"`
|
||||
ExpiresAt int64 `json:"expires_at"`
|
||||
}
|
||||
|
||||
// ImageTask is the API-safe task representation returned to callers.
|
||||
type ImageTask struct {
|
||||
ID string `json:"id"`
|
||||
TaskID string `json:"task_id"`
|
||||
Object string `json:"object"`
|
||||
Status string `json:"status"`
|
||||
HTTPStatus int `json:"http_status,omitempty"`
|
||||
ImageURL string `json:"image_url,omitempty"`
|
||||
Result json.RawMessage `json:"result,omitempty"`
|
||||
Error json.RawMessage `json:"error,omitempty"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
CompletedAt *int64 `json:"completed_at,omitempty"`
|
||||
ExpiresAt int64 `json:"expires_at"`
|
||||
}
|
||||
|
||||
type ImageTaskOwner struct {
|
||||
UserID int64
|
||||
APIKeyID int64
|
||||
}
|
||||
|
||||
type ImageTaskStore interface {
|
||||
Save(ctx context.Context, task *ImageTaskRecord, ttl time.Duration) error
|
||||
Get(ctx context.Context, id string) (*ImageTaskRecord, error)
|
||||
}
|
||||
|
||||
type ImageTaskService struct {
|
||||
store ImageTaskStore
|
||||
uploader *ImageResultUploader
|
||||
enabled bool
|
||||
ttl time.Duration
|
||||
executionTimeout time.Duration
|
||||
}
|
||||
|
||||
func NewImageTaskService(store ImageTaskStore) *ImageTaskService {
|
||||
return NewImageTaskServiceWithOptions(store, defaultImageTaskTTL, defaultImageTaskExecutionTimeout)
|
||||
}
|
||||
|
||||
func NewImageTaskServiceWithOptions(store ImageTaskStore, ttl, executionTimeout time.Duration) *ImageTaskService {
|
||||
if ttl <= 0 {
|
||||
ttl = defaultImageTaskTTL
|
||||
}
|
||||
if executionTimeout <= 0 {
|
||||
executionTimeout = defaultImageTaskExecutionTimeout
|
||||
}
|
||||
return &ImageTaskService{store: store, ttl: ttl, executionTimeout: executionTimeout}
|
||||
}
|
||||
|
||||
// NewImageTaskServiceWithUploader 构造一个已启用的图片任务服务:结果会先经 uploader
|
||||
// 转存到对象存储再落 Redis。uploader 为 nil 时不做转存(仅用于测试)。
|
||||
func NewImageTaskServiceWithUploader(store ImageTaskStore, uploader *ImageResultUploader, ttl, executionTimeout time.Duration) *ImageTaskService {
|
||||
s := NewImageTaskServiceWithOptions(store, ttl, executionTimeout)
|
||||
s.uploader = uploader
|
||||
s.enabled = true
|
||||
return s
|
||||
}
|
||||
|
||||
// Enabled 表示异步图片任务功能是否可用(总开关 + 凭证齐全)。
|
||||
// 关闭时 handler 直接返回 404,不创建任务、不写 Redis。
|
||||
func (s *ImageTaskService) Enabled() bool {
|
||||
return s != nil && s.enabled && s.store != nil
|
||||
}
|
||||
|
||||
func (s *ImageTaskService) ExecutionTimeout() time.Duration {
|
||||
if s == nil || s.executionTimeout <= 0 {
|
||||
return defaultImageTaskExecutionTimeout
|
||||
}
|
||||
return s.executionTimeout
|
||||
}
|
||||
|
||||
func (s *ImageTaskService) Create(ctx context.Context, owner ImageTaskOwner) (*ImageTask, error) {
|
||||
if s == nil || s.store == nil {
|
||||
return nil, ErrImageTaskUnavailable
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
task := &ImageTaskRecord{
|
||||
ID: "imgtask_" + strings.ReplaceAll(uuid.NewString(), "-", ""),
|
||||
UserID: owner.UserID,
|
||||
APIKeyID: owner.APIKeyID,
|
||||
Status: ImageTaskStatusProcessing,
|
||||
CreatedAt: now.Unix(),
|
||||
ExpiresAt: now.Add(s.ttl).Unix(),
|
||||
}
|
||||
if err := s.store.Save(ctx, task, s.ttl); err != nil {
|
||||
return nil, ErrImageTaskUnavailable.WithCause(err)
|
||||
}
|
||||
return imageTaskToPublic(task), nil
|
||||
}
|
||||
|
||||
func (s *ImageTaskService) Get(ctx context.Context, owner ImageTaskOwner, id string) (*ImageTask, error) {
|
||||
if s == nil || s.store == nil {
|
||||
return nil, ErrImageTaskUnavailable
|
||||
}
|
||||
task, err := s.store.Get(ctx, strings.TrimSpace(id))
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrImageTaskNotFound) {
|
||||
return nil, ErrImageTaskNotFound
|
||||
}
|
||||
return nil, ErrImageTaskUnavailable.WithCause(err)
|
||||
}
|
||||
if task.UserID != owner.UserID || task.APIKeyID != owner.APIKeyID {
|
||||
// Do not reveal whether a random task ID exists for another caller.
|
||||
return nil, ErrImageTaskNotFound
|
||||
}
|
||||
return imageTaskToPublic(task), nil
|
||||
}
|
||||
|
||||
func (s *ImageTaskService) Complete(ctx context.Context, id string, statusCode int, result json.RawMessage) error {
|
||||
if !json.Valid(result) {
|
||||
return s.Fail(ctx, id, http.StatusBadGateway, imageTaskErrorJSON("api_error", "upstream returned a non-JSON image response"))
|
||||
}
|
||||
if s.uploader != nil {
|
||||
rewritten, err := s.uploader.Rewrite(ctx, id, result)
|
||||
if err != nil {
|
||||
// 转存失败不回退存 base64,避免大 blob 撑爆 Redis:直接把任务标记为失败。
|
||||
logger.L().Error("image_task.offload_failed", zap.String("task_id", id), zap.Error(err))
|
||||
return s.Fail(ctx, id, http.StatusBadGateway, imageTaskErrorJSON("api_error", "failed to store generated image to object storage"))
|
||||
}
|
||||
result = rewritten
|
||||
}
|
||||
return s.finish(ctx, id, ImageTaskStatusCompleted, statusCode, result, nil)
|
||||
}
|
||||
|
||||
func (s *ImageTaskService) Fail(ctx context.Context, id string, statusCode int, taskErr json.RawMessage) error {
|
||||
if !json.Valid(taskErr) {
|
||||
taskErr = imageTaskErrorJSON("api_error", "image generation failed")
|
||||
}
|
||||
return s.finish(ctx, id, ImageTaskStatusFailed, statusCode, nil, taskErr)
|
||||
}
|
||||
|
||||
func (s *ImageTaskService) finish(ctx context.Context, id, status string, statusCode int, result, taskErr json.RawMessage) error {
|
||||
if s == nil || s.store == nil {
|
||||
return ErrImageTaskUnavailable
|
||||
}
|
||||
task, err := s.store.Get(ctx, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrImageTaskNotFound) {
|
||||
return ErrImageTaskNotFound
|
||||
}
|
||||
return ErrImageTaskUnavailable.WithCause(err)
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
completedAt := now.Unix()
|
||||
task.Status = status
|
||||
task.HTTPStatus = statusCode
|
||||
task.Result = result
|
||||
task.Error = taskErr
|
||||
task.CompletedAt = &completedAt
|
||||
task.ExpiresAt = now.Add(s.ttl).Unix()
|
||||
if err := s.store.Save(ctx, task, s.ttl); err != nil {
|
||||
return ErrImageTaskUnavailable.WithCause(err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func imageTaskToPublic(task *ImageTaskRecord) *ImageTask {
|
||||
if task == nil {
|
||||
return nil
|
||||
}
|
||||
return &ImageTask{
|
||||
ID: task.ID,
|
||||
TaskID: task.ID,
|
||||
Object: "image.generation.task",
|
||||
Status: task.Status,
|
||||
HTTPStatus: task.HTTPStatus,
|
||||
ImageURL: firstImageTaskURL(task.Result),
|
||||
Result: task.Result,
|
||||
Error: task.Error,
|
||||
CreatedAt: task.CreatedAt,
|
||||
CompletedAt: task.CompletedAt,
|
||||
ExpiresAt: task.ExpiresAt,
|
||||
}
|
||||
}
|
||||
|
||||
func firstImageTaskURL(result json.RawMessage) string {
|
||||
if len(result) == 0 || !json.Valid(result) {
|
||||
return ""
|
||||
}
|
||||
var response struct {
|
||||
Data []struct {
|
||||
URL string `json:"url"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if json.Unmarshal(result, &response) != nil || len(response.Data) == 0 {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(response.Data[0].URL)
|
||||
}
|
||||
|
||||
func imageTaskErrorJSON(errorType, message string) json.RawMessage {
|
||||
data, _ := json.Marshal(map[string]string{"type": errorType, "message": message})
|
||||
return data
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type imageTaskMemoryStore struct {
|
||||
task *ImageTaskRecord
|
||||
ttl time.Duration
|
||||
saveErr error
|
||||
getErr error
|
||||
}
|
||||
|
||||
func (s *imageTaskMemoryStore) Save(_ context.Context, task *ImageTaskRecord, ttl time.Duration) error {
|
||||
if s.saveErr != nil {
|
||||
return s.saveErr
|
||||
}
|
||||
copy := *task
|
||||
s.task = ©
|
||||
s.ttl = ttl
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *imageTaskMemoryStore) Get(_ context.Context, _ string) (*ImageTaskRecord, error) {
|
||||
if s.getErr != nil {
|
||||
return nil, s.getErr
|
||||
}
|
||||
if s.task == nil {
|
||||
return nil, ErrImageTaskNotFound
|
||||
}
|
||||
copy := *s.task
|
||||
return ©, nil
|
||||
}
|
||||
|
||||
func TestImageTaskServiceLifecycleAndOwnership(t *testing.T) {
|
||||
store := &imageTaskMemoryStore{}
|
||||
svc := NewImageTaskServiceWithOptions(store, time.Hour, 10*time.Minute)
|
||||
owner := ImageTaskOwner{UserID: 7, APIKeyID: 9}
|
||||
|
||||
created, err := svc.Create(context.Background(), owner)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ImageTaskStatusProcessing, created.Status)
|
||||
require.Equal(t, created.ID, created.TaskID)
|
||||
require.Equal(t, "image.generation.task", created.Object)
|
||||
require.Equal(t, time.Hour, store.ttl)
|
||||
require.Equal(t, owner.UserID, store.task.UserID)
|
||||
require.Equal(t, owner.APIKeyID, store.task.APIKeyID)
|
||||
|
||||
_, err = svc.Get(context.Background(), ImageTaskOwner{UserID: 7, APIKeyID: 10}, created.ID)
|
||||
require.ErrorIs(t, err, ErrImageTaskNotFound)
|
||||
|
||||
result := json.RawMessage(`{"created":123,"data":[{"url":"https://example.test/image.png"}]}`)
|
||||
require.NoError(t, svc.Complete(context.Background(), created.ID, http.StatusOK, result))
|
||||
|
||||
completed, err := svc.Get(context.Background(), owner, created.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ImageTaskStatusCompleted, completed.Status)
|
||||
require.Equal(t, http.StatusOK, completed.HTTPStatus)
|
||||
require.Equal(t, "https://example.test/image.png", completed.ImageURL)
|
||||
require.JSONEq(t, string(result), string(completed.Result))
|
||||
require.NotNil(t, completed.CompletedAt)
|
||||
}
|
||||
|
||||
func TestImageTaskServiceInvalidResultBecomesFailed(t *testing.T) {
|
||||
store := &imageTaskMemoryStore{}
|
||||
svc := NewImageTaskServiceWithOptions(store, time.Hour, time.Minute)
|
||||
created, err := svc.Create(context.Background(), ImageTaskOwner{UserID: 1, APIKeyID: 2})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, svc.Complete(context.Background(), created.ID, http.StatusOK, json.RawMessage(`not-json`)))
|
||||
got, err := svc.Get(context.Background(), ImageTaskOwner{UserID: 1, APIKeyID: 2}, created.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ImageTaskStatusFailed, got.Status)
|
||||
require.Equal(t, http.StatusBadGateway, got.HTTPStatus)
|
||||
require.Contains(t, string(got.Error), "non-JSON")
|
||||
}
|
||||
|
||||
func TestImageTaskServiceMapsStoreFailures(t *testing.T) {
|
||||
store := &imageTaskMemoryStore{saveErr: errors.New("redis down")}
|
||||
svc := NewImageTaskService(store)
|
||||
|
||||
_, err := svc.Create(context.Background(), ImageTaskOwner{UserID: 1, APIKeyID: 2})
|
||||
require.ErrorIs(t, err, ErrImageTaskUnavailable)
|
||||
}
|
||||
@@ -521,6 +521,22 @@ func ProvideAPIKeyAuthCacheInvalidator(apiKeyService *APIKeyService) APIKeyAuthC
|
||||
return apiKeyService
|
||||
}
|
||||
|
||||
// ProvideImageTaskService 构造异步图片任务服务。
|
||||
//
|
||||
// 对象存储是异步图片任务的启用前提:仅当 image_storage 开关打开且凭证齐全时,
|
||||
// 服务才启用,并挂上把结果转存到对象存储的 uploader;否则功能整体禁用
|
||||
// (handler 返回 404,不创建任务、不写 Redis),从而避免大 base64 结果撑爆 Redis。
|
||||
func ProvideImageTaskService(store ImageTaskStore, storage ImageStorage, cfg *config.Config) *ImageTaskService {
|
||||
if !cfg.ImageStorage.Active() {
|
||||
if cfg.ImageStorage.Enabled {
|
||||
logger.L().Warn("image_storage.enabled is true but object storage is not fully configured; async image tasks are disabled")
|
||||
}
|
||||
return NewImageTaskService(store)
|
||||
}
|
||||
uploader := NewImageResultUploader(storage, cfg.ImageStorage.Prefix, cfg.ImageStorage.MaxDownloadByte, nil)
|
||||
return NewImageTaskServiceWithUploader(store, uploader, defaultImageTaskTTL, defaultImageTaskExecutionTimeout)
|
||||
}
|
||||
|
||||
// ProvideBackupService creates and starts BackupService
|
||||
func ProvideBackupService(
|
||||
settingRepo SettingRepository,
|
||||
@@ -645,6 +661,7 @@ var ProviderSet = wire.NewSet(
|
||||
NewAdminService,
|
||||
NewGatewayService,
|
||||
NewOpenAIGatewayService,
|
||||
ProvideImageTaskService,
|
||||
ProvideBatchImageModelPricingResolver,
|
||||
NewBatchImagePublicService,
|
||||
NewBatchImageDownloadService,
|
||||
@@ -654,6 +671,7 @@ var ProviderSet = wire.NewSet(
|
||||
NewOAuthService,
|
||||
ProvideOpenAIOAuthService,
|
||||
NewGrokOAuthService,
|
||||
wire.Bind(new(GrokOAuthTokenService), new(*GrokOAuthService)),
|
||||
NewGeminiOAuthService,
|
||||
NewGeminiQuotaService,
|
||||
NewCompositeTokenCacheInvalidator,
|
||||
|
||||
@@ -1175,3 +1175,39 @@ update:
|
||||
# Leave empty for direct connection (recommended for overseas servers)
|
||||
# 留空表示直连(适用于海外服务器)
|
||||
proxy_url: ""
|
||||
|
||||
# =============================================================================
|
||||
# Image Storage (异步图片任务结果对象存储)
|
||||
# =============================================================================
|
||||
# 长耗时生图套 Cloudflare 会 524 超时,异步图片任务接口(/v1/images/generations/async
|
||||
# 等)先返回 task_id,再由客户端轮询 /v1/images/tasks/{task_id} 获取结果。
|
||||
#
|
||||
# 这里配置一个 S3 兼容对象存储(AWS S3 / Cloudflare R2 / 阿里云 OSS / MinIO 等),
|
||||
# 任务完成后把生成的图片上传到对象存储,只在 Redis 存一个短链接,避免 gpt-image-1 等
|
||||
# 返回的大 base64 结果把 Redis 撑爆。
|
||||
#
|
||||
# enabled 同时是异步图片任务功能的总开关:为 false 或凭证未配全时,异步生图接口
|
||||
# 整体返回 404、不创建任务、不写 Redis。
|
||||
#
|
||||
# 换其它厂商对象存储:只要它兼容 S3 API 即可直接用;如需完全自定义,实现
|
||||
# service.ImageStorage 接口(Save(ctx, key, contentType, data) -> url)即可。
|
||||
image_storage:
|
||||
enabled: false
|
||||
# S3 兼容端点。AWS 官方可留空;R2 形如 https://<account_id>.r2.cloudflarestorage.com
|
||||
endpoint: ""
|
||||
# 区域。Cloudflare R2 用 "auto"
|
||||
region: "auto"
|
||||
bucket: ""
|
||||
access_key_id: ""
|
||||
secret_access_key: ""
|
||||
# 对象 key 前缀
|
||||
prefix: "images/"
|
||||
# MinIO / 需要路径风格(path-style)访问的桶设为 true
|
||||
force_path_style: false
|
||||
# 若填写公开桶 / CDN 域名,则返回 public_base_url/key 永久直链;
|
||||
# 留空则返回带过期时间的 presigned 临时链接
|
||||
public_base_url: ""
|
||||
# public_base_url 为空时,presigned 链接的有效时长(小时)
|
||||
presign_expiry_hours: 24
|
||||
# 当上游返回的是图片 url 时,下载该图片再转存的字节上限(默认 32MB)
|
||||
max_download_bytes: 33554432
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
# Asynchronous Image Tasks
|
||||
|
||||
Asynchronous image tasks let clients submit long-running OpenAI-compatible image requests without keeping one HTTP connection open. This avoids proxy/CDN response timeouts such as Cloudflare 524 while preserving the existing image routing, billing, moderation, concurrency, and failover behavior.
|
||||
|
||||
## Endpoints
|
||||
|
||||
The authenticated gateway exposes both `/v1` paths and their existing no-prefix aliases:
|
||||
|
||||
```text
|
||||
POST /v1/images/generations/async
|
||||
POST /v1/images/edits/async
|
||||
GET /v1/images/tasks/{task_id}
|
||||
```
|
||||
|
||||
The aliases are `/images/generations/async`, `/images/edits/async`, and `/images/tasks/{task_id}`.
|
||||
|
||||
Only OpenAI and Grok groups are supported. Requests use the same JSON or multipart payload as the corresponding synchronous endpoint. Streaming image requests are rejected because a polled task returns one final JSON result.
|
||||
|
||||
## Enabling the feature (object storage)
|
||||
|
||||
Asynchronous image tasks are **disabled by default** and gated on object storage. When the switch is off — or the S3 credentials are incomplete — the async endpoints return `404` and never create a task or write to Redis. This is deliberate: without offloading, large `b64_json` results (several MB each, e.g. `gpt-image-1`) would accumulate in Redis and exhaust its memory.
|
||||
|
||||
Configure an S3-compatible object store (AWS S3, Cloudflare R2, Aliyun OSS, MinIO, …) in `config.yaml` (all keys also accept the `IMAGE_STORAGE_*` environment overrides):
|
||||
|
||||
```yaml
|
||||
image_storage:
|
||||
enabled: true
|
||||
endpoint: "https://<account_id>.r2.cloudflarestorage.com" # AWS 官方可留空
|
||||
region: "auto"
|
||||
bucket: "my-images"
|
||||
access_key_id: "..."
|
||||
secret_access_key: "..."
|
||||
prefix: "images/"
|
||||
force_path_style: false # MinIO/path-style buckets set true
|
||||
public_base_url: "" # set to return public_base_url/key直链; empty → presigned URL
|
||||
presign_expiry_hours: 24 # presigned link TTL when public_base_url is empty
|
||||
max_download_bytes: 33554432 # cap when re-hosting an upstream image URL (32MB)
|
||||
```
|
||||
|
||||
When a task completes, each generated image is uploaded to the bucket and the result is rewritten to a compact form: `data[].url` points at the stored object (a permanent `public_base_url/key` link, or a time-limited presigned URL) and `b64_json` is removed. Only this small JSON is stored in Redis. If an upload fails, the task is marked `failed` rather than persisting the raw base64.
|
||||
|
||||
To support a different vendor beyond the S3-compatible client, implement the `service.ImageStorage` interface (`Save(ctx, key, contentType, data) (url, error)`) and provide it in place of the S3 implementation.
|
||||
|
||||
## Submit a task
|
||||
|
||||
```bash
|
||||
curl -i https://api.example.com/v1/images/generations/async \
|
||||
-H 'Authorization: Bearer sk-...' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"model": "gpt-image-1",
|
||||
"prompt": "A lighthouse during a winter storm",
|
||||
"size": "1536x1024"
|
||||
}'
|
||||
```
|
||||
|
||||
The server stores the initial task in Redis and responds with `202 Accepted`:
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "imgtask_0123456789abcdef",
|
||||
"task_id": "imgtask_0123456789abcdef",
|
||||
"object": "image.generation.task",
|
||||
"status": "processing",
|
||||
"created_at": 1784092800,
|
||||
"expires_at": 1784179200,
|
||||
"poll_url": "/v1/images/tasks/imgtask_0123456789abcdef"
|
||||
}
|
||||
```
|
||||
|
||||
`Location` contains the polling path and `Retry-After: 3` provides the recommended polling interval.
|
||||
|
||||
## Poll a task
|
||||
|
||||
Use the same API key that submitted the task:
|
||||
|
||||
```bash
|
||||
curl https://api.example.com/v1/images/tasks/imgtask_0123456789abcdef \
|
||||
-H 'Authorization: Bearer sk-...'
|
||||
```
|
||||
|
||||
While work is in progress:
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "imgtask_0123456789abcdef",
|
||||
"task_id": "imgtask_0123456789abcdef",
|
||||
"object": "image.generation.task",
|
||||
"status": "processing",
|
||||
"created_at": 1784092800,
|
||||
"expires_at": 1784179200
|
||||
}
|
||||
```
|
||||
|
||||
On success, `result` mirrors the synchronous image API body, except each image has been offloaded to object storage: `data[].url` points at the stored object and `b64_json` is stripped (so both URL and base64 upstream formats end up as compact stored links):
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "imgtask_0123456789abcdef",
|
||||
"task_id": "imgtask_0123456789abcdef",
|
||||
"object": "image.generation.task",
|
||||
"status": "completed",
|
||||
"http_status": 200,
|
||||
"image_url": "https://...",
|
||||
"result": {
|
||||
"created": 1784092923,
|
||||
"data": [{"url": "https://..."}]
|
||||
},
|
||||
"created_at": 1784092800,
|
||||
"completed_at": 1784092923,
|
||||
"expires_at": 1784179323
|
||||
}
|
||||
```
|
||||
|
||||
For URL responses, `image_url` mirrors the first `data[].url` for simple clients. On failure, the task reaches `failed` and exposes the original OpenAI-compatible error object where available:
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "imgtask_0123456789abcdef",
|
||||
"task_id": "imgtask_0123456789abcdef",
|
||||
"object": "image.generation.task",
|
||||
"status": "failed",
|
||||
"http_status": 502,
|
||||
"error": {
|
||||
"type": "api_error",
|
||||
"message": "Upstream request failed"
|
||||
},
|
||||
"created_at": 1784092800,
|
||||
"completed_at": 1784092923,
|
||||
"expires_at": 1784179323
|
||||
}
|
||||
```
|
||||
|
||||
All submit and poll responses include `Cache-Control: no-store`, preventing a CDN from caching the `processing` state. Tasks and results expire 24 hours after their latest state update. A task executes for at most 30 minutes.
|
||||
|
||||
Task ownership is scoped to both user and API key. Unknown task IDs and IDs owned by another key both return `404`, avoiding task-existence disclosure. Polling remains available when the completed generation used the key's remaining balance; normal authentication, disabled-key, user, IP, and group checks still apply.
|
||||
Reference in New Issue
Block a user