diff --git a/.github/workflows/backend-ci.yml b/.github/workflows/backend-ci.yml
index 15ff97fe..fb4d0ce6 100644
--- a/.github/workflows/backend-ci.yml
+++ b/.github/workflows/backend-ci.yml
@@ -20,7 +20,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
- name: Unit tests
working-directory: backend
run: make test-unit
@@ -60,7 +60,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
- name: golangci-lint
uses: golangci/golangci-lint-action@v9
with:
diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml
index 80bc9850..7d48131a 100644
--- a/.github/workflows/release.yml
+++ b/.github/workflows/release.yml
@@ -115,7 +115,7 @@ jobs:
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
# Docker setup for GoReleaser
- name: Set up QEMU
diff --git a/.github/workflows/security-scan.yml b/.github/workflows/security-scan.yml
index ef8e59e5..e102b5f8 100644
--- a/.github/workflows/security-scan.yml
+++ b/.github/workflows/security-scan.yml
@@ -23,7 +23,7 @@ jobs:
cache-dependency-path: backend/go.sum
- name: Verify Go version
run: |
- go version | grep -q 'go1.26.3'
+ go version | grep -q 'go1.26.4'
- name: Run govulncheck
working-directory: backend
run: |
diff --git a/Dockerfile b/Dockerfile
index d556008b..f9a03a2b 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -7,7 +7,7 @@
# =============================================================================
ARG NODE_IMAGE=node:24-alpine
-ARG GOLANG_IMAGE=golang:1.26.3-alpine
+ARG GOLANG_IMAGE=golang:1.26.4-alpine
ARG ALPINE_IMAGE=alpine:3.21
ARG POSTGRES_IMAGE=postgres:18-alpine
ARG GOPROXY=https://goproxy.cn,direct
diff --git a/backend/Dockerfile b/backend/Dockerfile
index f153d686..26b1dc33 100644
--- a/backend/Dockerfile
+++ b/backend/Dockerfile
@@ -1,4 +1,4 @@
-FROM golang:1.26.3-alpine
+FROM golang:1.26.4-alpine
WORKDIR /app
diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION
index 66c01044..56ebc9e5 100644
--- a/backend/cmd/server/VERSION
+++ b/backend/cmd/server/VERSION
@@ -1 +1 @@
-0.1.131
+0.1.133
diff --git a/backend/cmd/server/wire.go b/backend/cmd/server/wire.go
index 9bfa2717..b474cfa1 100644
--- a/backend/cmd/server/wire.go
+++ b/backend/cmd/server/wire.go
@@ -98,6 +98,7 @@ func provideCleanup(
backupSvc *service.BackupService,
paymentOrderExpiry *service.PaymentOrderExpiryService,
channelMonitorRunner *service.ChannelMonitorRunner,
+ quotaFlusher *service.UserPlatformQuotaUsageFlusher,
) func() {
return func() {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
@@ -246,6 +247,12 @@ func provideCleanup(
}
return nil
}},
+ {"UserPlatformQuotaUsageFlusher", func() error {
+ if quotaFlusher != nil {
+ quotaFlusher.Stop()
+ }
+ return nil
+ }},
}
infraSteps := []cleanupStep{
diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go
index 465f5e25..814c07e6 100644
--- a/backend/cmd/server/wire_gen.go
+++ b/backend/cmd/server/wire_gen.go
@@ -91,36 +91,15 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
apiKeyHandler := handler.NewAPIKeyHandler(apiKeyService)
usageLogRepository := repository.NewUsageLogRepository(client, db)
usageService := service.NewUsageService(usageLogRepository, userRepository, client, apiKeyAuthCacheInvalidator)
- usageHandler := handler.NewUsageHandler(usageService, apiKeyService)
- redeemHandler := handler.NewRedeemHandler(redeemService)
- subscriptionHandler := handler.NewSubscriptionHandler(subscriptionService)
- announcementRepository := repository.NewAnnouncementRepository(client)
- announcementReadRepository := repository.NewAnnouncementReadRepository(client)
- announcementService := service.NewAnnouncementService(announcementRepository, announcementReadRepository, userRepository, userSubscriptionRepository)
- announcementHandler := handler.NewAnnouncementHandler(announcementService)
- channelMonitorRepository := repository.NewChannelMonitorRepository(client, db)
- channelMonitorService := service.ProvideChannelMonitorService(channelMonitorRepository, secretEncryptor)
- channelMonitorUserHandler := handler.NewChannelMonitorUserHandler(channelMonitorService, settingService)
- dashboardAggregationRepository := repository.NewDashboardAggregationRepository(db)
- dashboardStatsCache := repository.NewDashboardCache(redisClient, configConfig)
- dashboardService := service.NewDashboardService(usageLogRepository, dashboardAggregationRepository, dashboardStatsCache, configConfig)
- timingWheelService, err := service.ProvideTimingWheelService()
- if err != nil {
- return nil, err
- }
- dashboardAggregationService := service.ProvideDashboardAggregationService(dashboardAggregationRepository, timingWheelService, configConfig)
- dashboardHandler := admin.NewDashboardHandler(dashboardService, dashboardAggregationService)
+ opsRepository := repository.NewOpsRepository(db)
schedulerCache := repository.ProvideSchedulerCache(redisClient, configConfig)
accountRepository := repository.NewAccountRepository(client, db, schedulerCache)
- proxyExitInfoProber := repository.NewProxyExitInfoProber(configConfig)
- proxyLatencyCache := repository.NewProxyLatencyCache(redisClient)
- privacyClientFactory := providePrivacyClientFactory()
+ concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig)
+ concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig)
usageBillingRepository := repository.NewUsageBillingRepository(client, db)
gatewayCache := repository.NewGatewayCache(redisClient)
schedulerOutboxRepository := repository.NewSchedulerOutboxRepository(db)
schedulerSnapshotService := service.ProvideSchedulerSnapshotService(schedulerCache, schedulerOutboxRepository, accountRepository, groupRepository, configConfig)
- concurrencyCache := repository.ProvideConcurrencyCache(redisClient, configConfig)
- concurrencyService := service.ProvideConcurrencyService(concurrencyCache, accountRepository, configConfig)
pricingRemoteClient := repository.ProvidePricingRemoteClient(configConfig)
pricingService, err := service.ProvidePricingService(configConfig, pricingRemoteClient)
if err != nil {
@@ -134,44 +113,72 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
geminiTokenCache := repository.NewGeminiTokenCache(redisClient)
compositeTokenCacheInvalidator := service.NewCompositeTokenCacheInvalidator(geminiTokenCache)
rateLimitService := service.ProvideRateLimitService(accountRepository, usageLogRepository, configConfig, geminiQuotaService, tempUnschedCache, timeoutCounterCache, openAI403CounterCache, settingService, compositeTokenCacheInvalidator)
+ identityCache := repository.NewIdentityCache(redisClient)
+ identityService := service.NewIdentityService(identityCache)
httpUpstream := repository.NewHTTPUpstream(configConfig)
+ timingWheelService, err := service.ProvideTimingWheelService()
+ if err != nil {
+ return nil, err
+ }
deferredService := service.ProvideDeferredService(accountRepository, timingWheelService)
- openAIOAuthClient := repository.NewOpenAIOAuthClient()
- openAIOAuthService := service.NewOpenAIOAuthService(proxyRepository, openAIOAuthClient)
+ claudeOAuthClient := repository.NewClaudeOAuthClient()
+ oAuthService := service.NewOAuthService(proxyRepository, claudeOAuthClient)
oAuthRefreshAPI := service.ProvideOAuthRefreshAPI(accountRepository, geminiTokenCache)
- openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI)
+ claudeTokenProvider := service.ProvideClaudeTokenProvider(accountRepository, geminiTokenCache, oAuthService, oAuthRefreshAPI)
+ sessionLimitCache := repository.ProvideSessionLimitCache(redisClient, configConfig)
+ rpmCache := repository.NewRPMCache(redisClient)
+ digestSessionStore := service.NewDigestSessionStore()
+ tlsFingerprintProfileRepository := repository.NewTLSFingerprintProfileRepository(client)
+ tlsFingerprintProfileCache := repository.NewTLSFingerprintProfileCache(redisClient)
+ tlsFingerprintProfileService := service.NewTLSFingerprintProfileService(tlsFingerprintProfileRepository, tlsFingerprintProfileCache)
channelRepository := repository.NewChannelRepository(db)
channelService := service.NewChannelService(channelRepository, groupRepository, apiKeyAuthCacheInvalidator, pricingService)
modelPricingResolver := service.NewModelPricingResolver(channelService, billingService)
notificationEmailService := service.NewNotificationEmailService(settingRepository, emailService)
balanceNotifyService := service.ProvideBalanceNotifyService(emailService, settingRepository, accountRepository, notificationEmailService)
+ gatewayService := service.NewGatewayService(accountRepository, groupRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, identityService, httpUpstream, deferredService, claudeTokenProvider, sessionLimitCache, rpmCache, digestSessionStore, settingService, tlsFingerprintProfileService, channelService, modelPricingResolver, balanceNotifyService, serviceUserPlatformQuotaRepository)
+ openAIOAuthClient := repository.NewOpenAIOAuthClient()
+ privacyClientFactory := providePrivacyClientFactory()
+ openAIOAuthService := service.ProvideOpenAIOAuthService(proxyRepository, openAIOAuthClient, privacyClientFactory)
+ openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI)
openAIGatewayService := service.NewOpenAIGatewayService(accountRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, httpUpstream, deferredService, openAITokenProvider, modelPricingResolver, channelService, balanceNotifyService, settingService, serviceUserPlatformQuotaRepository)
- adminService := service.NewAdminService(userRepository, groupRepository, accountRepository, proxyRepository, apiKeyRepository, redeemCodeRepository, userGroupRateRepository, userRPMCache, billingCacheService, proxyExitInfoProber, proxyLatencyCache, apiKeyAuthCacheInvalidator, client, settingService, subscriptionService, userSubscriptionRepository, privacyClientFactory, openAIGatewayService)
- adminUserHandler := admin.NewUserHandler(adminService, concurrencyService, serviceUserPlatformQuotaRepository, billingCache)
- sessionLimitCache := repository.ProvideSessionLimitCache(redisClient, configConfig)
- rpmCache := repository.NewRPMCache(redisClient)
- groupCapacityService := service.NewGroupCapacityService(accountRepository, groupRepository, concurrencyService, sessionLimitCache, rpmCache)
- groupHandler := admin.NewGroupHandler(adminService, dashboardService, groupCapacityService)
- claudeOAuthClient := repository.NewClaudeOAuthClient()
- oAuthService := service.NewOAuthService(proxyRepository, claudeOAuthClient)
geminiOAuthClient := repository.NewGeminiOAuthClient(configConfig)
geminiCliCodeAssistClient := repository.NewGeminiCliCodeAssistClient()
driveClient := repository.NewGeminiDriveClient()
geminiOAuthService := service.NewGeminiOAuthService(proxyRepository, geminiOAuthClient, geminiCliCodeAssistClient, driveClient, configConfig)
- antigravityOAuthService := service.NewAntigravityOAuthService(proxyRepository)
- claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream)
- antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository)
- usageCache := service.NewUsageCache()
- identityCache := repository.NewIdentityCache(redisClient)
- tlsFingerprintProfileRepository := repository.NewTLSFingerprintProfileRepository(client)
- tlsFingerprintProfileCache := repository.NewTLSFingerprintProfileCache(redisClient)
- tlsFingerprintProfileService := service.NewTLSFingerprintProfileService(tlsFingerprintProfileRepository, tlsFingerprintProfileCache)
- accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, usageCache, identityCache, tlsFingerprintProfileService)
geminiTokenProvider := service.ProvideGeminiTokenProvider(accountRepository, geminiTokenCache, geminiOAuthService, oAuthRefreshAPI)
- claudeTokenProvider := service.ProvideClaudeTokenProvider(accountRepository, geminiTokenCache, oAuthService, oAuthRefreshAPI)
+ antigravityOAuthService := service.NewAntigravityOAuthService(proxyRepository)
antigravityTokenProvider := service.ProvideAntigravityTokenProvider(accountRepository, geminiTokenCache, antigravityOAuthService, oAuthRefreshAPI, tempUnschedCache)
internal500CounterCache := repository.NewInternal500CounterCache(redisClient)
antigravityGatewayService := service.NewAntigravityGatewayService(accountRepository, gatewayCache, schedulerSnapshotService, antigravityTokenProvider, rateLimitService, httpUpstream, settingService, internal500CounterCache)
+ geminiMessagesCompatService := service.NewGeminiMessagesCompatService(accountRepository, groupRepository, gatewayCache, schedulerSnapshotService, geminiTokenProvider, rateLimitService, httpUpstream, antigravityGatewayService, configConfig)
+ opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository)
+ opsService := service.ProvideOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink, settingService)
+ usageHandler := handler.NewUsageHandler(usageService, apiKeyService, opsService, settingService)
+ redeemHandler := handler.NewRedeemHandler(redeemService)
+ subscriptionHandler := handler.NewSubscriptionHandler(subscriptionService)
+ announcementRepository := repository.NewAnnouncementRepository(client)
+ announcementReadRepository := repository.NewAnnouncementReadRepository(client)
+ announcementService := service.NewAnnouncementService(announcementRepository, announcementReadRepository, userRepository, userSubscriptionRepository)
+ announcementHandler := handler.NewAnnouncementHandler(announcementService)
+ channelMonitorRepository := repository.NewChannelMonitorRepository(client, db)
+ channelMonitorService := service.ProvideChannelMonitorService(channelMonitorRepository, secretEncryptor)
+ channelMonitorUserHandler := handler.NewChannelMonitorUserHandler(channelMonitorService, settingService)
+ dashboardAggregationRepository := repository.NewDashboardAggregationRepository(db)
+ dashboardStatsCache := repository.NewDashboardCache(redisClient, configConfig)
+ dashboardService := service.NewDashboardService(usageLogRepository, dashboardAggregationRepository, dashboardStatsCache, configConfig)
+ dashboardAggregationService := service.ProvideDashboardAggregationService(dashboardAggregationRepository, timingWheelService, configConfig)
+ dashboardHandler := admin.NewDashboardHandler(dashboardService, dashboardAggregationService)
+ proxyExitInfoProber := repository.NewProxyExitInfoProber(configConfig)
+ proxyLatencyCache := repository.NewProxyLatencyCache(redisClient)
+ adminService := service.NewAdminService(userRepository, groupRepository, accountRepository, proxyRepository, apiKeyRepository, redeemCodeRepository, userGroupRateRepository, userRPMCache, billingCacheService, proxyExitInfoProber, proxyLatencyCache, apiKeyAuthCacheInvalidator, client, settingService, subscriptionService, userSubscriptionRepository, privacyClientFactory, openAIGatewayService)
+ adminUserHandler := admin.NewUserHandler(adminService, concurrencyService, serviceUserPlatformQuotaRepository, billingCache)
+ groupCapacityService := service.NewGroupCapacityService(accountRepository, groupRepository, concurrencyService, sessionLimitCache, rpmCache)
+ groupHandler := admin.NewGroupHandler(adminService, dashboardService, groupCapacityService)
+ claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream)
+ antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository)
+ usageCache := service.NewUsageCache()
+ accountUsageService := service.NewAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, usageCache, identityCache, tlsFingerprintProfileService)
accountTestService := service.NewAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService)
crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig)
accountHandler := admin.NewAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator)
@@ -189,13 +196,6 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
proxyHandler := admin.NewProxyHandler(adminService)
adminRedeemHandler := admin.NewRedeemHandler(adminService, redeemService)
promoHandler := admin.NewPromoHandler(promoService)
- opsRepository := repository.NewOpsRepository(db)
- identityService := service.NewIdentityService(identityCache)
- digestSessionStore := service.NewDigestSessionStore()
- gatewayService := service.NewGatewayService(accountRepository, groupRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, identityService, httpUpstream, deferredService, claudeTokenProvider, sessionLimitCache, rpmCache, digestSessionStore, settingService, tlsFingerprintProfileService, channelService, modelPricingResolver, balanceNotifyService, serviceUserPlatformQuotaRepository)
- geminiMessagesCompatService := service.NewGeminiMessagesCompatService(accountRepository, groupRepository, gatewayCache, schedulerSnapshotService, geminiTokenProvider, rateLimitService, httpUpstream, antigravityGatewayService, configConfig)
- opsSystemLogSink := service.ProvideOpsSystemLogSink(opsRepository)
- opsService := service.NewOpsService(opsRepository, settingRepository, configConfig, accountRepository, userRepository, concurrencyService, gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, opsSystemLogSink)
encryptionKey, err := payment.ProvideEncryptionKey(configConfig)
if err != nil {
return nil, err
@@ -269,7 +269,8 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
scheduledTestRunnerService := service.ProvideScheduledTestRunnerService(scheduledTestPlanRepository, scheduledTestService, accountTestService, rateLimitService, configConfig)
paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService)
channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService)
- v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner)
+ userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService)
+ v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher)
application := &Application{
Server: httpServer,
Cleanup: v,
@@ -324,6 +325,7 @@ func provideCleanup(
backupSvc *service.BackupService,
paymentOrderExpiry *service.PaymentOrderExpiryService,
channelMonitorRunner *service.ChannelMonitorRunner,
+ quotaFlusher *service.UserPlatformQuotaUsageFlusher,
) func() {
return func() {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
@@ -471,6 +473,12 @@ func provideCleanup(
}
return nil
}},
+ {"UserPlatformQuotaUsageFlusher", func() error {
+ if quotaFlusher != nil {
+ quotaFlusher.Stop()
+ }
+ return nil
+ }},
}
infraSteps := []cleanupStep{
diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go
index a44b2d5c..7f4e4773 100644
--- a/backend/cmd/server/wire_gen_test.go
+++ b/backend/cmd/server/wire_gen_test.go
@@ -77,6 +77,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
nil, // backupSvc
nil, // paymentOrderExpiry
nil, // channelMonitorRunner
+ nil, // quotaFlusher
)
require.NotPanics(t, func() {
diff --git a/backend/ent/group.go b/backend/ent/group.go
index a4f52c73..298df88a 100644
--- a/backend/ent/group.go
+++ b/backend/ent/group.go
@@ -85,6 +85,8 @@ type Group struct {
DefaultMappedModel string `json:"default_mapped_model,omitempty"`
// OpenAI Messages 调度模型配置:按 Claude 系列/精确模型映射到目标 GPT 模型
MessagesDispatchModelConfig domain.OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config,omitempty"`
+ // 自定义 /v1/models 展示列表配置;仅影响模型列表响应,不影响调度
+ ModelsListConfig domain.GroupModelsListConfig `json:"models_list_config,omitempty"`
// 分组 RPM 上限,0 表示不限制;设置后接管该分组用户的限流
RpmLimit int `json:"rpm_limit,omitempty"`
// Edges holds the relations/edges for other nodes in the graph.
@@ -193,7 +195,7 @@ func (*Group) scanValues(columns []string) ([]any, error) {
values := make([]any, len(columns))
for i := range columns {
switch columns[i] {
- case group.FieldModelRouting, group.FieldSupportedModelScopes, group.FieldMessagesDispatchModelConfig:
+ case group.FieldModelRouting, group.FieldSupportedModelScopes, group.FieldMessagesDispatchModelConfig, group.FieldModelsListConfig:
values[i] = new([]byte)
case group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldImageRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet:
values[i] = new(sql.NullBool)
@@ -440,6 +442,14 @@ func (_m *Group) assignValues(columns []string, values []any) error {
return fmt.Errorf("unmarshal field messages_dispatch_model_config: %w", err)
}
}
+ case group.FieldModelsListConfig:
+ if value, ok := values[i].(*[]byte); !ok {
+ return fmt.Errorf("unexpected type %T for field models_list_config", values[i])
+ } else if value != nil && len(*value) > 0 {
+ if err := json.Unmarshal(*value, &_m.ModelsListConfig); err != nil {
+ return fmt.Errorf("unmarshal field models_list_config: %w", err)
+ }
+ }
case group.FieldRpmLimit:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field rpm_limit", values[i])
@@ -641,6 +651,9 @@ func (_m *Group) String() string {
builder.WriteString("messages_dispatch_model_config=")
builder.WriteString(fmt.Sprintf("%v", _m.MessagesDispatchModelConfig))
builder.WriteString(", ")
+ builder.WriteString("models_list_config=")
+ builder.WriteString(fmt.Sprintf("%v", _m.ModelsListConfig))
+ builder.WriteString(", ")
builder.WriteString("rpm_limit=")
builder.WriteString(fmt.Sprintf("%v", _m.RpmLimit))
builder.WriteByte(')')
diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go
index 4e9ba6b6..ebe9bd7e 100644
--- a/backend/ent/group/group.go
+++ b/backend/ent/group/group.go
@@ -82,6 +82,8 @@ const (
FieldDefaultMappedModel = "default_mapped_model"
// FieldMessagesDispatchModelConfig holds the string denoting the messages_dispatch_model_config field in the database.
FieldMessagesDispatchModelConfig = "messages_dispatch_model_config"
+ // FieldModelsListConfig holds the string denoting the models_list_config field in the database.
+ FieldModelsListConfig = "models_list_config"
// FieldRpmLimit holds the string denoting the rpm_limit field in the database.
FieldRpmLimit = "rpm_limit"
// EdgeAPIKeys holds the string denoting the api_keys edge name in mutations.
@@ -192,6 +194,7 @@ var Columns = []string{
FieldRequirePrivacySet,
FieldDefaultMappedModel,
FieldMessagesDispatchModelConfig,
+ FieldModelsListConfig,
FieldRpmLimit,
}
@@ -276,6 +279,8 @@ var (
DefaultMappedModelValidator func(string) error
// DefaultMessagesDispatchModelConfig holds the default value on creation for the "messages_dispatch_model_config" field.
DefaultMessagesDispatchModelConfig domain.OpenAIMessagesDispatchModelConfig
+ // DefaultModelsListConfig holds the default value on creation for the "models_list_config" field.
+ DefaultModelsListConfig domain.GroupModelsListConfig
// DefaultRpmLimit holds the default value on creation for the "rpm_limit" field.
DefaultRpmLimit int
)
diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go
index 44b905bd..d5ed0c19 100644
--- a/backend/ent/group_create.go
+++ b/backend/ent/group_create.go
@@ -467,6 +467,20 @@ func (_c *GroupCreate) SetNillableMessagesDispatchModelConfig(v *domain.OpenAIMe
return _c
}
+// SetModelsListConfig sets the "models_list_config" field.
+func (_c *GroupCreate) SetModelsListConfig(v domain.GroupModelsListConfig) *GroupCreate {
+ _c.mutation.SetModelsListConfig(v)
+ return _c
+}
+
+// SetNillableModelsListConfig sets the "models_list_config" field if the given value is not nil.
+func (_c *GroupCreate) SetNillableModelsListConfig(v *domain.GroupModelsListConfig) *GroupCreate {
+ if v != nil {
+ _c.SetModelsListConfig(*v)
+ }
+ return _c
+}
+
// SetRpmLimit sets the "rpm_limit" field.
func (_c *GroupCreate) SetRpmLimit(v int) *GroupCreate {
_c.mutation.SetRpmLimit(v)
@@ -698,6 +712,10 @@ func (_c *GroupCreate) defaults() error {
v := group.DefaultMessagesDispatchModelConfig
_c.mutation.SetMessagesDispatchModelConfig(v)
}
+ if _, ok := _c.mutation.ModelsListConfig(); !ok {
+ v := group.DefaultModelsListConfig
+ _c.mutation.SetModelsListConfig(v)
+ }
if _, ok := _c.mutation.RpmLimit(); !ok {
v := group.DefaultRpmLimit
_c.mutation.SetRpmLimit(v)
@@ -798,6 +816,9 @@ func (_c *GroupCreate) check() error {
if _, ok := _c.mutation.MessagesDispatchModelConfig(); !ok {
return &ValidationError{Name: "messages_dispatch_model_config", err: errors.New(`ent: missing required field "Group.messages_dispatch_model_config"`)}
}
+ if _, ok := _c.mutation.ModelsListConfig(); !ok {
+ return &ValidationError{Name: "models_list_config", err: errors.New(`ent: missing required field "Group.models_list_config"`)}
+ }
if _, ok := _c.mutation.RpmLimit(); !ok {
return &ValidationError{Name: "rpm_limit", err: errors.New(`ent: missing required field "Group.rpm_limit"`)}
}
@@ -960,6 +981,10 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) {
_spec.SetField(group.FieldMessagesDispatchModelConfig, field.TypeJSON, value)
_node.MessagesDispatchModelConfig = value
}
+ if value, ok := _c.mutation.ModelsListConfig(); ok {
+ _spec.SetField(group.FieldModelsListConfig, field.TypeJSON, value)
+ _node.ModelsListConfig = value
+ }
if value, ok := _c.mutation.RpmLimit(); ok {
_spec.SetField(group.FieldRpmLimit, field.TypeInt, value)
_node.RpmLimit = value
@@ -1642,6 +1667,18 @@ func (u *GroupUpsert) UpdateMessagesDispatchModelConfig() *GroupUpsert {
return u
}
+// SetModelsListConfig sets the "models_list_config" field.
+func (u *GroupUpsert) SetModelsListConfig(v domain.GroupModelsListConfig) *GroupUpsert {
+ u.Set(group.FieldModelsListConfig, v)
+ return u
+}
+
+// UpdateModelsListConfig sets the "models_list_config" field to the value that was provided on create.
+func (u *GroupUpsert) UpdateModelsListConfig() *GroupUpsert {
+ u.SetExcluded(group.FieldModelsListConfig)
+ return u
+}
+
// SetRpmLimit sets the "rpm_limit" field.
func (u *GroupUpsert) SetRpmLimit(v int) *GroupUpsert {
u.Set(group.FieldRpmLimit, v)
@@ -2314,6 +2351,20 @@ func (u *GroupUpsertOne) UpdateMessagesDispatchModelConfig() *GroupUpsertOne {
})
}
+// SetModelsListConfig sets the "models_list_config" field.
+func (u *GroupUpsertOne) SetModelsListConfig(v domain.GroupModelsListConfig) *GroupUpsertOne {
+ return u.Update(func(s *GroupUpsert) {
+ s.SetModelsListConfig(v)
+ })
+}
+
+// UpdateModelsListConfig sets the "models_list_config" field to the value that was provided on create.
+func (u *GroupUpsertOne) UpdateModelsListConfig() *GroupUpsertOne {
+ return u.Update(func(s *GroupUpsert) {
+ s.UpdateModelsListConfig()
+ })
+}
+
// SetRpmLimit sets the "rpm_limit" field.
func (u *GroupUpsertOne) SetRpmLimit(v int) *GroupUpsertOne {
return u.Update(func(s *GroupUpsert) {
@@ -3155,6 +3206,20 @@ func (u *GroupUpsertBulk) UpdateMessagesDispatchModelConfig() *GroupUpsertBulk {
})
}
+// SetModelsListConfig sets the "models_list_config" field.
+func (u *GroupUpsertBulk) SetModelsListConfig(v domain.GroupModelsListConfig) *GroupUpsertBulk {
+ return u.Update(func(s *GroupUpsert) {
+ s.SetModelsListConfig(v)
+ })
+}
+
+// UpdateModelsListConfig sets the "models_list_config" field to the value that was provided on create.
+func (u *GroupUpsertBulk) UpdateModelsListConfig() *GroupUpsertBulk {
+ return u.Update(func(s *GroupUpsert) {
+ s.UpdateModelsListConfig()
+ })
+}
+
// SetRpmLimit sets the "rpm_limit" field.
func (u *GroupUpsertBulk) SetRpmLimit(v int) *GroupUpsertBulk {
return u.Update(func(s *GroupUpsert) {
diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go
index fe55982c..c10d60ec 100644
--- a/backend/ent/group_update.go
+++ b/backend/ent/group_update.go
@@ -616,6 +616,20 @@ func (_u *GroupUpdate) SetNillableMessagesDispatchModelConfig(v *domain.OpenAIMe
return _u
}
+// SetModelsListConfig sets the "models_list_config" field.
+func (_u *GroupUpdate) SetModelsListConfig(v domain.GroupModelsListConfig) *GroupUpdate {
+ _u.mutation.SetModelsListConfig(v)
+ return _u
+}
+
+// SetNillableModelsListConfig sets the "models_list_config" field if the given value is not nil.
+func (_u *GroupUpdate) SetNillableModelsListConfig(v *domain.GroupModelsListConfig) *GroupUpdate {
+ if v != nil {
+ _u.SetModelsListConfig(*v)
+ }
+ return _u
+}
+
// SetRpmLimit sets the "rpm_limit" field.
func (_u *GroupUpdate) SetRpmLimit(v int) *GroupUpdate {
_u.mutation.ResetRpmLimit()
@@ -1112,6 +1126,9 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) {
if value, ok := _u.mutation.MessagesDispatchModelConfig(); ok {
_spec.SetField(group.FieldMessagesDispatchModelConfig, field.TypeJSON, value)
}
+ if value, ok := _u.mutation.ModelsListConfig(); ok {
+ _spec.SetField(group.FieldModelsListConfig, field.TypeJSON, value)
+ }
if value, ok := _u.mutation.RpmLimit(); ok {
_spec.SetField(group.FieldRpmLimit, field.TypeInt, value)
}
@@ -2012,6 +2029,20 @@ func (_u *GroupUpdateOne) SetNillableMessagesDispatchModelConfig(v *domain.OpenA
return _u
}
+// SetModelsListConfig sets the "models_list_config" field.
+func (_u *GroupUpdateOne) SetModelsListConfig(v domain.GroupModelsListConfig) *GroupUpdateOne {
+ _u.mutation.SetModelsListConfig(v)
+ return _u
+}
+
+// SetNillableModelsListConfig sets the "models_list_config" field if the given value is not nil.
+func (_u *GroupUpdateOne) SetNillableModelsListConfig(v *domain.GroupModelsListConfig) *GroupUpdateOne {
+ if v != nil {
+ _u.SetModelsListConfig(*v)
+ }
+ return _u
+}
+
// SetRpmLimit sets the "rpm_limit" field.
func (_u *GroupUpdateOne) SetRpmLimit(v int) *GroupUpdateOne {
_u.mutation.ResetRpmLimit()
@@ -2538,6 +2569,9 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error)
if value, ok := _u.mutation.MessagesDispatchModelConfig(); ok {
_spec.SetField(group.FieldMessagesDispatchModelConfig, field.TypeJSON, value)
}
+ if value, ok := _u.mutation.ModelsListConfig(); ok {
+ _spec.SetField(group.FieldModelsListConfig, field.TypeJSON, value)
+ }
if value, ok := _u.mutation.RpmLimit(); ok {
_spec.SetField(group.FieldRpmLimit, field.TypeInt, value)
}
diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go
index 447f71ef..7abe4c60 100644
--- a/backend/ent/migrate/schema.go
+++ b/backend/ent/migrate/schema.go
@@ -669,6 +669,7 @@ var (
{Name: "require_privacy_set", Type: field.TypeBool, Default: false},
{Name: "default_mapped_model", Type: field.TypeString, Size: 100, Default: ""},
{Name: "messages_dispatch_model_config", Type: field.TypeJSON, SchemaType: map[string]string{"postgres": "jsonb"}},
+ {Name: "models_list_config", Type: field.TypeJSON, SchemaType: map[string]string{"postgres": "jsonb"}},
{Name: "rpm_limit", Type: field.TypeInt, Default: 0},
}
// GroupsTable holds the schema information for the "groups" table.
diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go
index 2e8fa7f4..003e25d5 100644
--- a/backend/ent/mutation.go
+++ b/backend/ent/mutation.go
@@ -14901,6 +14901,7 @@ type GroupMutation struct {
require_privacy_set *bool
default_mapped_model *string
messages_dispatch_model_config *domain.OpenAIMessagesDispatchModelConfig
+ models_list_config *domain.GroupModelsListConfig
rpm_limit *int
addrpm_limit *int
clearedFields map[string]struct{}
@@ -16619,6 +16620,42 @@ func (m *GroupMutation) ResetMessagesDispatchModelConfig() {
m.messages_dispatch_model_config = nil
}
+// SetModelsListConfig sets the "models_list_config" field.
+func (m *GroupMutation) SetModelsListConfig(dmlc domain.GroupModelsListConfig) {
+ m.models_list_config = &dmlc
+}
+
+// ModelsListConfig returns the value of the "models_list_config" field in the mutation.
+func (m *GroupMutation) ModelsListConfig() (r domain.GroupModelsListConfig, exists bool) {
+ v := m.models_list_config
+ if v == nil {
+ return
+ }
+ return *v, true
+}
+
+// OldModelsListConfig returns the old "models_list_config" field's value of the Group entity.
+// If the Group object wasn't provided to the builder, the object is fetched from the database.
+// An error is returned if the mutation operation is not UpdateOne, or the database query fails.
+func (m *GroupMutation) OldModelsListConfig(ctx context.Context) (v domain.GroupModelsListConfig, err error) {
+ if !m.op.Is(OpUpdateOne) {
+ return v, errors.New("OldModelsListConfig is only allowed on UpdateOne operations")
+ }
+ if m.id == nil || m.oldValue == nil {
+ return v, errors.New("OldModelsListConfig requires an ID field in the mutation")
+ }
+ oldValue, err := m.oldValue(ctx)
+ if err != nil {
+ return v, fmt.Errorf("querying old value for OldModelsListConfig: %w", err)
+ }
+ return oldValue.ModelsListConfig, nil
+}
+
+// ResetModelsListConfig resets all changes to the "models_list_config" field.
+func (m *GroupMutation) ResetModelsListConfig() {
+ m.models_list_config = nil
+}
+
// SetRpmLimit sets the "rpm_limit" field.
func (m *GroupMutation) SetRpmLimit(i int) {
m.rpm_limit = &i
@@ -17033,7 +17070,7 @@ func (m *GroupMutation) Type() string {
// order to get all numeric fields that were incremented/decremented, call
// AddedFields().
func (m *GroupMutation) Fields() []string {
- fields := make([]string, 0, 34)
+ fields := make([]string, 0, 35)
if m.created_at != nil {
fields = append(fields, group.FieldCreatedAt)
}
@@ -17133,6 +17170,9 @@ func (m *GroupMutation) Fields() []string {
if m.messages_dispatch_model_config != nil {
fields = append(fields, group.FieldMessagesDispatchModelConfig)
}
+ if m.models_list_config != nil {
+ fields = append(fields, group.FieldModelsListConfig)
+ }
if m.rpm_limit != nil {
fields = append(fields, group.FieldRpmLimit)
}
@@ -17210,6 +17250,8 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) {
return m.DefaultMappedModel()
case group.FieldMessagesDispatchModelConfig:
return m.MessagesDispatchModelConfig()
+ case group.FieldModelsListConfig:
+ return m.ModelsListConfig()
case group.FieldRpmLimit:
return m.RpmLimit()
}
@@ -17287,6 +17329,8 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e
return m.OldDefaultMappedModel(ctx)
case group.FieldMessagesDispatchModelConfig:
return m.OldMessagesDispatchModelConfig(ctx)
+ case group.FieldModelsListConfig:
+ return m.OldModelsListConfig(ctx)
case group.FieldRpmLimit:
return m.OldRpmLimit(ctx)
}
@@ -17529,6 +17573,13 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error {
}
m.SetMessagesDispatchModelConfig(v)
return nil
+ case group.FieldModelsListConfig:
+ v, ok := value.(domain.GroupModelsListConfig)
+ if !ok {
+ return fmt.Errorf("unexpected type %T for field %s", value, name)
+ }
+ m.SetModelsListConfig(v)
+ return nil
case group.FieldRpmLimit:
v, ok := value.(int)
if !ok {
@@ -17912,6 +17963,9 @@ func (m *GroupMutation) ResetField(name string) error {
case group.FieldMessagesDispatchModelConfig:
m.ResetMessagesDispatchModelConfig()
return nil
+ case group.FieldModelsListConfig:
+ m.ResetModelsListConfig()
+ return nil
case group.FieldRpmLimit:
m.ResetRpmLimit()
return nil
diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go
index aa6130f0..fdb837e8 100644
--- a/backend/ent/runtime/runtime.go
+++ b/backend/ent/runtime/runtime.go
@@ -870,8 +870,12 @@ func init() {
groupDescMessagesDispatchModelConfig := groupFields[29].Descriptor()
// group.DefaultMessagesDispatchModelConfig holds the default value on creation for the messages_dispatch_model_config field.
group.DefaultMessagesDispatchModelConfig = groupDescMessagesDispatchModelConfig.Default.(domain.OpenAIMessagesDispatchModelConfig)
+ // groupDescModelsListConfig is the schema descriptor for models_list_config field.
+ groupDescModelsListConfig := groupFields[30].Descriptor()
+ // group.DefaultModelsListConfig holds the default value on creation for the models_list_config field.
+ group.DefaultModelsListConfig = groupDescModelsListConfig.Default.(domain.GroupModelsListConfig)
// groupDescRpmLimit is the schema descriptor for rpm_limit field.
- groupDescRpmLimit := groupFields[30].Descriptor()
+ groupDescRpmLimit := groupFields[31].Descriptor()
// group.DefaultRpmLimit holds the default value on creation for the rpm_limit field.
group.DefaultRpmLimit = groupDescRpmLimit.Default.(int)
idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin()
diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go
index d47e8710..2a1715f8 100644
--- a/backend/ent/schema/group.go
+++ b/backend/ent/schema/group.go
@@ -155,6 +155,10 @@ func (Group) Fields() []ent.Field {
Default(domain.OpenAIMessagesDispatchModelConfig{}).
SchemaType(map[string]string{dialect.Postgres: "jsonb"}).
Comment("OpenAI Messages 调度模型配置:按 Claude 系列/精确模型映射到目标 GPT 模型"),
+ field.JSON("models_list_config", domain.GroupModelsListConfig{}).
+ Default(domain.GroupModelsListConfig{}).
+ SchemaType(map[string]string{dialect.Postgres: "jsonb"}).
+ Comment("自定义 /v1/models 展示列表配置;仅影响模型列表响应,不影响调度"),
// 分组级每分钟请求数上限(0 = 不限制)。设置后优先于用户级兜底生效。
field.Int("rpm_limit").
diff --git a/backend/go.mod b/backend/go.mod
index 587d5370..62be56c8 100644
--- a/backend/go.mod
+++ b/backend/go.mod
@@ -1,6 +1,6 @@
module github.com/Wei-Shaw/sub2api
-go 1.26.3
+go 1.26.4
require (
entgo.io/ent v0.14.5
diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go
index 7b275c83..13c541f9 100644
--- a/backend/internal/config/config.go
+++ b/backend/internal/config/config.go
@@ -652,6 +652,9 @@ type BillingConfig struct {
// - billing_cache_service.checkUserPlatformQuotaEligibility 首次缓存装载
// 读写两端必须共用同一 TTL,避免缓存生命周期不一致导致 quota 计数漂移。
UserPlatformQuotaCacheTTLSeconds int `mapstructure:"user_platform_quota_cache_ttl_seconds"`
+ // UserPlatformQuotaSentinelTTLSeconds sentinel(无 limit 占位)entry 的 TTL,
+ // 显著短于 quota cache 默认 86400s 以控 Redis 内存;默认 3600=1h。
+ UserPlatformQuotaSentinelTTLSeconds int `mapstructure:"user_platform_quota_sentinel_ttl_seconds"`
}
type CircuitBreakerConfig struct {
@@ -719,6 +722,8 @@ type GatewayConfig struct {
OpenAIPassthroughAllowTimeoutHeaders bool `mapstructure:"openai_passthrough_allow_timeout_headers"`
// OpenAIWS: OpenAI Responses WebSocket 配置(默认开启,可按需回滚到 HTTP)
OpenAIWS GatewayOpenAIWSConfig `mapstructure:"openai_ws"`
+ // OpenAIScheduler: OpenAI 高级调度器粘性逃逸配置
+ OpenAIScheduler GatewayOpenAISchedulerConfig `mapstructure:"openai_scheduler"`
// OpenAIHTTP2: OpenAI HTTP 上游协议策略(默认启用 HTTP/2,可按代理能力回退 HTTP/1.1)
OpenAIHTTP2 GatewayOpenAIHTTP2Config `mapstructure:"openai_http2"`
// ImageConcurrency: 图片生成独立并发限制配置(默认关闭)
@@ -885,6 +890,12 @@ type GatewayOpenAIWSConfig struct {
StoreDisabledForceNewConn bool `mapstructure:"store_disabled_force_new_conn"`
// PrewarmGenerateEnabled: 是否启用 WSv2 generate=false 预热(默认 false)
PrewarmGenerateEnabled bool `mapstructure:"prewarm_generate_enabled"`
+ // ClientReadLimitBytes: 入站客户端 WS 单帧读取上限。
+ ClientReadLimitBytes int64 `mapstructure:"client_read_limit_bytes"`
+ // HTTPBridgeEnabled: 首包过大时,保持客户端 WS,改用 HTTP Responses 上游。
+ HTTPBridgeEnabled bool `mapstructure:"http_bridge_enabled"`
+ // HTTPBridgeThresholdBytes: 触发 HTTP bridge 的入站 WS payload 阈值。
+ HTTPBridgeThresholdBytes int64 `mapstructure:"http_bridge_threshold_bytes"`
// Feature 开关:v2 优先于 v1
ResponsesWebsockets bool `mapstructure:"responses_websockets"`
@@ -951,6 +962,16 @@ type GatewayOpenAIWSSchedulerScoreWeights struct {
TTFT float64 `mapstructure:"ttft"`
}
+// GatewayOpenAISchedulerConfig OpenAI 高级调度器配置。
+type GatewayOpenAISchedulerConfig struct {
+ // StickyEscapeEnabled: 是否允许 session_hash sticky 在账号健康度劣化时临时逃逸
+ StickyEscapeEnabled bool `mapstructure:"sticky_escape_enabled"`
+ // StickyEscapeTTFTMs: TTFT EWMA 超过该阈值时跳过 sticky
+ StickyEscapeTTFTMs int `mapstructure:"sticky_escape_ttft_ms"`
+ // StickyEscapeErrorRate: 错误率 EWMA 超过该阈值时跳过 sticky
+ StickyEscapeErrorRate float64 `mapstructure:"sticky_escape_error_rate"`
+}
+
// GatewayUsageRecordConfig 使用量记录异步队列配置
type GatewayUsageRecordConfig struct {
// WorkerCount: worker 初始数量(自动扩缩容开启时作为初始并发上限)
@@ -1094,6 +1115,13 @@ type DatabaseConfig struct {
ConnMaxLifetimeMinutes int `mapstructure:"conn_max_lifetime_minutes"`
// ConnMaxIdleTimeMinutes: 空闲连接最大存活时间,及时释放不活跃连接
ConnMaxIdleTimeMinutes int `mapstructure:"conn_max_idle_time_minutes"`
+ // UserPlatformQuotaFlusherEnabled: 是否启用 user×platform 配额写聚合 flusher
+ UserPlatformQuotaFlusherEnabled bool `mapstructure:"user_platform_quota_flusher_enabled"`
+ // UserPlatformQuotaFlushIntervalMs: flusher 刷写间隔(毫秒)
+ UserPlatformQuotaFlushIntervalMs int `mapstructure:"user_platform_quota_flush_interval_ms"`
+ // UserPlatformQuotaFlushBatchSize: flusher 单批最大条数
+ // 建议 ≤ 6000(单条 UPSERT 原子上限)
+ UserPlatformQuotaFlushBatchSize int `mapstructure:"user_platform_quota_flush_batch_size"`
}
func (d *DatabaseConfig) DSN() string {
@@ -1372,6 +1400,15 @@ func load(allowMissingJWTSecret bool) (*Config, error) {
if err := viper.Unmarshal(&cfg); err != nil {
return nil, fmt.Errorf("unmarshal config error: %w", err)
}
+ if cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs == 0 {
+ cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs = 15000
+ }
+ if cfg.Gateway.OpenAIScheduler.StickyEscapeErrorRate == 0 {
+ cfg.Gateway.OpenAIScheduler.StickyEscapeErrorRate = 0.5
+ }
+ if !cfg.Gateway.OpenAIScheduler.StickyEscapeEnabled && !viper.IsSet("gateway.openai_scheduler.sticky_escape_enabled") {
+ cfg.Gateway.OpenAIScheduler.StickyEscapeEnabled = true
+ }
cfg.RunMode = NormalizeRunMode(cfg.RunMode)
cfg.Server.Mode = strings.ToLower(strings.TrimSpace(cfg.Server.Mode))
@@ -1581,6 +1618,7 @@ func setDefaults() {
viper.SetDefault("billing.circuit_breaker.reset_timeout_seconds", 30)
viper.SetDefault("billing.circuit_breaker.half_open_requests", 3)
viper.SetDefault("billing.user_platform_quota_cache_ttl_seconds", 86400)
+ viper.SetDefault("billing.user_platform_quota_sentinel_ttl_seconds", 3600)
// Turnstile
viper.SetDefault("turnstile.required", false)
@@ -1667,6 +1705,9 @@ func setDefaults() {
viper.SetDefault("database.max_idle_conns", 128)
viper.SetDefault("database.conn_max_lifetime_minutes", 30)
viper.SetDefault("database.conn_max_idle_time_minutes", 5)
+ viper.SetDefault("database.user_platform_quota_flusher_enabled", false)
+ viper.SetDefault("database.user_platform_quota_flush_interval_ms", 2000)
+ viper.SetDefault("database.user_platform_quota_flush_batch_size", 1000)
// Redis
viper.SetDefault("redis.host", "localhost")
@@ -1802,6 +1843,9 @@ func setDefaults() {
viper.SetDefault("gateway.openai_ws.store_disabled_conn_mode", "strict")
viper.SetDefault("gateway.openai_ws.store_disabled_force_new_conn", true)
viper.SetDefault("gateway.openai_ws.prewarm_generate_enabled", false)
+ viper.SetDefault("gateway.openai_ws.client_read_limit_bytes", 64*1024*1024)
+ viper.SetDefault("gateway.openai_ws.http_bridge_enabled", true)
+ viper.SetDefault("gateway.openai_ws.http_bridge_threshold_bytes", 15*1024*1024)
viper.SetDefault("gateway.openai_ws.responses_websockets", false)
viper.SetDefault("gateway.openai_ws.responses_websockets_v2", true)
viper.SetDefault("gateway.openai_ws.max_conns_per_account", 128)
@@ -2539,6 +2583,15 @@ func (c *Config) Validate() error {
if c.Gateway.OpenAIWS.PrewarmCooldownMS < 0 {
return fmt.Errorf("gateway.openai_ws.prewarm_cooldown_ms must be non-negative")
}
+ if c.Gateway.OpenAIWS.ClientReadLimitBytes <= 0 {
+ return fmt.Errorf("gateway.openai_ws.client_read_limit_bytes must be positive")
+ }
+ if c.Gateway.OpenAIWS.HTTPBridgeThresholdBytes < 0 {
+ return fmt.Errorf("gateway.openai_ws.http_bridge_threshold_bytes must be non-negative")
+ }
+ if c.Gateway.OpenAIWS.HTTPBridgeEnabled && c.Gateway.OpenAIWS.HTTPBridgeThresholdBytes == 0 {
+ return fmt.Errorf("gateway.openai_ws.http_bridge_threshold_bytes must be positive when http_bridge_enabled is true")
+ }
if c.Gateway.OpenAIWS.FallbackCooldownSeconds < 0 {
return fmt.Errorf("gateway.openai_ws.fallback_cooldown_seconds must be non-negative")
}
@@ -2613,6 +2666,12 @@ func (c *Config) Validate() error {
if weightSum <= 0 {
return fmt.Errorf("gateway.openai_ws.scheduler_score_weights must not all be zero")
}
+ if c.Gateway.OpenAIScheduler.StickyEscapeTTFTMs <= 0 {
+ return fmt.Errorf("gateway.openai_scheduler.sticky_escape_ttft_ms must be positive")
+ }
+ if c.Gateway.OpenAIScheduler.StickyEscapeErrorRate < 0 || c.Gateway.OpenAIScheduler.StickyEscapeErrorRate > 1 {
+ return fmt.Errorf("gateway.openai_scheduler.sticky_escape_error_rate must be between 0 and 1")
+ }
if c.Gateway.MaxLineSize < 0 {
return fmt.Errorf("gateway.max_line_size must be non-negative")
}
diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go
index a969e3d7..a4bfcd60 100644
--- a/backend/internal/config/config_test.go
+++ b/backend/internal/config/config_test.go
@@ -110,6 +110,15 @@ func TestLoadDefaultOpenAIWSConfig(t *testing.T) {
if cfg.Gateway.OpenAIWS.StickySessionTTLSeconds != 3600 {
t.Fatalf("Gateway.OpenAIWS.StickySessionTTLSeconds = %d, want 3600", cfg.Gateway.OpenAIWS.StickySessionTTLSeconds)
}
+ if !cfg.Gateway.OpenAIScheduler.StickyEscapeEnabled {
+ t.Fatalf("Gateway.OpenAIScheduler.StickyEscapeEnabled = false, want true")
+ }
+ if cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs != 15000 {
+ t.Fatalf("Gateway.OpenAIScheduler.StickyEscapeTTFTMs = %d, want 15000", cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs)
+ }
+ if cfg.Gateway.OpenAIScheduler.StickyEscapeErrorRate != 0.5 {
+ t.Fatalf("Gateway.OpenAIScheduler.StickyEscapeErrorRate = %v, want 0.5", cfg.Gateway.OpenAIScheduler.StickyEscapeErrorRate)
+ }
if !cfg.Gateway.OpenAIWS.SessionHashReadOldFallback {
t.Fatalf("Gateway.OpenAIWS.SessionHashReadOldFallback = false, want true")
}
@@ -134,6 +143,15 @@ func TestLoadDefaultOpenAIWSConfig(t *testing.T) {
if cfg.Gateway.OpenAIWS.PrewarmCooldownMS != 300 {
t.Fatalf("Gateway.OpenAIWS.PrewarmCooldownMS = %d, want 300", cfg.Gateway.OpenAIWS.PrewarmCooldownMS)
}
+ if cfg.Gateway.OpenAIWS.ClientReadLimitBytes != 64*1024*1024 {
+ t.Fatalf("Gateway.OpenAIWS.ClientReadLimitBytes = %d, want %d", cfg.Gateway.OpenAIWS.ClientReadLimitBytes, 64*1024*1024)
+ }
+ if !cfg.Gateway.OpenAIWS.HTTPBridgeEnabled {
+ t.Fatalf("Gateway.OpenAIWS.HTTPBridgeEnabled = false, want true")
+ }
+ if cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes != 15*1024*1024 {
+ t.Fatalf("Gateway.OpenAIWS.HTTPBridgeThresholdBytes = %d, want %d", cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes, 15*1024*1024)
+ }
if cfg.Gateway.OpenAIWS.RetryBackoffInitialMS != 120 {
t.Fatalf("Gateway.OpenAIWS.RetryBackoffInitialMS = %d, want 120", cfg.Gateway.OpenAIWS.RetryBackoffInitialMS)
}
@@ -1720,6 +1738,21 @@ func TestValidateConfig_OpenAIWSRules(t *testing.T) {
},
wantErr: "gateway.openai_ws.scheduler_score_weights must not all be zero",
},
+ {
+ name: "sticky_escape_ttft_ms 必须为正数",
+ mutate: func(c *Config) { c.Gateway.OpenAIScheduler.StickyEscapeTTFTMs = 0 },
+ wantErr: "gateway.openai_scheduler.sticky_escape_ttft_ms",
+ },
+ {
+ name: "sticky_escape_error_rate 不能小于 0",
+ mutate: func(c *Config) { c.Gateway.OpenAIScheduler.StickyEscapeErrorRate = -0.1 },
+ wantErr: "gateway.openai_scheduler.sticky_escape_error_rate",
+ },
+ {
+ name: "sticky_escape_error_rate 不能大于 1",
+ mutate: func(c *Config) { c.Gateway.OpenAIScheduler.StickyEscapeErrorRate = 1.1 },
+ wantErr: "gateway.openai_scheduler.sticky_escape_error_rate",
+ },
}
for _, tc := range cases {
diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go
index 27c543dd..7601f35b 100644
--- a/backend/internal/domain/constants.go
+++ b/backend/internal/domain/constants.go
@@ -72,6 +72,7 @@ const (
// 与前端 useModelWhitelist.ts 中的 antigravityDefaultMappings 保持一致
var DefaultAntigravityModelMapping = map[string]string{
// Claude 白名单
+ "claude-opus-4-8": "claude-opus-4-8", // 官方模型
"claude-opus-4-7": "claude-opus-4-7", // 官方模型
"claude-opus-4-6-thinking": "claude-opus-4-6-thinking", // 官方模型
"claude-opus-4-6": "claude-opus-4-6-thinking", // 简称映射
@@ -122,6 +123,7 @@ var DefaultAntigravityModelMapping = map[string]string{
// aws_region 自动调整为匹配的区域前缀(如 eu.、apac.、jp. 等)
var DefaultBedrockModelMapping = map[string]string{
// Claude Opus
+ "claude-opus-4-8": "us.anthropic.claude-opus-4-8-v1",
"claude-opus-4-7": "us.anthropic.claude-opus-4-7-v1",
"claude-opus-4-6-thinking": "us.anthropic.claude-opus-4-6-v1",
"claude-opus-4-6": "us.anthropic.claude-opus-4-6-v1",
diff --git a/backend/internal/domain/constants_test.go b/backend/internal/domain/constants_test.go
index de66137f..fe6272c5 100644
--- a/backend/internal/domain/constants_test.go
+++ b/backend/internal/domain/constants_test.go
@@ -24,3 +24,27 @@ func TestDefaultAntigravityModelMapping_ImageCompatibilityAliases(t *testing.T)
}
}
}
+
+func TestDefaultAntigravityModelMapping_ContainsOpus48(t *testing.T) {
+ t.Parallel()
+
+ got, ok := DefaultAntigravityModelMapping["claude-opus-4-8"]
+ if !ok {
+ t.Fatal("expected mapping for claude-opus-4-8 to exist")
+ }
+ if got != "claude-opus-4-8" {
+ t.Fatalf("unexpected claude-opus-4-8 mapping: got %q", got)
+ }
+}
+
+func TestDefaultBedrockModelMapping_ContainsOpus48(t *testing.T) {
+ t.Parallel()
+
+ got, ok := DefaultBedrockModelMapping["claude-opus-4-8"]
+ if !ok {
+ t.Fatal("expected Bedrock mapping for claude-opus-4-8 to exist")
+ }
+ if got != "us.anthropic.claude-opus-4-8-v1" {
+ t.Fatalf("unexpected Bedrock claude-opus-4-8 mapping: got %q", got)
+ }
+}
diff --git a/backend/internal/domain/models_list_config.go b/backend/internal/domain/models_list_config.go
new file mode 100644
index 00000000..3f050585
--- /dev/null
+++ b/backend/internal/domain/models_list_config.go
@@ -0,0 +1,7 @@
+package domain
+
+// GroupModelsListConfig controls the optional custom /v1/models response list.
+type GroupModelsListConfig struct {
+ Enabled bool `json:"enabled"`
+ Models []string `json:"models,omitempty"`
+}
diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go
index 4f566a8b..57195342 100644
--- a/backend/internal/handler/admin/account_handler.go
+++ b/backend/internal/handler/admin/account_handler.go
@@ -2131,6 +2131,56 @@ func (h *AccountHandler) SyncUpstreamModels(c *gin.Context) {
response.Success(c, gin.H{"models": models})
}
+// SyncUpstreamModelsPreview handles syncing live supported models using provided credentials (no account ID needed).
+// POST /api/v1/admin/accounts/models/sync-upstream-preview
+func (h *AccountHandler) SyncUpstreamModelsPreview(c *gin.Context) {
+ var req struct {
+ Platform string `json:"platform" binding:"required"`
+ Type string `json:"type" binding:"required"`
+ BaseURL string `json:"base_url"`
+ APIKey string `json:"api_key" binding:"required"`
+ }
+ if err := c.ShouldBindJSON(&req); err != nil {
+ response.BadRequest(c, "Invalid request: "+err.Error())
+ return
+ }
+
+ tempAccount := &service.Account{
+ Platform: req.Platform,
+ Type: req.Type,
+ Credentials: map[string]any{
+ "api_key": req.APIKey,
+ "base_url": req.BaseURL,
+ },
+ }
+
+ if h.accountTestService == nil {
+ response.InternalError(c, "Account test service is not configured")
+ return
+ }
+
+ models, err := h.accountTestService.FetchUpstreamSupportedModels(c.Request.Context(), tempAccount)
+ if err != nil {
+ var syncErr *service.UpstreamModelSyncError
+ if errors.As(err, &syncErr) {
+ switch syncErr.Kind {
+ case service.UpstreamModelSyncErrorConfiguration, service.UpstreamModelSyncErrorUnsupported:
+ response.BadRequest(c, syncErr.SafeMessage())
+ default:
+ slog.Warn("sync_upstream_models_preview_failed", "platform", req.Platform, "kind", syncErr.Kind)
+ response.Error(c, http.StatusBadGateway, syncErr.SafeMessage())
+ }
+ return
+ }
+
+ slog.Warn("sync_upstream_models_preview_failed", "platform", req.Platform)
+ response.Error(c, http.StatusBadGateway, "Failed to sync upstream models from upstream")
+ return
+ }
+
+ response.Success(c, gin.H{"models": models})
+}
+
// SetPrivacy handles setting privacy for a single OpenAI/Antigravity OAuth account
// POST /api/v1/admin/accounts/:id/set-privacy
func (h *AccountHandler) SetPrivacy(c *gin.Context) {
diff --git a/backend/internal/handler/admin/account_handler_list_test.go b/backend/internal/handler/admin/account_handler_list_test.go
new file mode 100644
index 00000000..4d628365
--- /dev/null
+++ b/backend/internal/handler/admin/account_handler_list_test.go
@@ -0,0 +1,52 @@
+package admin
+
+import (
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+ "time"
+
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+func setupAccountListRouter() (*gin.Engine, *stubAdminService) {
+ gin.SetMode(gin.TestMode)
+ router := gin.New()
+ adminSvc := newStubAdminService()
+ handler := NewAccountHandler(adminSvc, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
+ router.GET("/api/v1/admin/accounts", handler.List)
+ return router, adminSvc
+}
+
+func TestAccountHandlerListIncludesCreatedAt(t *testing.T) {
+ router, adminSvc := setupAccountListRouter()
+
+ rec := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/accounts?page=1&page_size=20&sort_by=created_at&sort_order=desc", nil)
+ router.ServeHTTP(rec, req)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+ require.Equal(t, "created_at", adminSvc.lastListAccounts.sortBy)
+
+ var payload struct {
+ Data struct {
+ Items []struct {
+ ID int64 `json:"id"`
+ CreatedAt string `json:"created_at"`
+ } `json:"items"`
+ } `json:"data"`
+ }
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
+ require.Len(t, payload.Data.Items, 1)
+
+ createdAt := payload.Data.Items[0].CreatedAt
+ require.NotEmpty(t, createdAt)
+ require.True(t, strings.HasSuffix(createdAt, "Z"), "created_at should be serialized as UTC")
+ parsed, err := time.Parse(time.RFC3339Nano, createdAt)
+ require.NoError(t, err)
+ _, offset := parsed.Zone()
+ require.Equal(t, 0, offset)
+}
diff --git a/backend/internal/handler/admin/admin_basic_handlers_test.go b/backend/internal/handler/admin/admin_basic_handlers_test.go
index 7b74bafc..bffddc8a 100644
--- a/backend/internal/handler/admin/admin_basic_handlers_test.go
+++ b/backend/internal/handler/admin/admin_basic_handlers_test.go
@@ -33,6 +33,7 @@ func setupAdminRouter() (*gin.Engine, *stubAdminService) {
router.GET("/api/v1/admin/groups", groupHandler.List)
router.GET("/api/v1/admin/groups/all", groupHandler.GetAll)
+ router.GET("/api/v1/admin/groups/:id/models-list-candidates", groupHandler.GetModelsListCandidates)
router.GET("/api/v1/admin/groups/:id", groupHandler.GetByID)
router.POST("/api/v1/admin/groups", groupHandler.Create)
router.PUT("/api/v1/admin/groups/:id", groupHandler.Update)
@@ -177,6 +178,12 @@ func TestGroupHandlerEndpoints(t *testing.T) {
router.ServeHTTP(rec, req)
require.Equal(t, http.StatusOK, rec.Code)
+ rec = httptest.NewRecorder()
+ req = httptest.NewRequest(http.MethodGet, "/api/v1/admin/groups/0/models-list-candidates?platform=openai", nil)
+ router.ServeHTTP(rec, req)
+ require.Equal(t, http.StatusOK, rec.Code)
+ require.Contains(t, rec.Body.String(), "gpt-5.5")
+
body, _ := json.Marshal(map[string]any{"name": "new", "platform": "anthropic", "subscription_type": "standard"})
rec = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodPost, "/api/v1/admin/groups", bytes.NewReader(body))
diff --git a/backend/internal/handler/admin/admin_service_stub_test.go b/backend/internal/handler/admin/admin_service_stub_test.go
index 65b71492..819f0cdc 100644
--- a/backend/internal/handler/admin/admin_service_stub_test.go
+++ b/backend/internal/handler/admin/admin_service_stub_test.go
@@ -160,6 +160,10 @@ func (s *stubAdminService) GetUser(ctx context.Context, id int64) (*service.User
return &user, nil
}
+func (s *stubAdminService) GetUserIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ return s.GetUser(ctx, id)
+}
+
func (s *stubAdminService) CreateUser(ctx context.Context, input *service.CreateUserInput) (*service.User, error) {
user := service.User{ID: 100, Email: input.Email, Status: service.StatusActive}
return &user, nil
@@ -265,6 +269,13 @@ func (s *stubAdminService) GetGroup(ctx context.Context, id int64) (*service.Gro
return &group, nil
}
+func (s *stubAdminService) GetGroupModelsListCandidates(ctx context.Context, id int64, platform string) ([]string, error) {
+ if platform == service.PlatformOpenAI {
+ return []string{"gpt-5.5", "gpt-5.4"}, nil
+ }
+ return []string{"claude-sonnet-4-6"}, nil
+}
+
func (s *stubAdminService) CreateGroup(ctx context.Context, input *service.CreateGroupInput) (*service.Group, error) {
group := service.Group{ID: 200, Name: input.Name, Status: service.StatusActive}
return &group, nil
diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go
index 3667bbcd..102ee02f 100644
--- a/backend/internal/handler/admin/group_handler.go
+++ b/backend/internal/handler/admin/group_handler.go
@@ -113,6 +113,7 @@ type CreateGroupRequest struct {
RequirePrivacySet bool `json:"require_privacy_set"`
DefaultMappedModel string `json:"default_mapped_model"`
MessagesDispatchModelConfig service.OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config"`
+ ModelsListConfig service.GroupModelsListConfig `json:"models_list_config"`
// 分组 RPM 上限(0 = 不限制)
RPMLimit int `json:"rpm_limit"`
// 从指定分组复制账号(创建后自动绑定)
@@ -122,7 +123,7 @@ type CreateGroupRequest struct {
// UpdateGroupRequest represents update group request
type UpdateGroupRequest struct {
Name string `json:"name"`
- Description string `json:"description"`
+ Description *string `json:"description"`
Platform string `json:"platform" binding:"omitempty,oneof=anthropic openai gemini antigravity"`
RateMultiplier *float64 `json:"rate_multiplier"`
IsExclusive *bool `json:"is_exclusive"`
@@ -153,6 +154,7 @@ type UpdateGroupRequest struct {
RequirePrivacySet *bool `json:"require_privacy_set"`
DefaultMappedModel *string `json:"default_mapped_model"`
MessagesDispatchModelConfig *service.OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config"`
+ ModelsListConfig *service.GroupModelsListConfig `json:"models_list_config"`
// 分组 RPM 上限(0 = 不限制);nil 表示未提供不改动
RPMLimit *int `json:"rpm_limit"`
// 从指定分组复制账号(同步操作:先清空当前分组的账号绑定,再绑定源分组的账号)
@@ -238,6 +240,28 @@ func (h *GroupHandler) GetByID(c *gin.Context) {
response.Success(c, dto.GroupFromServiceAdmin(group))
}
+// GetModelsListCandidates handles getting candidate model IDs for custom /v1/models list.
+// GET /api/v1/admin/groups/:id/models-list-candidates
+func (h *GroupHandler) GetModelsListCandidates(c *gin.Context) {
+ groupID, err := strconv.ParseInt(c.Param("id"), 10, 64)
+ if err != nil || groupID < 0 {
+ response.BadRequest(c, "Invalid group ID")
+ return
+ }
+
+ models, err := h.adminService.GetGroupModelsListCandidates(
+ c.Request.Context(),
+ groupID,
+ c.Query("platform"),
+ )
+ if err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+
+ response.Success(c, gin.H{"models": models})
+}
+
// Create handles creating a new group
// POST /api/v1/admin/groups
func (h *GroupHandler) Create(c *gin.Context) {
@@ -275,6 +299,7 @@ func (h *GroupHandler) Create(c *gin.Context) {
RequirePrivacySet: req.RequirePrivacySet,
DefaultMappedModel: req.DefaultMappedModel,
MessagesDispatchModelConfig: req.MessagesDispatchModelConfig,
+ ModelsListConfig: req.ModelsListConfig,
RPMLimit: req.RPMLimit,
CopyAccountsFromGroupIDs: req.CopyAccountsFromGroupIDs,
})
@@ -330,6 +355,7 @@ func (h *GroupHandler) Update(c *gin.Context) {
RequirePrivacySet: req.RequirePrivacySet,
DefaultMappedModel: req.DefaultMappedModel,
MessagesDispatchModelConfig: req.MessagesDispatchModelConfig,
+ ModelsListConfig: req.ModelsListConfig,
RPMLimit: req.RPMLimit,
CopyAccountsFromGroupIDs: req.CopyAccountsFromGroupIDs,
})
diff --git a/backend/internal/handler/admin/ops_handler.go b/backend/internal/handler/admin/ops_handler.go
index 418c302f..b9558b97 100644
--- a/backend/internal/handler/admin/ops_handler.go
+++ b/backend/internal/handler/admin/ops_handler.go
@@ -110,6 +110,9 @@ func (h *OpsHandler) GetErrorLogs(c *gin.Context) {
filter.Source = strings.TrimSpace(c.Query("error_source"))
filter.Query = strings.TrimSpace(c.Query("q"))
filter.UserQuery = strings.TrimSpace(c.Query("user_query"))
+ // Model 过滤:admin 走精确匹配(ModelFuzzy 默认 false,保持管理端语义)。
+ // buildOpsErrorLogsWhere 以 COALESCE(requested_model, model) 比对。
+ filter.Model = strings.TrimSpace(c.Query("model"))
// Force request errors: client-visible status >= 400.
// buildOpsErrorLogsWhere already applies this for non-upstream phase.
@@ -137,6 +140,22 @@ func (h *OpsHandler) GetErrorLogs(c *gin.Context) {
filter.AccountID = &id
}
+ if v := strings.TrimSpace(c.Query("user_id")); v != "" {
+ id, err := strconv.ParseInt(v, 10, 64)
+ if err != nil || id <= 0 {
+ response.BadRequest(c, "Invalid user_id")
+ return
+ }
+ filter.UserID = &id
+ }
+ if v := strings.TrimSpace(c.Query("api_key_id")); v != "" {
+ id, err := strconv.ParseInt(v, 10, 64)
+ if err != nil || id <= 0 {
+ response.BadRequest(c, "Invalid api_key_id")
+ return
+ }
+ filter.APIKeyID = &id
+ }
if v := strings.TrimSpace(c.Query("resolved")); v != "" {
switch strings.ToLower(v) {
case "1", "true", "yes":
@@ -211,6 +230,9 @@ func (h *OpsHandler) ListRequestErrors(c *gin.Context) {
filter.Source = strings.TrimSpace(c.Query("error_source"))
filter.Query = strings.TrimSpace(c.Query("q"))
filter.UserQuery = strings.TrimSpace(c.Query("user_query"))
+ // Model 过滤:admin 走精确匹配(ModelFuzzy 默认 false,保持管理端语义)。
+ // buildOpsErrorLogsWhere 以 COALESCE(requested_model, model) 比对。
+ filter.Model = strings.TrimSpace(c.Query("model"))
// Force request errors: client-visible status >= 400.
// buildOpsErrorLogsWhere already applies this for non-upstream phase.
diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go
index 5e7a9cf5..47116dfb 100644
--- a/backend/internal/handler/admin/setting_handler.go
+++ b/backend/internal/handler/admin/setting_handler.go
@@ -256,6 +256,7 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
RewriteMessageCacheControl: settings.RewriteMessageCacheControl,
AntigravityUserAgentVersion: settings.AntigravityUserAgentVersion,
OpenAICodexUserAgent: settings.OpenAICodexUserAgent,
+ OpenAIAllowClaudeCodeCodexPlugin: settings.OpenAIAllowClaudeCodeCodexPlugin,
WebSearchEmulationEnabled: settings.WebSearchEmulationEnabled,
PaymentVisibleMethodAlipaySource: settings.PaymentVisibleMethodAlipaySource,
PaymentVisibleMethodWxpaySource: settings.PaymentVisibleMethodWxpaySource,
@@ -296,6 +297,8 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
AvailableChannelsEnabled: settings.AvailableChannelsEnabled,
AffiliateEnabled: settings.AffiliateEnabled,
+
+ AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests,
}
// OpenAI fast policy (stored under a dedicated setting key)
@@ -600,6 +603,7 @@ type UpdateSettingsRequest struct {
RewriteMessageCacheControl *bool `json:"rewrite_message_cache_control"`
AntigravityUserAgentVersion *string `json:"antigravity_user_agent_version"`
OpenAICodexUserAgent *string `json:"openai_codex_user_agent"`
+ OpenAIAllowClaudeCodeCodexPlugin *bool `json:"openai_allow_claude_code_codex_plugin"`
// Payment visible method routing
PaymentVisibleMethodAlipaySource *string `json:"payment_visible_method_alipay_source"`
@@ -672,6 +676,8 @@ type UpdateSettingsRequest struct {
AuthSourceGitHubPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_github_platform_quotas"`
AuthSourceGooglePlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_google_platform_quotas"`
AuthSourceDingTalkPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_dingtalk_platform_quotas"`
+
+ AllowUserViewErrorRequests *bool `json:"allow_user_view_error_requests"`
}
// UpdateSettings 更新系统设置
@@ -1635,6 +1641,12 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
MaxClaudeCodeVersion: req.MaxClaudeCodeVersion,
AllowUngroupedKeyScheduling: req.AllowUngroupedKeyScheduling,
BackendModeEnabled: req.BackendModeEnabled,
+ AllowUserViewErrorRequests: func() bool {
+ if req.AllowUserViewErrorRequests != nil {
+ return *req.AllowUserViewErrorRequests
+ }
+ return previousSettings.AllowUserViewErrorRequests
+ }(),
OpsMonitoringEnabled: func() bool {
if req.OpsMonitoringEnabled != nil {
return *req.OpsMonitoringEnabled
@@ -1701,6 +1713,12 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
}
return previousSettings.OpenAICodexUserAgent
}(),
+ OpenAIAllowClaudeCodeCodexPlugin: func() bool {
+ if req.OpenAIAllowClaudeCodeCodexPlugin != nil {
+ return *req.OpenAIAllowClaudeCodeCodexPlugin
+ }
+ return previousSettings.OpenAIAllowClaudeCodeCodexPlugin
+ }(),
PaymentVisibleMethodAlipaySource: func() string {
if req.PaymentVisibleMethodAlipaySource != nil {
return strings.TrimSpace(*req.PaymentVisibleMethodAlipaySource)
@@ -2077,6 +2095,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
RewriteMessageCacheControl: updatedSettings.RewriteMessageCacheControl,
AntigravityUserAgentVersion: updatedSettings.AntigravityUserAgentVersion,
OpenAICodexUserAgent: updatedSettings.OpenAICodexUserAgent,
+ OpenAIAllowClaudeCodeCodexPlugin: updatedSettings.OpenAIAllowClaudeCodeCodexPlugin,
PaymentVisibleMethodAlipaySource: updatedSettings.PaymentVisibleMethodAlipaySource,
PaymentVisibleMethodWxpaySource: updatedSettings.PaymentVisibleMethodWxpaySource,
PaymentVisibleMethodAlipayEnabled: updatedSettings.PaymentVisibleMethodAlipayEnabled,
@@ -2117,7 +2136,8 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
AffiliateEnabled: updatedSettings.AffiliateEnabled,
- RiskControlEnabled: updatedSettings.RiskControlEnabled,
+ RiskControlEnabled: updatedSettings.RiskControlEnabled,
+ AllowUserViewErrorRequests: updatedSettings.AllowUserViewErrorRequests,
}
if fastPolicy, err := h.settingService.GetOpenAIFastPolicySettings(c.Request.Context()); err != nil {
slog.Error("openai_fast_policy_settings_get_failed", "error", err)
@@ -2546,6 +2566,9 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings,
if before.OpenAICodexUserAgent != after.OpenAICodexUserAgent {
changed = append(changed, "openai_codex_user_agent")
}
+ if before.OpenAIAllowClaudeCodeCodexPlugin != after.OpenAIAllowClaudeCodeCodexPlugin {
+ changed = append(changed, "openai_allow_claude_code_codex_plugin")
+ }
if before.PaymentVisibleMethodAlipaySource != after.PaymentVisibleMethodAlipaySource {
changed = append(changed, "payment_visible_method_alipay_source")
}
diff --git a/backend/internal/handler/admin/system_handler.go b/backend/internal/handler/admin/system_handler.go
index 3e2022c7..fb6c0ef7 100644
--- a/backend/internal/handler/admin/system_handler.go
+++ b/backend/internal/handler/admin/system_handler.go
@@ -2,6 +2,7 @@ package admin
import (
"context"
+ "errors"
"net/http"
"strconv"
"strings"
@@ -17,12 +18,18 @@ import (
// SystemHandler handles system-related operations
type SystemHandler struct {
- updateSvc *service.UpdateService
+ updateSvc systemUpdateService
lockSvc *service.SystemOperationLockService
}
+type systemUpdateService interface {
+ CheckUpdate(ctx context.Context, force bool) (*service.UpdateInfo, error)
+ PerformUpdate(ctx context.Context) error
+ Rollback() error
+}
+
// NewSystemHandler creates a new SystemHandler
-func NewSystemHandler(updateSvc *service.UpdateService, lockSvc *service.SystemOperationLockService) *SystemHandler {
+func NewSystemHandler(updateSvc systemUpdateService, lockSvc *service.SystemOperationLockService) *SystemHandler {
return &SystemHandler{
updateSvc: updateSvc,
lockSvc: lockSvc,
@@ -67,6 +74,21 @@ func (h *SystemHandler) PerformUpdate(c *gin.Context) {
}()
if err := h.updateSvc.PerformUpdate(ctx); err != nil {
+ if errors.Is(err, service.ErrNoUpdateAvailable) {
+ info, checkErr := h.updateSvc.CheckUpdate(ctx, false)
+ if checkErr != nil {
+ releaseReason = "SYSTEM_UPDATE_FAILED"
+ return nil, checkErr
+ }
+ succeeded = true
+ return gin.H{
+ "message": "Already up to date",
+ "already_up_to_date": true,
+ "current_version": info.CurrentVersion,
+ "latest_version": info.LatestVersion,
+ "operation_id": lock.OperationID(),
+ }, nil
+ }
releaseReason = "SYSTEM_UPDATE_FAILED"
return nil, err
}
diff --git a/backend/internal/handler/admin/system_handler_test.go b/backend/internal/handler/admin/system_handler_test.go
new file mode 100644
index 00000000..0f33a452
--- /dev/null
+++ b/backend/internal/handler/admin/system_handler_test.go
@@ -0,0 +1,144 @@
+//go:build unit
+
+package admin
+
+import (
+ "context"
+ "encoding/json"
+ "errors"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+type systemHandlerUpdateServiceStub struct {
+ performErr error
+ updateInfo *service.UpdateInfo
+ checkErr error
+ checkForces []bool
+ performCall int
+}
+
+func (s *systemHandlerUpdateServiceStub) CheckUpdate(_ context.Context, force bool) (*service.UpdateInfo, error) {
+ s.checkForces = append(s.checkForces, force)
+ return s.updateInfo, s.checkErr
+}
+
+func (s *systemHandlerUpdateServiceStub) PerformUpdate(context.Context) error {
+ s.performCall++
+ return s.performErr
+}
+
+func (s *systemHandlerUpdateServiceStub) Rollback() error {
+ return nil
+}
+
+type systemUpdateResponseEnvelope struct {
+ Code int `json:"code"`
+ Message string `json:"message"`
+ Data struct {
+ Message string `json:"message"`
+ AlreadyUpToDate bool `json:"already_up_to_date"`
+ CurrentVersion string `json:"current_version"`
+ LatestVersion string `json:"latest_version"`
+ OperationID string `json:"operation_id"`
+ } `json:"data"`
+}
+
+type systemUpdateErrorEnvelope struct {
+ Code int `json:"code"`
+ Message string `json:"message"`
+}
+
+func newSystemHandlerTestRouter(t *testing.T, updateSvc *systemHandlerUpdateServiceStub, repo *memoryIdempotencyRepoStub) *gin.Engine {
+ t.Helper()
+ gin.SetMode(gin.TestMode)
+ service.SetDefaultIdempotencyCoordinator(nil)
+ t.Cleanup(func() {
+ service.SetDefaultIdempotencyCoordinator(nil)
+ })
+
+ lockSvc := service.NewSystemOperationLockService(repo, service.IdempotencyConfig{
+ ProcessingTimeout: time.Second,
+ SystemOperationTTL: time.Minute,
+ })
+ handler := NewSystemHandler(updateSvc, lockSvc)
+
+ router := gin.New()
+ router.POST("/api/v1/admin/system/update", handler.PerformUpdate)
+ return router
+}
+
+func requireSystemLockStatus(t *testing.T, repo *memoryIdempotencyRepoStub, wantStatus string) {
+ t.Helper()
+ repo.mu.Lock()
+ defer repo.mu.Unlock()
+
+ for _, record := range repo.data {
+ if record.Status == wantStatus {
+ return
+ }
+ }
+ t.Fatalf("system lock status %q not found in records: %#v", wantStatus, repo.data)
+}
+
+func TestSystemHandlerPerformUpdateAlreadyUpToDateReturnsOK(t *testing.T) {
+ updateSvc := &systemHandlerUpdateServiceStub{
+ performErr: service.ErrNoUpdateAvailable,
+ updateInfo: &service.UpdateInfo{
+ CurrentVersion: "0.1.132",
+ LatestVersion: "0.1.132",
+ HasUpdate: false,
+ },
+ }
+ repo := newMemoryIdempotencyRepoStub()
+ router := newSystemHandlerTestRouter(t, updateSvc, repo)
+
+ rec := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/system/update", nil)
+ req.Header.Set("Idempotency-Key", "already-up-to-date")
+ router.ServeHTTP(rec, req)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+ require.Equal(t, 1, updateSvc.performCall)
+ require.Equal(t, []bool{false}, updateSvc.checkForces)
+ requireSystemLockStatus(t, repo, service.IdempotencyStatusSucceeded)
+
+ var body systemUpdateResponseEnvelope
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &body))
+ require.Equal(t, 0, body.Code)
+ require.Equal(t, "success", body.Message)
+ require.Equal(t, "Already up to date", body.Data.Message)
+ require.True(t, body.Data.AlreadyUpToDate)
+ require.Equal(t, "0.1.132", body.Data.CurrentVersion)
+ require.Equal(t, "0.1.132", body.Data.LatestVersion)
+ require.NotEmpty(t, body.Data.OperationID)
+}
+
+func TestSystemHandlerPerformUpdateFailureStillReturnsInternalError(t *testing.T) {
+ updateSvc := &systemHandlerUpdateServiceStub{
+ performErr: errors.New("download failed"),
+ }
+ repo := newMemoryIdempotencyRepoStub()
+ router := newSystemHandlerTestRouter(t, updateSvc, repo)
+
+ rec := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/system/update", nil)
+ req.Header.Set("Idempotency-Key", "real-failure")
+ router.ServeHTTP(rec, req)
+
+ require.Equal(t, http.StatusInternalServerError, rec.Code)
+ require.Equal(t, 1, updateSvc.performCall)
+ require.Empty(t, updateSvc.checkForces)
+ requireSystemLockStatus(t, repo, service.IdempotencyStatusFailedRetryable)
+
+ var body systemUpdateErrorEnvelope
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &body))
+ require.Equal(t, http.StatusInternalServerError, body.Code)
+ require.Equal(t, "internal error", body.Message)
+}
diff --git a/backend/internal/handler/admin/usage_handler.go b/backend/internal/handler/admin/usage_handler.go
index 0857a138..11a4aeb8 100644
--- a/backend/internal/handler/admin/usage_handler.go
+++ b/backend/internal/handler/admin/usage_handler.go
@@ -325,10 +325,24 @@ func (h *UsageHandler) Stats(c *gin.Context) {
EndTime: &endTime,
}
- stats, err := h.usageService.GetStatsWithFilters(c.Request.Context(), filters)
- if err != nil {
- response.ErrorFrom(c, err)
- return
+ var stats *usagestats.UsageStats
+ // nocache: 绕过缓存直接回源,刷新者本人拿最新;不回写缓存(管理台"我刷新我自己拿最新"语义,非全局失效)。
+ if parseBoolQueryWithDefault(c.Query("nocache"), false) {
+ s, err := h.usageService.GetStatsWithFilters(c.Request.Context(), filters)
+ if err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+ stats = s
+ c.Header("X-Usage-Stats-Cache", "bypass")
+ } else {
+ s, hit, err := h.getStatsCached(c.Request.Context(), filters)
+ if err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+ stats = s
+ c.Header("X-Usage-Stats-Cache", cacheStatusValue(hit))
}
response.Success(c, stats)
@@ -344,23 +358,25 @@ func (h *UsageHandler) SearchUsers(c *gin.Context) {
}
// Limit to 30 results
- users, _, err := h.adminService.ListUsers(c.Request.Context(), 1, 30, service.UserListFilters{Search: keyword}, "email", "asc")
+ users, _, err := h.adminService.ListUsers(c.Request.Context(), 1, 30, service.UserListFilters{Search: keyword, IncludeDeleted: true}, "email", "asc")
if err != nil {
response.ErrorFrom(c, err)
return
}
- // Return simplified user list (only id and email)
+ // Return simplified user list (only id, email and deleted flag)
type SimpleUser struct {
- ID int64 `json:"id"`
- Email string `json:"email"`
+ ID int64 `json:"id"`
+ Email string `json:"email"`
+ Deleted bool `json:"deleted"`
}
result := make([]SimpleUser, len(users))
for i, u := range users {
result[i] = SimpleUser{
- ID: u.ID,
- Email: u.Email,
+ ID: u.ID,
+ Email: u.Email,
+ Deleted: u.DeletedAt != nil,
}
}
diff --git a/backend/internal/handler/admin/usage_handler_search_users_test.go b/backend/internal/handler/admin/usage_handler_search_users_test.go
new file mode 100644
index 00000000..ca435012
--- /dev/null
+++ b/backend/internal/handler/admin/usage_handler_search_users_test.go
@@ -0,0 +1,56 @@
+package admin
+
+import (
+ "context"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+// 捕获 ListUsers 入参、返回一个已删用户的 admin service 桩。
+type searchUsersAdminStub struct {
+ service.AdminService
+ gotFilters service.UserListFilters
+}
+
+func (s *searchUsersAdminStub) ListUsers(ctx context.Context, page, pageSize int, filters service.UserListFilters, sortBy, sortOrder string) ([]service.User, int64, error) {
+ s.gotFilters = filters
+ ts := time.Date(2026, 5, 28, 0, 0, 0, 0, time.UTC)
+ return []service.User{
+ {ID: 1, Email: "active@test.com"},
+ {ID: 2, Email: "deleted@test.com", DeletedAt: &ts},
+ }, 2, nil
+}
+
+func TestAdminUsageSearchUsers_IncludesDeletedAndFlags(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ stub := &searchUsersAdminStub{}
+ handler := NewUsageHandler(nil, nil, stub, nil)
+ router := gin.New()
+ router.GET("/admin/usage/search-users", handler.SearchUsers)
+
+ req := httptest.NewRequest(http.MethodGet, "/admin/usage/search-users?q=test", nil)
+ rec := httptest.NewRecorder()
+ router.ServeHTTP(rec, req)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+ require.True(t, stub.gotFilters.IncludeDeleted, "SearchUsers 必须请求 IncludeDeleted")
+
+ var resp struct {
+ Data []struct {
+ ID int64 `json:"id"`
+ Email string `json:"email"`
+ Deleted bool `json:"deleted"`
+ } `json:"data"`
+ }
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &resp))
+ require.Len(t, resp.Data, 2)
+ require.False(t, resp.Data[0].Deleted)
+ require.True(t, resp.Data[1].Deleted, "已删用户必须标记 deleted=true")
+}
diff --git a/backend/internal/handler/admin/usage_query_cache.go b/backend/internal/handler/admin/usage_query_cache.go
new file mode 100644
index 00000000..b288a95b
--- /dev/null
+++ b/backend/internal/handler/admin/usage_query_cache.go
@@ -0,0 +1,62 @@
+package admin
+
+import (
+ "context"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
+)
+
+// 与 dashboard 查询缓存同款:30s TTL 进程内缓存,仅服务 /admin/usage/stats 读路径。
+var usageStatsCache = newSnapshotCache(30 * time.Second)
+
+type usageStatsCacheKeyData struct {
+ StartTime string `json:"start_time"`
+ EndTime string `json:"end_time"`
+ UserID int64 `json:"user_id"`
+ APIKeyID int64 `json:"api_key_id"`
+ AccountID int64 `json:"account_id"`
+ GroupID int64 `json:"group_id"`
+ Model string `json:"model"`
+ BillingMode string `json:"billing_mode"`
+ RequestType *int16 `json:"request_type"`
+ Stream *bool `json:"stream"`
+ BillingType *int8 `json:"billing_type"`
+}
+
+func usageStatsCacheKey(filters usagestats.UsageLogFilters) string {
+ start := ""
+ if filters.StartTime != nil {
+ start = filters.StartTime.UTC().Format(time.RFC3339)
+ }
+ end := ""
+ if filters.EndTime != nil {
+ end = filters.EndTime.UTC().Format(time.RFC3339)
+ }
+ return mustMarshalDashboardCacheKey(usageStatsCacheKeyData{
+ StartTime: start,
+ EndTime: end,
+ UserID: filters.UserID,
+ APIKeyID: filters.APIKeyID,
+ AccountID: filters.AccountID,
+ GroupID: filters.GroupID,
+ Model: filters.Model,
+ BillingMode: filters.BillingMode,
+ RequestType: filters.RequestType,
+ Stream: filters.Stream,
+ BillingType: filters.BillingType,
+ })
+}
+
+// getStatsCached 命中则返回缓存,未命中则回源 usageService 并写缓存。
+func (h *UsageHandler) getStatsCached(ctx context.Context, filters usagestats.UsageLogFilters) (*usagestats.UsageStats, bool, error) {
+ key := usageStatsCacheKey(filters)
+ entry, hit, err := usageStatsCache.GetOrLoad(key, func() (any, error) {
+ return h.usageService.GetStatsWithFilters(ctx, filters)
+ })
+ if err != nil {
+ return nil, hit, err
+ }
+ stats, err := snapshotPayloadAs[*usagestats.UsageStats](entry.Payload)
+ return stats, hit, err
+}
diff --git a/backend/internal/handler/admin/usage_query_cache_test.go b/backend/internal/handler/admin/usage_query_cache_test.go
new file mode 100644
index 00000000..857e507a
--- /dev/null
+++ b/backend/internal/handler/admin/usage_query_cache_test.go
@@ -0,0 +1,28 @@
+package admin
+
+import (
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUsageStatsCacheKey_StableAndDistinct(t *testing.T) {
+ start := time.Date(2026, 5, 29, 0, 0, 0, 0, time.UTC)
+ end := time.Date(2026, 5, 31, 0, 0, 0, 0, time.UTC)
+ base := usagestats.UsageLogFilters{StartTime: &start, EndTime: &end, Model: "claude-3"}
+
+ k1 := usageStatsCacheKey(base)
+ k2 := usageStatsCacheKey(base)
+ require.NotEmpty(t, k1)
+ require.Equal(t, k1, k2, "same filters must produce same key")
+
+ other := base
+ other.Model = "gpt-4o"
+ require.NotEqual(t, k1, usageStatsCacheKey(other), "different model must change key")
+
+ withUser := base
+ withUser.UserID = 7
+ require.NotEqual(t, k1, usageStatsCacheKey(withUser), "different user must change key")
+}
diff --git a/backend/internal/handler/admin/user_handler.go b/backend/internal/handler/admin/user_handler.go
index 32a21692..a21fe55a 100644
--- a/backend/internal/handler/admin/user_handler.go
+++ b/backend/internal/handler/admin/user_handler.go
@@ -48,14 +48,14 @@ func NewUserHandler(
// CreateUserRequest represents admin create user request
type CreateUserRequest struct {
- Email string `json:"email" binding:"required,email"`
- Password string `json:"password" binding:"required,min=6"`
- Username string `json:"username"`
- Notes string `json:"notes"`
- Balance float64 `json:"balance"`
- Concurrency int `json:"concurrency"`
- RPMLimit int `json:"rpm_limit"`
- AllowedGroups []int64 `json:"allowed_groups"`
+ Email string `json:"email" binding:"required,email"`
+ Password string `json:"password" binding:"required,min=6"`
+ Username string `json:"username"`
+ Notes string `json:"notes"`
+ Balance *float64 `json:"balance"`
+ Concurrency int `json:"concurrency"`
+ RPMLimit int `json:"rpm_limit"`
+ AllowedGroups []int64 `json:"allowed_groups"`
}
// UpdateUserRequest represents admin update user request
@@ -195,7 +195,12 @@ func (h *UserHandler) GetByID(c *gin.Context) {
return
}
- user, err := h.adminService.GetUser(c.Request.Context(), userID)
+ var user *service.User
+ if c.Query("include_deleted") == "true" {
+ user, err = h.adminService.GetUserIncludeDeleted(c.Request.Context(), userID)
+ } else {
+ user, err = h.adminService.GetUser(c.Request.Context(), userID)
+ }
if err != nil {
response.ErrorFrom(c, err)
return
@@ -743,7 +748,7 @@ func (h *UserHandler) UpdateUserPlatformQuotas(c *gin.Context) {
if h.billingCache != nil {
for _, p := range service.AllowedQuotaPlatforms {
if err := h.billingCache.DeleteUserPlatformQuotaCache(ctx, userID, p); err != nil {
- slog.Warn("quota cache invalidation failed", "user_id", userID, "platform", p, "err", err)
+ slog.Error("ALERT: quota cache invalidation failed after UpsertForUser; limit 生效可能延迟至 sentinel TTL(最长 1h),需人工确认或重试失效", "user_id", userID, "platform", p, "err", err)
}
}
}
@@ -827,7 +832,7 @@ func (h *UserHandler) ResetUserPlatformQuotaWindow(c *gin.Context) {
if h.billingCache != nil {
if err := h.billingCache.DeleteUserPlatformQuotaCache(ctx, userID, req.Platform); err != nil {
- slog.Warn("quota cache invalidation failed", "user_id", userID, "platform", req.Platform, "err", err)
+ slog.Error("ALERT: quota cache invalidation failed after ResetExpiredWindow; 窗口重置可能延迟至 sentinel TTL(最长 1h)", "user_id", userID, "platform", req.Platform, "err", err)
}
}
diff --git a/backend/internal/handler/admin/user_handler_get_deleted_test.go b/backend/internal/handler/admin/user_handler_get_deleted_test.go
new file mode 100644
index 00000000..1b3070cd
--- /dev/null
+++ b/backend/internal/handler/admin/user_handler_get_deleted_test.go
@@ -0,0 +1,51 @@
+package admin
+
+import (
+ "context"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+type getByIDAdminStub struct {
+ service.AdminService
+}
+
+func (s *getByIDAdminStub) GetUser(_ context.Context, _ int64) (*service.User, error) {
+ return nil, service.ErrUserNotFound
+}
+
+func (s *getByIDAdminStub) GetUserIncludeDeleted(_ context.Context, id int64) (*service.User, error) {
+ return &service.User{ID: id, Email: "del@test.com"}, nil
+}
+
+func setupGetByIDRouter(svc service.AdminService) *gin.Engine {
+ gin.SetMode(gin.TestMode)
+ r := gin.New()
+ h := NewUserHandler(svc, nil, nil, nil)
+ r.GET("/admin/users/:id", h.GetByID)
+ return r
+}
+
+func TestAdminUserGetByID_IncludeDeleted(t *testing.T) {
+ svc := &getByIDAdminStub{AdminService: newStubAdminService()}
+ router := setupGetByIDRouter(svc)
+
+ t.Run("normal path returns 404 for deleted user", func(t *testing.T) {
+ w := httptest.NewRecorder()
+ req, _ := http.NewRequest(http.MethodGet, "/admin/users/7", nil)
+ router.ServeHTTP(w, req)
+ require.Equal(t, http.StatusNotFound, w.Code)
+ })
+
+ t.Run("include_deleted=true returns 200", func(t *testing.T) {
+ w := httptest.NewRecorder()
+ req, _ := http.NewRequest(http.MethodGet, "/admin/users/7?include_deleted=true", nil)
+ router.ServeHTTP(w, req)
+ require.Equal(t, http.StatusOK, w.Code)
+ })
+}
diff --git a/backend/internal/handler/auth_oauth_pending_flow_test.go b/backend/internal/handler/auth_oauth_pending_flow_test.go
index 70fb160a..2f8f4e58 100644
--- a/backend/internal/handler/auth_oauth_pending_flow_test.go
+++ b/backend/internal/handler/auth_oauth_pending_flow_test.go
@@ -2914,6 +2914,10 @@ func (r *oauthPendingFlowUserRepo) DisableTotp(ctx context.Context, userID int64
Exec(ctx)
}
+func (r *oauthPendingFlowUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ return r.GetByID(ctx, id)
+}
+
func oauthPendingFlowServiceUser(entity *dbent.User) *service.User {
if entity == nil {
return nil
diff --git a/backend/internal/handler/concurrency_error_response.go b/backend/internal/handler/concurrency_error_response.go
new file mode 100644
index 00000000..52abf735
--- /dev/null
+++ b/backend/internal/handler/concurrency_error_response.go
@@ -0,0 +1,27 @@
+package handler
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "net/http"
+)
+
+const statusClientClosedRequest = 499
+
+func concurrencyErrorResponse(err error, slotType string) (int, string, string) {
+ var concurrencyErr *ConcurrencyError
+ if errors.As(err, &concurrencyErr) {
+ if concurrencyErr.SlotType != "" {
+ slotType = concurrencyErr.SlotType
+ }
+ return http.StatusTooManyRequests, "rate_limit_error",
+ fmt.Sprintf("Concurrency limit exceeded for %s, please retry later", slotType)
+ }
+
+ if errors.Is(err, context.Canceled) {
+ return statusClientClosedRequest, "api_error", "context canceled"
+ }
+
+ return http.StatusServiceUnavailable, "api_error", "Service temporarily unavailable, please retry later"
+}
diff --git a/backend/internal/handler/concurrency_error_response_test.go b/backend/internal/handler/concurrency_error_response_test.go
new file mode 100644
index 00000000..a2e6b9ab
--- /dev/null
+++ b/backend/internal/handler/concurrency_error_response_test.go
@@ -0,0 +1,63 @@
+package handler
+
+import (
+ "context"
+ "errors"
+ "net/http"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestConcurrencyErrorResponse(t *testing.T) {
+ tests := []struct {
+ name string
+ err error
+ slotType string
+ wantStatus int
+ wantType string
+ wantMessage string
+ }{
+ {
+ name: "true concurrency timeout remains rate limit",
+ err: &ConcurrencyError{SlotType: "account", IsTimeout: true},
+ slotType: "user",
+ wantStatus: http.StatusTooManyRequests,
+ wantType: "rate_limit_error",
+ wantMessage: "Concurrency limit exceeded for account, please retry later",
+ },
+ {
+ name: "client cancellation is not classified as concurrency limit",
+ err: context.Canceled,
+ slotType: "user",
+ wantStatus: statusClientClosedRequest,
+ wantType: "api_error",
+ wantMessage: "context canceled",
+ },
+ {
+ name: "deadline exceeded is service unavailable",
+ err: context.DeadlineExceeded,
+ slotType: "user",
+ wantStatus: http.StatusServiceUnavailable,
+ wantType: "api_error",
+ wantMessage: "Service temporarily unavailable, please retry later",
+ },
+ {
+ name: "redis acquire error is service unavailable",
+ err: errors.New("redis unavailable"),
+ slotType: "user",
+ wantStatus: http.StatusServiceUnavailable,
+ wantType: "api_error",
+ wantMessage: "Service temporarily unavailable, please retry later",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ status, errType, message := concurrencyErrorResponse(tt.err, tt.slotType)
+ require.Equal(t, tt.wantStatus, status)
+ require.Equal(t, tt.wantType, errType)
+ require.Equal(t, tt.wantMessage, message)
+ })
+ }
+}
diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go
index 2c71be9d..86f98f15 100644
--- a/backend/internal/handler/dto/mappers.go
+++ b/backend/internal/handler/dto/mappers.go
@@ -30,6 +30,7 @@ func UserFromServiceShallow(u *service.User) *User {
BalanceNotifyExtraEmails: NotifyEmailEntriesFromService(u.BalanceNotifyExtraEmails),
TotalRecharged: u.TotalRecharged,
RPMLimit: u.RPMLimit,
+ DeletedAt: u.DeletedAt,
}
}
@@ -147,6 +148,7 @@ func GroupFromServiceAdmin(g *service.Group) *AdminGroup {
MCPXMLInject: g.MCPXMLInject,
DefaultMappedModel: g.DefaultMappedModel,
MessagesDispatchModelConfig: g.MessagesDispatchModelConfig,
+ ModelsListConfig: g.ModelsListConfig,
SupportedModelScopes: g.SupportedModelScopes,
AccountCount: g.AccountCount,
ActiveAccountCount: g.ActiveAccountCount,
diff --git a/backend/internal/handler/dto/mappers_deleted_user_test.go b/backend/internal/handler/dto/mappers_deleted_user_test.go
new file mode 100644
index 00000000..8ce5388e
--- /dev/null
+++ b/backend/internal/handler/dto/mappers_deleted_user_test.go
@@ -0,0 +1,20 @@
+package dto
+
+import (
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUserFromServiceShallow_MapsDeletedAt(t *testing.T) {
+ ts := time.Date(2026, 5, 28, 10, 0, 0, 0, time.UTC)
+
+ deleted := UserFromServiceShallow(&service.User{ID: 1, Email: "d@test.com", DeletedAt: &ts})
+ require.NotNil(t, deleted.DeletedAt)
+ require.Equal(t, ts, *deleted.DeletedAt)
+
+ active := UserFromServiceShallow(&service.User{ID: 2, Email: "a@test.com"})
+ require.Nil(t, active.DeletedAt, "active user must have nil DeletedAt")
+}
diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go
index d2b7fb2b..90ffc7e0 100644
--- a/backend/internal/handler/dto/settings.go
+++ b/backend/internal/handler/dto/settings.go
@@ -186,6 +186,7 @@ type SystemSettings struct {
RewriteMessageCacheControl bool `json:"rewrite_message_cache_control"`
AntigravityUserAgentVersion string `json:"antigravity_user_agent_version"`
OpenAICodexUserAgent string `json:"openai_codex_user_agent"`
+ OpenAIAllowClaudeCodeCodexPlugin bool `json:"openai_allow_claude_code_codex_plugin"`
// Web Search Emulation
WebSearchEmulationEnabled bool `json:"web_search_emulation_enabled"`
@@ -252,6 +253,9 @@ type SystemSettings struct {
// 系统全局默认平台配额(key = platform,nil/缺省 = 不限制)
DefaultPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"default_platform_quotas,omitempty"`
+
+ // 允许终端用户在用量页查看自己的失败请求
+ AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"`
}
type DefaultSubscriptionSetting struct {
@@ -316,6 +320,8 @@ type PublicSettings struct {
AffiliateEnabled bool `json:"affiliate_enabled"`
RiskControlEnabled bool `json:"risk_control_enabled"`
+
+ AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"`
}
type LoginAgreementDocument struct {
diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go
index 31828375..08dc6572 100644
--- a/backend/internal/handler/dto/types.go
+++ b/backend/internal/handler/dto/types.go
@@ -20,6 +20,7 @@ type User struct {
LastActiveAt *time.Time `json:"last_active_at,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
+ DeletedAt *time.Time `json:"deleted_at,omitempty"`
// 余额不足通知
BalanceNotifyEnabled bool `json:"balance_notify_enabled"`
@@ -138,6 +139,7 @@ type AdminGroup struct {
// OpenAI Messages 调度配置(仅 openai 平台使用)
DefaultMappedModel string `json:"default_mapped_model"`
MessagesDispatchModelConfig domain.OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config"`
+ ModelsListConfig domain.GroupModelsListConfig `json:"models_list_config"`
// 支持的模型系列(仅 antigravity 平台使用)
SupportedModelScopes []string `json:"supported_model_scopes"`
diff --git a/backend/internal/handler/endpoint.go b/backend/internal/handler/endpoint.go
index db29618a..0d6f4b3c 100644
--- a/backend/internal/handler/endpoint.go
+++ b/backend/internal/handler/endpoint.go
@@ -17,6 +17,7 @@ import (
const (
EndpointMessages = "/v1/messages"
EndpointChatCompletions = "/v1/chat/completions"
+ EndpointEmbeddings = "/v1/embeddings"
EndpointResponses = "/v1/responses"
EndpointImagesGenerations = "/v1/images/generations"
EndpointImagesEdits = "/v1/images/edits"
@@ -42,6 +43,8 @@ const (
func NormalizeInboundEndpoint(path string) string {
path = strings.TrimSpace(path)
switch {
+ case strings.Contains(path, EndpointEmbeddings):
+ return EndpointEmbeddings
case strings.Contains(path, EndpointChatCompletions):
return EndpointChatCompletions
case strings.Contains(path, EndpointMessages):
@@ -75,7 +78,7 @@ func DeriveUpstreamEndpoint(inbound, rawRequestPath, platform string) string {
switch platform {
case service.PlatformOpenAI:
- if inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits {
+ if inbound == EndpointEmbeddings || inbound == EndpointImagesGenerations || inbound == EndpointImagesEdits {
return inbound
}
// OpenAI forwards everything to the Responses API.
diff --git a/backend/internal/handler/endpoint_test.go b/backend/internal/handler/endpoint_test.go
index 369c5fa7..42b6d6e7 100644
--- a/backend/internal/handler/endpoint_test.go
+++ b/backend/internal/handler/endpoint_test.go
@@ -24,6 +24,7 @@ func TestNormalizeInboundEndpoint(t *testing.T) {
// Direct canonical paths.
{"/v1/messages", EndpointMessages},
{"/v1/chat/completions", EndpointChatCompletions},
+ {"/v1/embeddings", EndpointEmbeddings},
{"/v1/responses", EndpointResponses},
{"/v1/images/generations", EndpointImagesGenerations},
{"/v1/images/edits", EndpointImagesEdits},
@@ -77,6 +78,7 @@ func TestDeriveUpstreamEndpoint(t *testing.T) {
{"openai responses nested", EndpointResponses, "/openai/v1/responses/compact/detail", service.PlatformOpenAI, "/v1/responses/compact/detail"},
{"openai from messages", EndpointMessages, "/v1/messages", service.PlatformOpenAI, EndpointResponses},
{"openai from completions", EndpointChatCompletions, "/v1/chat/completions", service.PlatformOpenAI, EndpointResponses},
+ {"openai embeddings", EndpointEmbeddings, "/v1/embeddings", service.PlatformOpenAI, EndpointEmbeddings},
{"openai image generations", EndpointImagesGenerations, "/v1/images/generations", service.PlatformOpenAI, EndpointImagesGenerations},
{"openai image edits", EndpointImagesEdits, "/openai/v1/images/edits", service.PlatformOpenAI, EndpointImagesEdits},
diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go
index 87a935fd..8853bffb 100644
--- a/backend/internal/handler/gateway_handler.go
+++ b/backend/internal/handler/gateway_handler.go
@@ -154,7 +154,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
setOpsRequestContext(c, "", false)
- parsedReq, err := service.ParseGatewayRequest(body, domain.PlatformAnthropic)
+ bodyRef := service.NewRequestBodyRef(body)
+ parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic)
if err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
@@ -440,7 +441,17 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 记录 Forward 前已写入字节数,Forward 后若增加则说明 SSE 内容已发,禁止 failover
writerSizeBeforeForward := c.Writer.Size()
if account.Platform == service.PlatformAntigravity {
- result, err = h.antigravityGatewayService.ForwardGemini(requestCtx, c, account, reqModel, "generateContent", reqStream, body, hasBoundSession)
+ result, err = h.antigravityGatewayService.ForwardGemini(
+ requestCtx,
+ c,
+ account,
+ reqModel,
+ "generateContent",
+ reqStream,
+ body,
+ hasBoundSession,
+ service.WithForwardGeminiSession(derefGroupID(apiKey.GroupID), sessionKey),
+ )
} else {
result, err = h.geminiCompatService.Forward(requestCtx, c, account, body)
}
@@ -509,11 +520,12 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
}
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
+ // ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
+ forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
- ParsedRequest: parsedReq,
QuotaPlatform: quotaPlatform,
APIKey: apiKey,
User: apiKey.User,
@@ -524,7 +536,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
UserAgent: userAgent,
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
- ForceCacheBilling: fs.ForceCacheBilling,
+ ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
}); err != nil {
@@ -562,6 +574,12 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
retryWithFallback := false
for {
+ attemptParsedReq, err := parsedReq.CloneForBody(body)
+ if err != nil {
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
+ return
+ }
+
// 选择支持该模型的账号
reqLog.Info("sticky.selecting_account",
zap.String("session_key", sessionKey),
@@ -693,7 +711,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// ===== 用户消息串行队列 START =====
var queueRelease func()
- umqMode := h.getUserMsgQueueMode(account, parsedReq)
+ umqMode := h.getUserMsgQueueMode(account, attemptParsedReq)
switch umqMode {
case config.UMQModeSerialize:
@@ -740,20 +758,26 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 用 wrapReleaseOnDone 确保 context 取消时自动释放(仅 serialize 模式有 queueRelease)
queueRelease = wrapReleaseOnDone(c.Request.Context(), queueRelease)
// 注入回调到 ParsedRequest:使用外层 wrapper 以便提前清理 AfterFunc
- parsedReq.OnUpstreamAccepted = queueRelease
+ attemptParsedReq.OnUpstreamAccepted = queueRelease
// ===== 用户消息串行队列 END =====
- // 应用渠道模型映射到请求
+ // 渠道模型映射只作用于本次账号尝试,避免 failover 后污染原始 ParsedRequest。
if channelMapping.Mapped {
- parsedReq.Model = channelMapping.MappedModel
- parsedReq.Body = h.gatewayService.ReplaceModelInBody(parsedReq.Body, channelMapping.MappedModel)
+ attemptParsedReq.Model = channelMapping.MappedModel
+ if err := attemptParsedReq.ReplaceBody(h.gatewayService.ReplaceModelInBody(attemptParsedReq.Body.Bytes(), channelMapping.MappedModel)); err != nil {
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
+ return
+ }
}
// Bedrock CC 兼容:渠道模型映射后,清理 Anthropic API 专有字段、注入 Bedrock 必需字段
- parsedReq.Body = h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), parsedReq.Body, parsedReq.Model, account, apiKey.GroupID)
- body = parsedReq.Body
+ if err := attemptParsedReq.ReplaceBody(h.gatewayService.ApplyBedrockCCCompat(c.Request.Context(), attemptParsedReq.Body.Bytes(), attemptParsedReq.Model, account, apiKey.GroupID)); err != nil {
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
+ return
+ }
+ attemptBody := attemptParsedReq.Body.Bytes()
// 转发请求 - 根据账号平台分流
- c.Set("parsed_request", parsedReq)
+ c.Set("parsed_request", attemptParsedReq)
var result *service.ForwardResult
requestCtx := c.Request.Context()
if fs.SwitchCount > 0 {
@@ -762,9 +786,9 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 记录 Forward 前已写入字节数,Forward 后若增加则说明 SSE 内容已发,禁止 failover
writerSizeBeforeForward := c.Writer.Size()
if account.Platform == service.PlatformAntigravity && account.Type != service.AccountTypeAPIKey {
- result, err = h.antigravityGatewayService.Forward(requestCtx, c, account, body, hasBoundSession)
+ result, err = h.antigravityGatewayService.Forward(requestCtx, c, account, attemptBody, hasBoundSession)
} else {
- result, err = h.gatewayService.Forward(requestCtx, c, account, parsedReq)
+ result, err = h.gatewayService.Forward(requestCtx, c, account, attemptParsedReq)
}
// 兜底释放串行锁(正常情况已通过回调提前释放)
@@ -772,7 +796,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
queueRelease()
}
// 清理回调引用,防止 failover 重试时旧回调被错误调用
- parsedReq.OnUpstreamAccepted = nil
+ attemptParsedReq.OnUpstreamAccepted = nil
if accountReleaseFunc != nil {
accountReleaseFunc()
@@ -895,20 +919,22 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// 捕获请求信息(用于异步记录,避免在 goroutine 中访问 gin.Context)
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
- requestPayloadHash := service.HashUsageRequestPayload(body)
+ // Forward 内部可能继续改写 body,usage 去重指纹必须使用最终上游接受的当前 body。
+ requestPayloadHash := service.HashUsageRequestPayload(attemptParsedReq.Body.Bytes())
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
if result.ReasoningEffort == nil {
- result.ReasoningEffort = service.NormalizeClaudeOutputEffort(parsedReq.OutputEffort)
+ result.ReasoningEffort = service.NormalizeClaudeOutputEffort(attemptParsedReq.OutputEffort)
}
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
+ // ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
+ forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), currentAPIKey)
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
- ParsedRequest: parsedReq,
QuotaPlatform: quotaPlatform,
APIKey: currentAPIKey,
User: currentAPIKey.User,
@@ -919,7 +945,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
UserAgent: userAgent,
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
- ForceCacheBilling: fs.ForceCacheBilling,
+ ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
}); err != nil {
@@ -961,22 +987,14 @@ func (h *GatewayHandler) Models(c *gin.Context) {
// Get available models from account configurations for the selected group platform.
availableModels := h.gatewayService.GetAvailableModels(c.Request.Context(), groupID, platform)
+ if apiKey != nil && apiKey.Group != nil && apiKey.Group.CustomModelsListEnabled() {
+ availableModels = filterModelsByCustomList(availableModels, defaultModelIDsForPlatform(platform), apiKey.Group.ModelsListConfig.Models)
+ writeCustomModelsList(c, platform, availableModels)
+ return
+ }
if len(availableModels) > 0 {
- // Build model list from whitelist
- models := make([]claude.Model, 0, len(availableModels))
- for _, modelID := range availableModels {
- models = append(models, claude.Model{
- ID: modelID,
- Type: "model",
- DisplayName: modelID,
- CreatedAt: "2024-01-01T00:00:00Z",
- })
- }
- c.JSON(http.StatusOK, gin.H{
- "object": "list",
- "data": models,
- })
+ writeModelsList(c, availableModels)
return
}
@@ -1003,6 +1021,134 @@ func (h *GatewayHandler) Models(c *gin.Context) {
})
}
+func writeModelsList(c *gin.Context, modelIDs []string) {
+ models := make([]claude.Model, 0, len(modelIDs))
+ for _, modelID := range modelIDs {
+ models = append(models, claude.Model{
+ ID: modelID,
+ Type: "model",
+ DisplayName: modelID,
+ CreatedAt: "2024-01-01T00:00:00Z",
+ })
+ }
+ c.JSON(http.StatusOK, gin.H{
+ "object": "list",
+ "data": models,
+ })
+}
+
+func writeCustomModelsList(c *gin.Context, platform string, modelIDs []string) {
+ if platform == service.PlatformOpenAI {
+ writeOpenAIModelsList(c, modelIDs)
+ return
+ }
+ writeModelsList(c, modelIDs)
+}
+
+func writeOpenAIModelsList(c *gin.Context, modelIDs []string) {
+ defaultsByID := make(map[string]openai.Model, len(openai.DefaultModels))
+ for _, model := range openai.DefaultModels {
+ defaultsByID[model.ID] = model
+ }
+
+ models := make([]openai.Model, 0, len(modelIDs))
+ for _, modelID := range modelIDs {
+ if model, ok := defaultsByID[modelID]; ok {
+ models = append(models, model)
+ continue
+ }
+ models = append(models, openai.Model{
+ ID: modelID,
+ Object: "model",
+ Created: 1704067200,
+ OwnedBy: "openai",
+ Type: "model",
+ DisplayName: modelID,
+ })
+ }
+ c.JSON(http.StatusOK, gin.H{
+ "object": "list",
+ "data": models,
+ })
+}
+
+func filterModelsByCustomList(availableModels, fallbackModels, selectedModels []string) []string {
+ if len(selectedModels) == 0 {
+ return availableModels
+ }
+ source := availableModels
+ if len(source) == 0 {
+ source = fallbackModels
+ }
+ if len(source) == 0 {
+ return nil
+ }
+
+ allowed := make([]string, 0, len(source))
+ for _, model := range source {
+ model = strings.TrimSpace(model)
+ if model != "" {
+ allowed = append(allowed, model)
+ }
+ }
+
+ seen := make(map[string]struct{}, len(selectedModels))
+ filtered := make([]string, 0, len(selectedModels))
+ for _, model := range selectedModels {
+ model = strings.TrimSpace(model)
+ if model == "" {
+ continue
+ }
+ if !customModelsListAllowsModel(allowed, model) {
+ continue
+ }
+ if _, ok := seen[model]; ok {
+ continue
+ }
+ seen[model] = struct{}{}
+ filtered = append(filtered, model)
+ }
+ return filtered
+}
+
+func customModelsListAllowsModel(availablePatterns []string, model string) bool {
+ for _, pattern := range availablePatterns {
+ if pattern == model {
+ return true
+ }
+ if strings.HasSuffix(pattern, "*") && strings.HasPrefix(model, strings.TrimSuffix(pattern, "*")) {
+ return true
+ }
+ }
+ return false
+}
+
+func defaultModelIDsForPlatform(platform string) []string {
+ switch platform {
+ case service.PlatformOpenAI:
+ return openai.DefaultModelIDs()
+ case service.PlatformGemini:
+ ids := make([]string, 0, len(geminicli.DefaultModels))
+ for _, model := range geminicli.DefaultModels {
+ ids = append(ids, model.ID)
+ }
+ return ids
+ case service.PlatformAntigravity:
+ models := antigravity.DefaultModels()
+ ids := make([]string, 0, len(models))
+ for _, model := range models {
+ ids = append(ids, model.ID)
+ }
+ return ids
+ default:
+ ids := make([]string, 0, len(claude.DefaultModels))
+ for _, model := range claude.DefaultModels {
+ ids = append(ids, model.ID)
+ }
+ return ids
+ }
+}
+
// AntigravityModels 返回 Antigravity 支持的全部模型
// GET /antigravity/models
func (h *GatewayHandler) AntigravityModels(c *gin.Context) {
@@ -1351,10 +1497,10 @@ func (h *GatewayHandler) calculateSubscriptionRemaining(group *service.Group, su
return min
}
-// handleConcurrencyError handles concurrency-related errors with proper 429 response
+// handleConcurrencyError handles concurrency-related acquire errors.
func (h *GatewayHandler) handleConcurrencyError(c *gin.Context, err error, slotType string, streamStarted bool) {
- h.handleStreamingAwareError(c, http.StatusTooManyRequests, "rate_limit_error",
- fmt.Sprintf("Concurrency limit exceeded for %s, please retry later", slotType), streamStarted)
+ status, errType, message := concurrencyErrorResponse(err, slotType)
+ h.handleStreamingAwareError(c, status, errType, message, streamStarted)
}
func (h *GatewayHandler) handleFailoverExhausted(c *gin.Context, failoverErr *service.UpstreamFailoverError, platform string, streamStarted bool) {
@@ -1563,7 +1709,8 @@ func (h *GatewayHandler) CountTokens(c *gin.Context) {
setOpsRequestContext(c, "", false)
- parsedReq, err := service.ParseGatewayRequest(body, domain.PlatformAnthropic)
+ bodyRef := service.NewRequestBodyRef(body)
+ parsedReq, err := service.ParseGatewayRequest(bodyRef, domain.PlatformAnthropic)
if err != nil {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
return
@@ -1936,10 +2083,11 @@ func (h *GatewayHandler) maybeLogCompatibilityFallbackMetrics(reqLog *zap.Logger
)
}
-func (h *GatewayHandler) submitUsageRecordTask(task service.UsageRecordTask) {
+func (h *GatewayHandler) submitUsageRecordTask(parent context.Context, task service.UsageRecordTask) {
if task == nil {
return
}
+ task = wrapUsageRecordTaskContext(parent, task)
if h.usageRecordWorkerPool != nil {
h.usageRecordWorkerPool.Submit(task)
return
diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go
index acbdc261..719700aa 100644
--- a/backend/internal/handler/gateway_handler_chat_completions.go
+++ b/backend/internal/handler/gateway_handler_chat_completions.go
@@ -151,9 +151,10 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
}
// Parse request for session hash
- parsedReq, _ := service.ParseGatewayRequest(body, "chat_completions")
+ bodyRef := service.NewRequestBodyRef(body)
+ parsedReq, _ := service.ParseGatewayRequest(bodyRef, "chat_completions")
if parsedReq == nil {
- parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: body}
+ parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: bodyRef}
}
parsedReq.SessionContext = &service.SessionContext{
ClientIP: ip.GetClientIP(c),
@@ -292,7 +293,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
QuotaPlatform: quotaPlatform,
diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go
index 6a083f31..49f80d19 100644
--- a/backend/internal/handler/gateway_handler_responses.go
+++ b/backend/internal/handler/gateway_handler_responses.go
@@ -156,9 +156,10 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
}
// Parse request for session hash
- parsedReq, _ := service.ParseGatewayRequest(body, "responses")
+ bodyRef := service.NewRequestBodyRef(body)
+ parsedReq, _ := service.ParseGatewayRequest(bodyRef, "responses")
if parsedReq == nil {
- parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: body}
+ parsedReq = &service.ParsedRequest{Model: reqModel, Stream: reqStream, Body: bodyRef}
}
parsedReq.SessionContext = &service.SessionContext{
ClientIP: ip.GetClientIP(c),
@@ -267,7 +268,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
QuotaPlatform: quotaPlatform,
diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go
index 09e6c09b..4b6a47eb 100644
--- a/backend/internal/handler/gateway_helper.go
+++ b/backend/internal/handler/gateway_helper.go
@@ -18,18 +18,12 @@ import (
// claudeCodeValidator is a singleton validator for Claude Code client detection
var claudeCodeValidator = service.NewClaudeCodeValidator()
-const claudeCodeParsedRequestContextKey = "claude_code_parsed_request"
-
// SetClaudeCodeClientContext 检查请求是否来自 Claude Code 客户端,并设置到 context 中
// 返回更新后的 context
func SetClaudeCodeClientContext(c *gin.Context, body []byte, parsedReq *service.ParsedRequest) {
if c == nil || c.Request == nil {
return
}
- if parsedReq != nil {
- c.Set(claudeCodeParsedRequestContextKey, parsedReq)
- }
-
ua := c.GetHeader("User-Agent")
// Fast path:非 Claude CLI UA 直接判定 false,避免热路径二次 JSON 反序列化。
if !claudeCodeValidator.ValidateUserAgent(ua) {
@@ -45,9 +39,6 @@ func SetClaudeCodeClientContext(c *gin.Context, body []byte, parsedReq *service.
} else {
// 仅在确认为 Claude CLI 且 messages 路径时再做 body 解析。
bodyMap := claudeCodeBodyMapFromParsedRequest(parsedReq)
- if bodyMap == nil {
- bodyMap = claudeCodeBodyMapFromContextCache(c)
- }
if bodyMap == nil && len(body) > 0 {
_ = json.Unmarshal(body, &bodyMap)
}
@@ -74,8 +65,12 @@ func claudeCodeBodyMapFromParsedRequest(parsedReq *service.ParsedRequest) map[st
bodyMap := map[string]any{
"model": parsedReq.Model,
}
- if parsedReq.System != nil || parsedReq.HasSystem {
- bodyMap["system"] = parsedReq.System
+ if parsedReq.HasSystem {
+ if system, ok := parsedReq.SystemValue(); ok {
+ bodyMap["system"] = system
+ } else {
+ bodyMap["system"] = nil
+ }
}
if parsedReq.MetadataUserID != "" {
bodyMap["metadata"] = map[string]any{"user_id": parsedReq.MetadataUserID}
@@ -83,26 +78,6 @@ func claudeCodeBodyMapFromParsedRequest(parsedReq *service.ParsedRequest) map[st
return bodyMap
}
-func claudeCodeBodyMapFromContextCache(c *gin.Context) map[string]any {
- if c == nil {
- return nil
- }
- if cached, ok := c.Get(service.OpenAIParsedRequestBodyKey); ok {
- if bodyMap, ok := cached.(map[string]any); ok {
- return bodyMap
- }
- }
- if cached, ok := c.Get(claudeCodeParsedRequestContextKey); ok {
- switch v := cached.(type) {
- case *service.ParsedRequest:
- return claudeCodeBodyMapFromParsedRequest(v)
- case service.ParsedRequest:
- return claudeCodeBodyMapFromParsedRequest(&v)
- }
- }
- return nil
-}
-
// 并发槽位等待相关常量
//
// 性能优化说明:
@@ -336,6 +311,9 @@ func (h *ConcurrencyHelper) waitForSlotWithPingTimeout(c *gin.Context, slotType
for {
select {
case <-ctx.Done():
+ if parentErr := c.Request.Context().Err(); parentErr != nil {
+ return nil, parentErr
+ }
return nil, &ConcurrencyError{
SlotType: slotType,
IsTimeout: true,
diff --git a/backend/internal/handler/gateway_helper_hotpath_test.go b/backend/internal/handler/gateway_helper_hotpath_test.go
index 4a677199..a6b6a429 100644
--- a/backend/internal/handler/gateway_helper_hotpath_test.go
+++ b/backend/internal/handler/gateway_helper_hotpath_test.go
@@ -177,7 +177,7 @@ func TestSetClaudeCodeClientContext_FastPathAndStrictPath(t *testing.T) {
})
}
-func TestSetClaudeCodeClientContext_ReuseParsedRequestAndContextCache(t *testing.T) {
+func TestSetClaudeCodeClientContext_ReuseParsedRequest(t *testing.T) {
t.Run("reuse parsed request without body unmarshal", func(t *testing.T) {
c, _ := newHelperTestContext(http.MethodPost, "/v1/messages")
c.Request.Header.Set("User-Agent", "claude-cli/1.0.1")
@@ -185,36 +185,13 @@ func TestSetClaudeCodeClientContext_ReuseParsedRequestAndContextCache(t *testing
c.Request.Header.Set("anthropic-beta", "message-batches-2024-09-24")
c.Request.Header.Set("anthropic-version", "2023-06-01")
- parsedReq := &service.ParsedRequest{
- Model: "claude-3-5-sonnet-20241022",
- System: []any{
- map[string]any{"text": "You are Claude Code, Anthropic's official CLI for Claude."},
- },
- MetadataUserID: "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
- }
+ parsedReq, err := service.ParseGatewayRequest(service.NewRequestBodyRef(validClaudeCodeBodyJSON()), "")
+ require.NoError(t, err)
// body 非法 JSON,如果函数复用 parsedReq 成功则仍应判定为 Claude Code。
SetClaudeCodeClientContext(c, []byte(`{invalid`), parsedReq)
require.True(t, service.IsClaudeCodeClient(c.Request.Context()))
})
-
- t.Run("reuse context cache without body unmarshal", func(t *testing.T) {
- c, _ := newHelperTestContext(http.MethodPost, "/v1/messages")
- c.Request.Header.Set("User-Agent", "claude-cli/1.0.1")
- c.Request.Header.Set("X-App", "claude-code")
- c.Request.Header.Set("anthropic-beta", "message-batches-2024-09-24")
- c.Request.Header.Set("anthropic-version", "2023-06-01")
- c.Set(service.OpenAIParsedRequestBodyKey, map[string]any{
- "model": "claude-3-5-sonnet-20241022",
- "system": []any{
- map[string]any{"text": "You are Claude Code, Anthropic's official CLI for Claude."},
- },
- "metadata": map[string]any{"user_id": "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa"},
- })
-
- SetClaudeCodeClientContext(c, []byte(`{invalid`), nil)
- require.True(t, service.IsClaudeCodeClient(c.Request.Context()))
- })
}
func TestWaitForSlotWithPingTimeout_AccountAndUserAcquire(t *testing.T) {
@@ -280,6 +257,25 @@ func TestWaitForSlotWithPingTimeout_TimeoutAndStreamPing(t *testing.T) {
})
}
+func TestWaitForSlotWithPingTimeout_ParentContextCanceled(t *testing.T) {
+ cache := &helperConcurrencyCacheStub{
+ accountSeq: []bool{false},
+ }
+ concurrency := service.NewConcurrencyService(cache)
+ helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond)
+ c, _ := newHelperTestContext(http.MethodPost, "/v1/messages")
+ reqCtx, cancel := context.WithCancel(c.Request.Context())
+ c.Request = c.Request.WithContext(reqCtx)
+ cancel()
+
+ streamStarted := false
+ release, err := helper.waitForSlotWithPingTimeout(c, "account", 101, 2, time.Second, false, &streamStarted, true)
+ require.Nil(t, release)
+ require.ErrorIs(t, err, context.Canceled)
+ var cErr *ConcurrencyError
+ require.False(t, errors.As(err, &cErr))
+}
+
func TestWaitForSlotWithPingTimeout_AcquireError(t *testing.T) {
errCache := &helperConcurrencyCacheStubWithError{
err: errors.New("redis unavailable"),
diff --git a/backend/internal/handler/gateway_models_test.go b/backend/internal/handler/gateway_models_test.go
index 78b07a1a..c5238f2a 100644
--- a/backend/internal/handler/gateway_models_test.go
+++ b/backend/internal/handler/gateway_models_test.go
@@ -25,7 +25,11 @@ type gatewayModelsResponseForTest struct {
}
type gatewayModelItemForTest struct {
- ID string `json:"id"`
+ ID string `json:"id"`
+ Object string `json:"object"`
+ Created int64 `json:"created"`
+ OwnedBy string `json:"owned_by"`
+ CreatedAt string `json:"created_at"`
}
func (s *gatewayModelsAccountRepoStub) ListSchedulableByGroupID(ctx context.Context, groupID int64) ([]service.Account, error) {
@@ -127,6 +131,267 @@ func TestGatewayModels_GeminiGroupFiltersMappedModelsByPlatform(t *testing.T) {
require.Equal(t, []string{"gemini-2.5-flash"}, modelIDsForTest(got.Data))
}
+func TestGatewayModels_CustomModelsListDisabledKeepsOriginalModels(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(22)
+ h := newGatewayModelsHandlerForTest(
+ &gatewayModelsAccountRepoStub{
+ byGroup: map[int64][]service.Account{
+ groupID: {
+ {
+ ID: 1,
+ Platform: service.PlatformOpenAI,
+ Credentials: map[string]any{
+ "model_mapping": map[string]any{
+ "gpt-5.5": "gpt-5.5",
+ "gpt-5.4": "gpt-5.4",
+ },
+ },
+ },
+ },
+ },
+ },
+ )
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
+ c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
+ Group: &service.Group{
+ ID: groupID,
+ Platform: service.PlatformOpenAI,
+ ModelsListConfig: service.GroupModelsListConfig{
+ Enabled: false,
+ Models: []string{"gpt-5.5"},
+ },
+ },
+ })
+
+ h.Models(c)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+
+ var got gatewayModelsResponseForTest
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
+ require.Equal(t, []string{"gpt-5.4", "gpt-5.5"}, modelIDsForTest(got.Data))
+}
+
+func TestGatewayModels_CustomModelsListFiltersAndOrdersMappedModels(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(23)
+ h := newGatewayModelsHandlerForTest(
+ &gatewayModelsAccountRepoStub{
+ byGroup: map[int64][]service.Account{
+ groupID: {
+ {
+ ID: 1,
+ Platform: service.PlatformOpenAI,
+ Credentials: map[string]any{
+ "model_mapping": map[string]any{
+ "gpt-5.4": "gpt-5.4",
+ "gpt-5.5": "gpt-5.5",
+ "legacy-gpt-2024": "legacy-gpt-2024",
+ },
+ },
+ },
+ },
+ },
+ },
+ )
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
+ c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
+ Group: &service.Group{
+ ID: groupID,
+ Platform: service.PlatformOpenAI,
+ ModelsListConfig: service.GroupModelsListConfig{
+ Enabled: true,
+ Models: []string{"gpt-5.5", "missing-model", "gpt-5.4"},
+ },
+ },
+ })
+
+ h.Models(c)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+
+ var got gatewayModelsResponseForTest
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
+ require.Equal(t, []string{"gpt-5.5", "gpt-5.4"}, modelIDsForTest(got.Data))
+}
+
+func TestGatewayModels_CustomModelsListKeepsConcreteModelAllowedByWildcardMapping(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(26)
+ h := newGatewayModelsHandlerForTest(
+ &gatewayModelsAccountRepoStub{
+ byGroup: map[int64][]service.Account{
+ groupID: {
+ {
+ ID: 1,
+ Platform: service.PlatformAnthropic,
+ Credentials: map[string]any{
+ "model_mapping": map[string]any{
+ "claude-*": "claude-sonnet-4-6",
+ },
+ },
+ },
+ },
+ },
+ },
+ )
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
+ c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
+ Group: &service.Group{
+ ID: groupID,
+ Platform: service.PlatformAnthropic,
+ ModelsListConfig: service.GroupModelsListConfig{
+ Enabled: true,
+ Models: []string{"claude-sonnet-4-6"},
+ },
+ },
+ })
+
+ h.Models(c)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+
+ var got gatewayModelsResponseForTest
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
+ require.Equal(t, []string{"claude-sonnet-4-6"}, modelIDsForTest(got.Data))
+}
+
+func TestGatewayModels_CustomModelsListCanReturnEmptyWhenSelectionsUnavailable(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(24)
+ h := newGatewayModelsHandlerForTest(
+ &gatewayModelsAccountRepoStub{
+ byGroup: map[int64][]service.Account{
+ groupID: {
+ {
+ ID: 1,
+ Platform: service.PlatformOpenAI,
+ Credentials: map[string]any{
+ "model_mapping": map[string]any{
+ "gpt-5.4": "gpt-5.4",
+ },
+ },
+ },
+ },
+ },
+ },
+ )
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
+ c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
+ Group: &service.Group{
+ ID: groupID,
+ Platform: service.PlatformOpenAI,
+ ModelsListConfig: service.GroupModelsListConfig{
+ Enabled: true,
+ Models: []string{"gpt-5.5"},
+ },
+ },
+ })
+
+ h.Models(c)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+
+ var got gatewayModelsResponseForTest
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
+ require.Empty(t, modelIDsForTest(got.Data))
+}
+
+func TestGatewayModels_CustomModelsListFiltersDefaultFallbackModels(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(25)
+ h := newGatewayModelsHandlerForTest(
+ &gatewayModelsAccountRepoStub{
+ byGroup: map[int64][]service.Account{
+ groupID: {
+ {ID: 1, Platform: service.PlatformOpenAI},
+ },
+ },
+ },
+ )
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
+ c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
+ Group: &service.Group{
+ ID: groupID,
+ Platform: service.PlatformOpenAI,
+ ModelsListConfig: service.GroupModelsListConfig{
+ Enabled: true,
+ Models: []string{"gpt-5.5", "legacy-gpt-2024", "gpt-5.4"},
+ },
+ },
+ })
+
+ h.Models(c)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+
+ var got gatewayModelsResponseForTest
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
+ require.Equal(t, []string{"gpt-5.5", "gpt-5.4"}, modelIDsForTest(got.Data))
+}
+
+func TestGatewayModels_OpenAICustomModelsListKeepsOpenAIResponseShapeForDefaultFallback(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(27)
+ h := newGatewayModelsHandlerForTest(
+ &gatewayModelsAccountRepoStub{
+ byGroup: map[int64][]service.Account{
+ groupID: {
+ {ID: 1, Platform: service.PlatformOpenAI},
+ },
+ },
+ },
+ )
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodGet, "/v1/models", nil)
+ c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
+ Group: &service.Group{
+ ID: groupID,
+ Platform: service.PlatformOpenAI,
+ ModelsListConfig: service.GroupModelsListConfig{
+ Enabled: true,
+ Models: []string{"gpt-5.5", "gpt-5.4"},
+ },
+ },
+ })
+
+ h.Models(c)
+
+ require.Equal(t, http.StatusOK, rec.Code)
+
+ var got gatewayModelsResponseForTest
+ require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
+ require.Equal(t, []string{"gpt-5.5", "gpt-5.4"}, modelIDsForTest(got.Data))
+ require.Equal(t, "model", got.Data[0].Object)
+ require.NotZero(t, got.Data[0].Created)
+ require.Equal(t, "openai", got.Data[0].OwnedBy)
+ require.Empty(t, got.Data[0].CreatedAt)
+}
+
func modelIDsForTest(models []gatewayModelItemForTest) []string {
ids := make([]string, 0, len(models))
for _, model := range models {
diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go
index 27ea4404..b7781eec 100644
--- a/backend/internal/handler/gemini_v1beta_handler.go
+++ b/backend/internal/handler/gemini_v1beta_handler.go
@@ -262,7 +262,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
sessionHash := extractGeminiCLISessionHash(c, body)
if sessionHash == "" {
// Fallback: 使用通用的会话哈希生成逻辑(适用于其他客户端)
- parsedReq, _ := service.ParseGatewayRequest(body, domain.PlatformGemini)
+ parsedReq, _ := service.ParseGatewayRequest(service.NewRequestBodyRef(body), domain.PlatformGemini)
if parsedReq != nil {
parsedReq.SessionContext = &service.SessionContext{
ClientIP: ip.GetClientIP(c),
@@ -477,8 +477,19 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
if fs.SwitchCount > 0 {
requestCtx = service.WithAccountSwitchCount(requestCtx, fs.SwitchCount, h.metadataBridgeEnabled())
}
+ sessionGroupID := derefGroupID(apiKey.GroupID)
if account.Platform == service.PlatformAntigravity && account.Type != service.AccountTypeAPIKey {
- result, err = h.antigravityGatewayService.ForwardGemini(requestCtx, c, account, modelName, action, stream, body, hasBoundSession)
+ result, err = h.antigravityGatewayService.ForwardGemini(
+ requestCtx,
+ c,
+ account,
+ modelName,
+ action,
+ stream,
+ body,
+ hasBoundSession,
+ service.WithForwardGeminiSession(sessionGroupID, sessionKey),
+ )
} else {
result, err = h.geminiCompatService.ForwardNative(requestCtx, c, account, modelName, action, stream, body)
}
@@ -527,8 +538,10 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
requestPayloadHash := service.HashUsageRequestPayload(body)
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
+ // ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
+ forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsageWithLongContext(ctx, &service.RecordUsageLongContextInput{
Result: result,
QuotaPlatform: quotaPlatform,
@@ -543,7 +556,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
RequestPayloadHash: requestPayloadHash,
LongContextThreshold: 200000, // Gemini 200K 阈值
LongContextMultiplier: 2.0, // 超出部分双倍计费
- ForceCacheBilling: fs.ForceCacheBilling,
+ ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
}); err != nil {
diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go
index 17f0d47e..d5865620 100644
--- a/backend/internal/handler/openai_chat_completions.go
+++ b/backend/internal/handler/openai_chat_completions.go
@@ -127,7 +127,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
for {
reqLog.Debug("openai_chat_completions.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs)))
- selection, scheduleDecision, err := h.gatewayService.SelectAccountWithScheduler(
+ selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(),
apiKey.GroupID,
"",
@@ -135,6 +135,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
reqModel,
failedAccountIDs,
service.OpenAIUpstreamTransportAny,
+ service.OpenAIEndpointCapabilityChatCompletions,
false,
)
if err != nil {
@@ -273,7 +274,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveRawCCUpstreamEndpoint(c, account)
- h.submitOpenAIUsageRecordTask(result, func(ctx context.Context) {
+ h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
Result: result,
APIKey: apiKey,
diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go
new file mode 100644
index 00000000..20f90735
--- /dev/null
+++ b/backend/internal/handler/openai_embeddings.go
@@ -0,0 +1,247 @@
+package handler
+
+import (
+ "context"
+ "errors"
+ "net/http"
+ "strconv"
+ "strings"
+ "time"
+
+ pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/ip"
+ "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"
+ "github.com/tidwall/gjson"
+ "go.uber.org/zap"
+)
+
+// Embeddings handles the OpenAI-compatible Embeddings API.
+// POST /v1/embeddings
+func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
+ streamStarted := false
+ requestStart := time.Now()
+
+ apiKey, ok := middleware2.GetAPIKeyFromContext(c)
+ if !ok {
+ h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key")
+ return
+ }
+
+ subject, ok := middleware2.GetAuthSubjectFromContext(c)
+ if !ok {
+ h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found")
+ return
+ }
+ reqLog := requestLogger(
+ c,
+ "handler.openai_gateway.embeddings",
+ zap.Int64("user_id", subject.UserID),
+ zap.Int64("api_key_id", apiKey.ID),
+ zap.Any("group_id", apiKey.GroupID),
+ )
+ if !h.ensureResponsesDependencies(c, reqLog) {
+ return
+ }
+
+ body, err := pkghttputil.ReadRequestBodyWithPrealloc(c.Request)
+ if err != nil {
+ if maxErr, ok := extractMaxBytesError(err); ok {
+ h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit))
+ return
+ }
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body")
+ return
+ }
+ if len(body) == 0 {
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Request body is empty")
+ return
+ }
+ if !gjson.ValidBytes(body) {
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to parse request body")
+ return
+ }
+
+ modelResult := gjson.GetBytes(body, "model")
+ if !modelResult.Exists() || modelResult.Type != gjson.String || strings.TrimSpace(modelResult.String()) == "" {
+ h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "model is required")
+ return
+ }
+ reqModel := modelResult.String()
+ reqLog = reqLog.With(zap.String("model", reqModel))
+ setOpsRequestContext(c, reqModel, false)
+ setOpsEndpointContext(c, "", int16(service.RequestTypeSync))
+
+ channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, reqModel)
+
+ subscription, _ := middleware2.GetSubscriptionFromContext(c)
+ service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds())
+
+ userReleaseFunc, acquired := h.acquireResponsesUserSlot(c, subject.UserID, subject.Concurrency, false, &streamStarted, reqLog)
+ if !acquired {
+ return
+ }
+ if userReleaseFunc != nil {
+ defer userReleaseFunc()
+ }
+
+ if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil {
+ reqLog.Info("openai_embeddings.billing_check_failed", zap.Error(err))
+ status, code, message, retryAfter := billingErrorDetails(err)
+ if retryAfter > 0 {
+ c.Header("Retry-After", strconv.Itoa(retryAfter))
+ }
+ h.errorResponse(c, status, code, message)
+ return
+ }
+
+ failedAccountIDs := make(map[int64]struct{})
+ var lastFailoverErr *service.UpstreamFailoverError
+ switchCount := 0
+ maxAccountSwitches := h.maxAccountSwitches
+ if maxAccountSwitches <= 0 {
+ maxAccountSwitches = 3
+ }
+ routingStart := time.Now()
+
+ for {
+ selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
+ c.Request.Context(),
+ apiKey.GroupID,
+ "",
+ "",
+ reqModel,
+ failedAccountIDs,
+ service.OpenAIUpstreamTransportHTTPSSE,
+ service.OpenAIEndpointCapabilityEmbeddings,
+ false,
+ )
+ if err != nil {
+ reqLog.Warn("openai_embeddings.account_select_failed",
+ zap.Error(err),
+ zap.Int("excluded_account_count", len(failedAccountIDs)),
+ )
+ if len(failedAccountIDs) == 0 {
+ markOpsRoutingCapacityLimitedIfNoAvailable(c, err)
+ h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "Service temporarily unavailable")
+ return
+ }
+ if lastFailoverErr != nil {
+ h.handleFailoverExhausted(c, lastFailoverErr, false)
+ } else {
+ h.errorResponse(c, http.StatusBadGateway, "api_error", "Upstream request failed")
+ }
+ return
+ }
+ if selection == nil || selection.Account == nil {
+ markOpsRoutingCapacityLimited(c)
+ h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available accounts")
+ return
+ }
+ account := selection.Account
+ setOpsSelectedAccount(c, account.ID, account.Platform)
+
+ accountReleaseFunc, accountAcquired := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, false, &streamStarted, reqLog)
+ if !accountAcquired {
+ return
+ }
+
+ service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds())
+ forwardStart := time.Now()
+
+ forwardBody := body
+ if channelMapping.Mapped {
+ forwardBody = h.gatewayService.ReplaceModelInBody(body, channelMapping.MappedModel)
+ }
+ writerSizeBeforeForward := c.Writer.Size()
+ result, err := func() (*service.OpenAIForwardResult, error) {
+ defer func() {
+ if accountReleaseFunc != nil {
+ accountReleaseFunc()
+ }
+ }()
+ return h.gatewayService.ForwardEmbeddings(c.Request.Context(), c, account, forwardBody, "")
+ }()
+
+ forwardDurationMs := time.Since(forwardStart).Milliseconds()
+ upstreamLatencyMs, _ := getContextInt64(c, service.OpsUpstreamLatencyMsKey)
+ responseLatencyMs := forwardDurationMs
+ if upstreamLatencyMs > 0 && forwardDurationMs > upstreamLatencyMs {
+ responseLatencyMs = forwardDurationMs - upstreamLatencyMs
+ }
+ service.SetOpsLatencyMs(c, service.OpsResponseLatencyMsKey, responseLatencyMs)
+
+ if err != nil {
+ var failoverErr *service.UpstreamFailoverError
+ if errors.As(err, &failoverErr) {
+ if c.Writer.Size() != writerSizeBeforeForward {
+ h.handleFailoverExhausted(c, failoverErr, true)
+ return
+ }
+ h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
+ h.gatewayService.RecordOpenAIAccountSwitch()
+ failedAccountIDs[account.ID] = struct{}{}
+ lastFailoverErr = failoverErr
+ if switchCount >= maxAccountSwitches {
+ h.handleFailoverExhausted(c, failoverErr, false)
+ return
+ }
+ switchCount++
+ reqLog.Warn("openai_embeddings.upstream_failover_switching",
+ zap.Int64("account_id", account.ID),
+ zap.Int("upstream_status", failoverErr.StatusCode),
+ zap.Int("switch_count", switchCount),
+ zap.Int("max_switches", maxAccountSwitches),
+ )
+ continue
+ }
+ h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
+ if c.Writer.Size() == writerSizeBeforeForward {
+ h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
+ }
+ reqLog.Warn("openai_embeddings.forward_failed",
+ zap.Int64("account_id", account.ID),
+ zap.Error(err),
+ )
+ return
+ }
+
+ h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, nil)
+ userAgent := c.GetHeader("User-Agent")
+ clientIP := ip.GetClientIP(c)
+ inboundEndpoint := GetInboundEndpoint(c)
+ upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
+
+ h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
+ if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
+ Result: result,
+ APIKey: apiKey,
+ User: apiKey.User,
+ Account: account,
+ Subscription: subscription,
+ InboundEndpoint: inboundEndpoint,
+ UpstreamEndpoint: upstreamEndpoint,
+ UserAgent: userAgent,
+ IPAddress: clientIP,
+ APIKeyService: h.apiKeyService,
+ ChannelUsageFields: channelMapping.ToUsageFields(reqModel, result.UpstreamModel),
+ }); err != nil {
+ logger.L().With(
+ zap.String("component", "handler.openai_gateway.embeddings"),
+ zap.Int64("user_id", subject.UserID),
+ zap.Int64("api_key_id", apiKey.ID),
+ zap.Any("group_id", apiKey.GroupID),
+ zap.String("model", reqModel),
+ zap.Int64("account_id", account.ID),
+ ).Error("openai_embeddings.record_usage_failed", zap.Error(err))
+ }
+ })
+ reqLog.Debug("openai_embeddings.request_completed",
+ zap.Int64("account_id", account.ID),
+ zap.Int("switch_count", switchCount),
+ )
+ return
+ }
+}
diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go
index 88ece8e7..f3d4caf0 100644
--- a/backend/internal/handler/openai_gateway_handler.go
+++ b/backend/internal/handler/openai_gateway_handler.go
@@ -12,6 +12,7 @@ import (
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil"
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
@@ -46,6 +47,31 @@ func resolveOpenAIMessagesDispatchMappedModel(apiKey *service.APIKey, requestedM
return strings.TrimSpace(apiKey.Group.ResolveMessagesDispatchModel(requestedModel))
}
+func usageRecordContext(parent context.Context, base context.Context) context.Context {
+ if base == nil {
+ base = context.Background()
+ }
+ if parent == nil {
+ return base
+ }
+ if clientRequestID, _ := parent.Value(ctxkey.ClientRequestID).(string); strings.TrimSpace(clientRequestID) != "" {
+ base = context.WithValue(base, ctxkey.ClientRequestID, strings.TrimSpace(clientRequestID))
+ }
+ if requestID, _ := parent.Value(ctxkey.RequestID).(string); strings.TrimSpace(requestID) != "" {
+ base = context.WithValue(base, ctxkey.RequestID, strings.TrimSpace(requestID))
+ }
+ return base
+}
+
+func wrapUsageRecordTaskContext(parent context.Context, task service.UsageRecordTask) service.UsageRecordTask {
+ if task == nil {
+ return nil
+ }
+ return func(ctx context.Context) {
+ task(usageRecordContext(parent, ctx))
+ }
+}
+
// NewOpenAIGatewayHandler creates a new OpenAIGatewayHandler
func NewOpenAIGatewayHandler(
gatewayService *service.OpenAIGatewayService,
@@ -266,7 +292,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
for {
// Select account supporting the requested model
reqLog.Debug("openai.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs)))
- selection, scheduleDecision, err := h.gatewayService.SelectAccountWithScheduler(
+ selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(),
apiKey.GroupID,
previousResponseID,
@@ -274,6 +300,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
reqModel,
failedAccountIDs,
service.OpenAIUpstreamTransportAny,
+ service.OpenAIEndpointCapabilityChatCompletions,
requireCompact,
)
if err != nil {
@@ -437,7 +464,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
- h.submitOpenAIUsageRecordTask(result, func(ctx context.Context) {
+ h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
Result: result,
APIKey: apiKey,
@@ -675,7 +702,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
currentRoutingModel = effectiveMappedModel
}
reqLog.Debug("openai_messages.account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs)))
- selection, scheduleDecision, err := h.gatewayService.SelectAccountWithScheduler(
+ selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
c.Request.Context(),
apiKey.GroupID,
"", // no previous_response_id
@@ -683,6 +710,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
currentRoutingModel,
failedAccountIDs,
service.OpenAIUpstreamTransportAny,
+ service.OpenAIEndpointCapabilityChatCompletions,
false,
)
if err != nil {
@@ -821,7 +849,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
- h.submitOpenAIUsageRecordTask(result, func(ctx context.Context) {
+ h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
Result: result,
APIKey: apiKey,
@@ -920,19 +948,12 @@ func (h *OpenAIGatewayHandler) validateFunctionCallOutputRequest(c *gin.Context,
return true
}
- var reqBody map[string]any
- if err := json.Unmarshal(body, &reqBody); err != nil {
- // 保持原有容错语义:解析失败时跳过预校验,沿用后续上游校验结果。
- return true
- }
-
- c.Set(service.OpenAIParsedRequestBodyKey, reqBody)
- validation := service.ValidateFunctionCallOutputContext(reqBody)
+ validation := service.ValidateFunctionCallOutputContextBytes(body)
if !validation.HasFunctionCallOutput {
return true
}
- previousResponseID, _ := reqBody["previous_response_id"].(string)
+ previousResponseID := gjson.GetBytes(body, "previous_response_id").String()
if strings.TrimSpace(previousResponseID) != "" || validation.HasToolCallContext {
return true
}
@@ -1146,7 +1167,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
defer func() {
_ = wsConn.CloseNow()
}()
- wsConn.SetReadLimit(16 * 1024 * 1024)
+ wsConn.SetReadLimit(service.ResolveOpenAIWSClientReadLimitBytes(h.cfg))
ctx := c.Request.Context()
readCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
@@ -1209,11 +1230,14 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
var currentUserRelease func()
var currentAccountRelease func()
- releaseTurnSlots := func() {
+ releaseAccountSlot := func() {
if currentAccountRelease != nil {
currentAccountRelease()
currentAccountRelease = nil
}
+ }
+ releaseTurnSlots := func() {
+ releaseAccountSlot()
if currentUserRelease != nil {
currentUserRelease()
currentUserRelease = nil
@@ -1233,6 +1257,23 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
return
}
currentUserRelease = wrapReleaseOnDone(ctx, userReleaseFunc)
+ ensureUserSlotHeld := func() bool {
+ if currentUserRelease != nil {
+ return true
+ }
+ userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency)
+ if err != nil {
+ reqLog.Warn("openai.websocket_user_slot_reacquire_failed", zap.Error(err))
+ closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to acquire user concurrency slot")
+ return false
+ }
+ if !userAcquired {
+ closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "too many concurrent requests, please retry later")
+ return false
+ }
+ currentUserRelease = wrapReleaseOnDone(ctx, userReleaseFunc)
+ return true
+ }
subscription, _ := middleware2.GetSubscriptionFromContext(c)
if err := h.billingCacheService.CheckBillingEligibility(ctx, apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil {
@@ -1246,195 +1287,249 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
firstMessage,
openAIWSIngressFallbackSessionSeed(subject.UserID, apiKey.ID, apiKey.GroupID),
)
- selection, scheduleDecision, err := h.gatewayService.SelectAccountWithScheduler(
- ctx,
- apiKey.GroupID,
- previousResponseID,
- sessionHash,
- reqModel,
- nil,
- service.OpenAIUpstreamTransportResponsesWebsocketV2,
- false,
- )
- if err != nil {
- reqLog.Warn("openai.websocket_account_select_failed", zap.Error(err))
- closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "no available account")
- return
- }
- if selection == nil || selection.Account == nil {
- closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "no available account")
- return
- }
+ maxAccountSwitches := h.maxAccountSwitches
+ switchCount := 0
+ failedAccountIDs := make(map[int64]struct{})
+ var lastFailoverErr *service.UpstreamFailoverError
- account := selection.Account
- accountMaxConcurrency := account.Concurrency
- if selection.WaitPlan != nil && selection.WaitPlan.MaxConcurrency > 0 {
- accountMaxConcurrency = selection.WaitPlan.MaxConcurrency
- }
- accountReleaseFunc := selection.ReleaseFunc
- if !selection.Acquired {
- if selection.WaitPlan == nil {
- closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "account is busy, please retry later")
- return
- }
- fastReleaseFunc, fastAcquired, err := h.concurrencyHelper.TryAcquireAccountSlot(
+ for {
+ reqLog.Debug("openai.websocket_account_selecting", zap.Int("excluded_account_count", len(failedAccountIDs)))
+ selection, scheduleDecision, err := h.gatewayService.SelectAccountWithSchedulerForCapability(
ctx,
- account.ID,
- selection.WaitPlan.MaxConcurrency,
+ apiKey.GroupID,
+ previousResponseID,
+ sessionHash,
+ reqModel,
+ failedAccountIDs,
+ service.OpenAIUpstreamTransportResponsesWebsocketV2,
+ service.OpenAIEndpointCapabilityChatCompletions,
+ false,
)
if err != nil {
- reqLog.Warn("openai.websocket_account_slot_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err))
- closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to acquire account concurrency slot")
+ reqLog.Warn("openai.websocket_account_select_failed",
+ zap.Error(err),
+ zap.Int("excluded_account_count", len(failedAccountIDs)),
+ )
+ if lastFailoverErr != nil {
+ closeOpenAIWSFailoverExhausted(wsConn, lastFailoverErr)
+ } else {
+ closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "no available account")
+ }
return
}
- if !fastAcquired {
- closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "account is busy, please retry later")
+ if selection == nil || selection.Account == nil {
+ if lastFailoverErr != nil {
+ closeOpenAIWSFailoverExhausted(wsConn, lastFailoverErr)
+ } else {
+ closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "no available account")
+ }
return
}
- accountReleaseFunc = fastReleaseFunc
- }
- currentAccountRelease = wrapReleaseOnDone(ctx, accountReleaseFunc)
- if err := h.gatewayService.BindStickySession(ctx, apiKey.GroupID, sessionHash, account.ID); err != nil {
- reqLog.Warn("openai.websocket_bind_sticky_session_failed", zap.Int64("account_id", account.ID), zap.Error(err))
- }
- token, _, err := h.gatewayService.GetAccessToken(ctx, account)
- if err != nil {
- reqLog.Warn("openai.websocket_get_access_token_failed", zap.Int64("account_id", account.ID), zap.Error(err))
- closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to get access token")
- return
- }
-
- reqLog.Debug("openai.websocket_account_selected",
- zap.Int64("account_id", account.ID),
- zap.String("account_name", account.Name),
- zap.String("schedule_layer", scheduleDecision.Layer),
- zap.Int("candidate_count", scheduleDecision.CandidateCount),
- )
-
- hooks := &service.OpenAIWSIngressHooks{
- InitialRequestModel: reqModel,
- BeforeRequest: func(turn int, payload []byte, originalModel string) error {
- if turn == 1 {
- return nil
- }
- if !gjson.ValidBytes(payload) {
- return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", errors.New("invalid json"))
- }
- model := strings.TrimSpace(originalModel)
- if model == "" {
- model = strings.TrimSpace(gjson.GetBytes(payload, "model").String())
- }
- if model == "" {
- model = reqModel
- }
- if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, model, payload); decision != nil && decision.Blocked {
- writeContentModerationWSError(ctx, wsConn, decision)
- return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, decision.Message, nil)
- }
- return nil
- },
- BeforeTurn: func(turn int) error {
- if turn == 1 {
- return nil
- }
- // 防御式清理:避免异常路径下旧槽位覆盖导致泄漏。
- releaseTurnSlots()
- // 非首轮 turn 需要重新抢占并发槽位,避免长连接空闲占槽。
- userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency)
- if err != nil {
- return service.NewOpenAIWSClientCloseError(coderws.StatusInternalError, "failed to acquire user concurrency slot", err)
- }
- if !userAcquired {
- return service.NewOpenAIWSClientCloseError(coderws.StatusTryAgainLater, "too many concurrent requests, please retry later", nil)
- }
- accountReleaseFunc, accountAcquired, err := h.concurrencyHelper.TryAcquireAccountSlot(ctx, account.ID, accountMaxConcurrency)
- if err != nil {
- if userReleaseFunc != nil {
- userReleaseFunc()
- }
- return service.NewOpenAIWSClientCloseError(coderws.StatusInternalError, "failed to acquire account concurrency slot", err)
- }
- if !accountAcquired {
- if userReleaseFunc != nil {
- userReleaseFunc()
- }
- return service.NewOpenAIWSClientCloseError(coderws.StatusTryAgainLater, "account is busy, please retry later", nil)
- }
- currentUserRelease = wrapReleaseOnDone(ctx, userReleaseFunc)
- currentAccountRelease = wrapReleaseOnDone(ctx, accountReleaseFunc)
- return nil
- },
- AfterTurn: func(turn int, result *service.OpenAIForwardResult, turnErr error) {
- releaseTurnSlots()
- if turnErr != nil {
- if result == nil || result.ImageCount <= 0 {
- return
- }
- reqLog.Warn("openai.websocket_partial_error_with_image_result",
- zap.Int64("account_id", account.ID),
- zap.Int("image_count", result.ImageCount),
- zap.Error(turnErr),
- )
- }
- if result == nil {
+ account := selection.Account
+ accountMaxConcurrency := account.Concurrency
+ if selection.WaitPlan != nil && selection.WaitPlan.MaxConcurrency > 0 {
+ accountMaxConcurrency = selection.WaitPlan.MaxConcurrency
+ }
+ accountReleaseFunc := selection.ReleaseFunc
+ if !selection.Acquired {
+ if selection.WaitPlan == nil {
+ closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "account is busy, please retry later")
return
}
- if account.Type == service.AccountTypeOAuth {
- h.gatewayService.UpdateCodexUsageSnapshotFromHeaders(ctx, account.ID, result.ResponseHeaders)
+ fastReleaseFunc, fastAcquired, err := h.concurrencyHelper.TryAcquireAccountSlot(
+ ctx,
+ account.ID,
+ selection.WaitPlan.MaxConcurrency,
+ )
+ if err != nil {
+ reqLog.Warn("openai.websocket_account_slot_acquire_failed", zap.Int64("account_id", account.ID), zap.Error(err))
+ closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to acquire account concurrency slot")
+ return
}
- h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, result.FirstTokenMs)
- inboundEndpoint := GetInboundEndpoint(c)
- upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
- h.submitOpenAIUsageRecordTask(result, func(taskCtx context.Context) {
- if err := h.gatewayService.RecordUsage(taskCtx, &service.OpenAIRecordUsageInput{
- Result: result,
- APIKey: apiKey,
- User: apiKey.User,
- Account: account,
- Subscription: subscription,
- InboundEndpoint: inboundEndpoint,
- UpstreamEndpoint: upstreamEndpoint,
- UserAgent: userAgent,
- IPAddress: clientIP,
- RequestPayloadHash: service.HashUsageRequestPayload(firstMessage),
- APIKeyService: h.apiKeyService,
- ChannelUsageFields: channelMappingWS.ToUsageFields(reqModel, result.UpstreamModel),
- }); err != nil {
- reqLog.Error("openai.websocket_record_usage_failed",
- zap.Int64("account_id", account.ID),
- zap.String("request_id", result.RequestID),
- zap.Error(err),
- )
- }
- })
- },
- }
+ if !fastAcquired {
+ closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "account is busy, please retry later")
+ return
+ }
+ accountReleaseFunc = fastReleaseFunc
+ }
+ currentAccountRelease = wrapReleaseOnDone(ctx, accountReleaseFunc)
+ if err := h.gatewayService.BindStickySession(ctx, apiKey.GroupID, sessionHash, account.ID); err != nil {
+ reqLog.Warn("openai.websocket_bind_sticky_session_failed", zap.Int64("account_id", account.ID), zap.Error(err))
+ }
- // 应用渠道模型映射到 WebSocket 首条消息
- wsFirstMessage := firstMessage
- if channelMappingWS.Mapped {
- wsFirstMessage = h.gatewayService.ReplaceModelInBody(firstMessage, channelMappingWS.MappedModel)
- }
-
- if err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks); err != nil {
- h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
- closeStatus, closeReason := summarizeWSCloseErrorForLog(err)
- reqLog.Warn("openai.websocket_proxy_failed",
- zap.Int64("account_id", account.ID),
- zap.Error(err),
- zap.String("close_status", closeStatus),
- zap.String("close_reason", closeReason),
- )
- var closeErr *service.OpenAIWSClientCloseError
- if errors.As(err, &closeErr) {
- closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason())
+ token, _, err := h.gatewayService.GetAccessToken(ctx, account)
+ if err != nil {
+ reqLog.Warn("openai.websocket_get_access_token_failed", zap.Int64("account_id", account.ID), zap.Error(err))
+ closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "failed to get access token")
return
}
- closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "upstream websocket proxy failed")
+
+ reqLog.Debug("openai.websocket_account_selected",
+ zap.Int64("account_id", account.ID),
+ zap.String("account_name", account.Name),
+ zap.String("schedule_layer", scheduleDecision.Layer),
+ zap.Int("candidate_count", scheduleDecision.CandidateCount),
+ )
+
+ var requestPayloadHash string
+ hooks := &service.OpenAIWSIngressHooks{
+ InitialRequestModel: reqModel,
+ BeforeRequest: func(turn int, payload []byte, originalModel string) error {
+ if turn == 1 {
+ return nil
+ }
+ if !gjson.ValidBytes(payload) {
+ return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", errors.New("invalid json"))
+ }
+ model := strings.TrimSpace(originalModel)
+ if model == "" {
+ model = strings.TrimSpace(gjson.GetBytes(payload, "model").String())
+ }
+ if model == "" {
+ model = reqModel
+ }
+ if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, model, payload); decision != nil && decision.Blocked {
+ writeContentModerationWSError(ctx, wsConn, decision)
+ return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, decision.Message, nil)
+ }
+ return nil
+ },
+ BeforeTurn: func(turn int) error {
+ if turn == 1 {
+ return nil
+ }
+ // 防御式清理:避免异常路径下旧槽位覆盖导致泄漏。
+ releaseTurnSlots()
+ // 非首轮 turn 需要重新抢占并发槽位,避免长连接空闲占槽。
+ userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, subject.UserID, subject.Concurrency)
+ if err != nil {
+ return service.NewOpenAIWSClientCloseError(coderws.StatusInternalError, "failed to acquire user concurrency slot", err)
+ }
+ if !userAcquired {
+ return service.NewOpenAIWSClientCloseError(coderws.StatusTryAgainLater, "too many concurrent requests, please retry later", nil)
+ }
+ accountReleaseFunc, accountAcquired, err := h.concurrencyHelper.TryAcquireAccountSlot(ctx, account.ID, accountMaxConcurrency)
+ if err != nil {
+ if userReleaseFunc != nil {
+ userReleaseFunc()
+ }
+ return service.NewOpenAIWSClientCloseError(coderws.StatusInternalError, "failed to acquire account concurrency slot", err)
+ }
+ if !accountAcquired {
+ if userReleaseFunc != nil {
+ userReleaseFunc()
+ }
+ return service.NewOpenAIWSClientCloseError(coderws.StatusTryAgainLater, "account is busy, please retry later", nil)
+ }
+ currentUserRelease = wrapReleaseOnDone(ctx, userReleaseFunc)
+ currentAccountRelease = wrapReleaseOnDone(ctx, accountReleaseFunc)
+ return nil
+ },
+ AfterTurn: func(turn int, result *service.OpenAIForwardResult, turnErr error) {
+ releaseTurnSlots()
+ if turnErr != nil {
+ if result == nil || result.ImageCount <= 0 {
+ return
+ }
+ reqLog.Warn("openai.websocket_partial_error_with_image_result",
+ zap.Int64("account_id", account.ID),
+ zap.Int("image_count", result.ImageCount),
+ zap.Error(turnErr),
+ )
+ }
+ if result == nil {
+ return
+ }
+ if account.Type == service.AccountTypeOAuth {
+ h.gatewayService.UpdateCodexUsageSnapshotFromHeaders(ctx, account.ID, result.ResponseHeaders)
+ }
+ h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, true, result.FirstTokenMs)
+ inboundEndpoint := GetInboundEndpoint(c)
+ upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
+ h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) {
+ if err := h.gatewayService.RecordUsage(taskCtx, &service.OpenAIRecordUsageInput{
+ Result: result,
+ APIKey: apiKey,
+ User: apiKey.User,
+ Account: account,
+ Subscription: subscription,
+ InboundEndpoint: inboundEndpoint,
+ UpstreamEndpoint: upstreamEndpoint,
+ UserAgent: userAgent,
+ IPAddress: clientIP,
+ RequestPayloadHash: requestPayloadHash,
+ APIKeyService: h.apiKeyService,
+ ChannelUsageFields: channelMappingWS.ToUsageFields(reqModel, result.UpstreamModel),
+ }); err != nil {
+ reqLog.Error("openai.websocket_record_usage_failed",
+ zap.Int64("account_id", account.ID),
+ zap.String("request_id", result.RequestID),
+ zap.Error(err),
+ )
+ }
+ })
+ },
+ }
+
+ // 应用渠道模型映射到 WebSocket 首条消息
+ wsFirstMessage := firstMessage
+ if channelMappingWS.Mapped {
+ wsFirstMessage = h.gatewayService.ReplaceModelInBody(firstMessage, channelMappingWS.MappedModel)
+ }
+
+ // WebSocket 首包可能很大,hash 必须在 hooks 外算成字符串,避免 AfterTurn 闭包保活请求体。
+ requestPayloadHash = service.HashUsageRequestPayload(wsFirstMessage)
+
+ if err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks); err != nil {
+ var failoverErr *service.UpstreamFailoverError
+ if errors.As(err, &failoverErr) {
+ h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
+ releaseAccountSlot()
+ failedAccountIDs[account.ID] = struct{}{}
+ lastFailoverErr = failoverErr
+ if switchCount >= maxAccountSwitches {
+ closeOpenAIWSFailoverExhausted(wsConn, failoverErr)
+ return
+ }
+ switchCount++
+ if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount) {
+ closeOpenAIWSFailoverExhausted(wsConn, failoverErr)
+ return
+ }
+ h.gatewayService.RecordOpenAIAccountSwitch()
+ reqLog.Warn("openai.websocket_upstream_failover_switching",
+ zap.Int64("account_id", account.ID),
+ zap.Int("upstream_status", failoverErr.StatusCode),
+ zap.Int("switch_count", switchCount),
+ zap.Int("max_switches", maxAccountSwitches),
+ )
+ if !ensureUserSlotHeld() {
+ return
+ }
+ continue
+ }
+
+ h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, false, nil)
+ closeStatus, closeReason := summarizeWSCloseErrorForLog(err)
+ reqLog.Warn("openai.websocket_proxy_failed",
+ zap.Int64("account_id", account.ID),
+ zap.Error(err),
+ zap.String("close_status", closeStatus),
+ zap.String("close_reason", closeReason),
+ )
+ var closeErr *service.OpenAIWSClientCloseError
+ if errors.As(err, &closeErr) {
+ closeOpenAIClientWS(wsConn, closeErr.StatusCode(), closeErr.Reason())
+ return
+ }
+ closeOpenAIClientWS(wsConn, coderws.StatusInternalError, "upstream websocket proxy failed")
+ return
+ }
+ reqLog.Info("openai.websocket_ingress_closed", zap.Int64("account_id", account.ID))
return
}
- reqLog.Info("openai.websocket_ingress_closed", zap.Int64("account_id", account.ID))
+
}
func (h *OpenAIGatewayHandler) recoverResponsesPanic(c *gin.Context, streamStarted *bool) {
@@ -1540,10 +1635,11 @@ func getContextInt64(c *gin.Context, key string) (int64, bool) {
}
}
-func (h *OpenAIGatewayHandler) submitUsageRecordTask(task service.UsageRecordTask) {
+func (h *OpenAIGatewayHandler) submitUsageRecordTask(parent context.Context, task service.UsageRecordTask) {
if task == nil {
return
}
+ task = wrapUsageRecordTaskContext(parent, task)
if h.usageRecordWorkerPool != nil {
h.usageRecordWorkerPool.Submit(task)
return
@@ -1562,18 +1658,19 @@ func (h *OpenAIGatewayHandler) submitUsageRecordTask(task service.UsageRecordTas
task(ctx)
}
-func (h *OpenAIGatewayHandler) submitOpenAIUsageRecordTask(result *service.OpenAIForwardResult, task service.UsageRecordTask) {
+func (h *OpenAIGatewayHandler) submitOpenAIUsageRecordTask(parent context.Context, result *service.OpenAIForwardResult, task service.UsageRecordTask) {
if result != nil && result.ImageCount > 0 {
- h.submitMandatoryUsageRecordTask(task)
+ h.submitMandatoryUsageRecordTask(parent, task)
return
}
- h.submitUsageRecordTask(task)
+ h.submitUsageRecordTask(parent, task)
}
-func (h *OpenAIGatewayHandler) submitMandatoryUsageRecordTask(task service.UsageRecordTask) {
+func (h *OpenAIGatewayHandler) submitMandatoryUsageRecordTask(parent context.Context, task service.UsageRecordTask) {
if task == nil {
return
}
+ task = wrapUsageRecordTaskContext(parent, task)
if h.usageRecordWorkerPool != nil {
if mode := h.usageRecordWorkerPool.Submit(task); mode != service.UsageRecordSubmitModeDropped {
return
@@ -1616,10 +1713,10 @@ func (h *OpenAIGatewayHandler) acquireImageGenerationSlot(c *gin.Context, stream
return nil, false
}
-// handleConcurrencyError handles concurrency-related errors with proper 429 response
+// handleConcurrencyError handles concurrency-related acquire errors.
func (h *OpenAIGatewayHandler) handleConcurrencyError(c *gin.Context, err error, slotType string, streamStarted bool) {
- h.handleStreamingAwareError(c, http.StatusTooManyRequests, "rate_limit_error",
- fmt.Sprintf("Concurrency limit exceeded for %s, please retry later", slotType), streamStarted)
+ status, errType, message := concurrencyErrorResponse(err, slotType)
+ h.handleStreamingAwareError(c, status, errType, message, streamStarted)
}
func (h *OpenAIGatewayHandler) handleFailoverExhausted(c *gin.Context, failoverErr *service.UpstreamFailoverError, streamStarted bool) {
@@ -1800,6 +1897,23 @@ func closeOpenAIClientWS(conn *coderws.Conn, status coderws.StatusCode, reason s
_ = conn.CloseNow()
}
+func closeOpenAIWSFailoverExhausted(conn *coderws.Conn, failoverErr *service.UpstreamFailoverError) {
+ if failoverErr == nil {
+ closeOpenAIClientWS(conn, coderws.StatusInternalError, "upstream websocket proxy failed")
+ return
+ }
+ switch failoverErr.StatusCode {
+ case http.StatusTooManyRequests:
+ closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, "upstream rate limit exceeded, please retry later")
+ case 529, http.StatusInternalServerError, http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout:
+ closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, "upstream service temporarily unavailable")
+ case http.StatusUnauthorized, http.StatusForbidden:
+ closeOpenAIClientWS(conn, coderws.StatusPolicyViolation, "upstream websocket authentication failed")
+ default:
+ closeOpenAIClientWS(conn, coderws.StatusInternalError, "upstream websocket proxy failed")
+ }
+}
+
func writeContentModerationWSError(ctx context.Context, conn *coderws.Conn, decision *service.ContentModerationDecision) {
if conn == nil || decision == nil {
return
diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go
index b304640e..b3fb35ee 100644
--- a/backend/internal/handler/openai_gateway_handler_test.go
+++ b/backend/internal/handler/openai_gateway_handler_test.go
@@ -7,6 +7,7 @@ import (
"net/http"
"net/http/httptest"
"strings"
+ "sync"
"testing"
"time"
@@ -740,16 +741,31 @@ func (r *contentModerationHandlerSettingRepo) Delete(ctx context.Context, key st
}
type contentModerationHandlerTestRepo struct {
+ mu sync.Mutex
logs []service.ContentModerationLog
}
func (r *contentModerationHandlerTestRepo) CreateLog(ctx context.Context, log *service.ContentModerationLog) error {
if log != nil {
+ r.mu.Lock()
+ defer r.mu.Unlock()
r.logs = append(r.logs, *log)
}
return nil
}
+func (r *contentModerationHandlerTestRepo) resetLogs() {
+ r.mu.Lock()
+ defer r.mu.Unlock()
+ r.logs = nil
+}
+
+func (r *contentModerationHandlerTestRepo) logSnapshot() []service.ContentModerationLog {
+ r.mu.Lock()
+ defer r.mu.Unlock()
+ return append([]service.ContentModerationLog(nil), r.logs...)
+}
+
func (r *contentModerationHandlerTestRepo) ListLogs(ctx context.Context, filter service.ContentModerationLogFilter) ([]service.ContentModerationLog, *pagination.PaginationResult, error) {
return nil, nil, nil
}
@@ -808,7 +824,10 @@ func TestOpenAIResponsesWebSocket_ContentModerationBlocksFirstFrame(t *testing.T
})
require.NoError(t, err)
require.True(t, decision.Blocked)
- repo.logs = nil
+ require.Eventually(t, func() bool {
+ return len(repo.logSnapshot()) == 1
+ }, time.Second, 10*time.Millisecond)
+ repo.resetLogs()
h := &OpenAIGatewayHandler{
gatewayService: &service.OpenAIGatewayService{},
billingCacheService: &service.BillingCacheService{},
@@ -848,10 +867,14 @@ func TestOpenAIResponsesWebSocket_ContentModerationBlocksFirstFrame(t *testing.T
require.Equal(t, coderws.StatusPolicyViolation, closeErr.Code)
require.Contains(t, closeErr.Reason, "内容审计测试阻断")
}
- require.Len(t, repo.logs, 1)
- require.True(t, repo.logs[0].Flagged)
- require.Equal(t, service.ContentModerationActionBlock, repo.logs[0].Action)
- require.Equal(t, "bad prompt", repo.logs[0].InputExcerpt)
+ var logs []service.ContentModerationLog
+ require.Eventually(t, func() bool {
+ logs = repo.logSnapshot()
+ return len(logs) == 1
+ }, time.Second, 10*time.Millisecond)
+ require.True(t, logs[0].Flagged)
+ require.Equal(t, service.ContentModerationActionBlock, logs[0].Action)
+ require.Equal(t, "bad prompt", logs[0].InputExcerpt)
}
func TestOpenAIResponsesWebSocket_PassthroughUsageLogPersistsUserAgentAndReasoningEffort(t *testing.T) {
@@ -1075,6 +1098,52 @@ func (s *openAIWSUsageHandlerAccountRepoStub) GetByID(ctx context.Context, id in
return &account, nil
}
+type openAIWSFailoverHandlerAccountRepoStub struct {
+ service.AccountRepository
+ accounts []service.Account
+ rateLimitedIDs []int64
+}
+
+func (s *openAIWSFailoverHandlerAccountRepoStub) ListSchedulableByPlatform(ctx context.Context, platform string) ([]service.Account, error) {
+ out := make([]service.Account, 0, len(s.accounts))
+ for _, account := range s.accounts {
+ if account.Platform == platform && account.IsSchedulable() {
+ out = append(out, account)
+ }
+ }
+ return out, nil
+}
+
+func (s *openAIWSFailoverHandlerAccountRepoStub) ListSchedulableByGroupIDAndPlatform(ctx context.Context, groupID int64, platform string) ([]service.Account, error) {
+ return s.ListSchedulableByPlatform(ctx, platform)
+}
+
+func (s *openAIWSFailoverHandlerAccountRepoStub) ListSchedulableUngroupedByPlatform(ctx context.Context, platform string) ([]service.Account, error) {
+ return s.ListSchedulableByPlatform(ctx, platform)
+}
+
+func (s *openAIWSFailoverHandlerAccountRepoStub) GetByID(ctx context.Context, id int64) (*service.Account, error) {
+ for _, account := range s.accounts {
+ if account.ID == id {
+ acc := account
+ return &acc, nil
+ }
+ }
+ return nil, nil
+}
+
+func (s *openAIWSFailoverHandlerAccountRepoStub) SetRateLimited(ctx context.Context, id int64, resetAt time.Time) error {
+ s.rateLimitedIDs = append(s.rateLimitedIDs, id)
+ for i := range s.accounts {
+ if s.accounts[i].ID == id {
+ reset := resetAt
+ s.accounts[i].RateLimitResetAt = &reset
+ break
+ }
+ }
+ return nil
+}
+
type openAIWSUsageHandlerUsageLogRepoStub struct {
service.UsageLogRepository
created chan *service.UsageLog
@@ -1107,6 +1176,201 @@ func (s *openAIWSUsageHandlerChannelRepoStub) GetGroupPlatforms(ctx context.Cont
return out, nil
}
+func TestOpenAIResponsesWebSocket_FailoverOnUpstreamUsageLimitEvent(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ firstHitCh := make(chan []byte, 1)
+ secondHitCh := make(chan []byte, 1)
+
+ firstUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
+ if err != nil {
+ return
+ }
+ defer func() { _ = conn.CloseNow() }()
+
+ readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
+ _, payload, readErr := conn.Read(readCtx)
+ cancelRead()
+ if readErr == nil {
+ firstHitCh <- payload
+ }
+
+ writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
+ _ = conn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"error","error":{"code":"rate_limit_exceeded","type":"usage_limit_reached","message":"The usage limit has been reached"}}`))
+ cancelWrite()
+ }))
+ defer firstUpstream.Close()
+
+ secondUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
+ if err != nil {
+ return
+ }
+ defer func() { _ = conn.CloseNow() }()
+
+ readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
+ _, payload, readErr := conn.Read(readCtx)
+ cancelRead()
+ if readErr == nil {
+ secondHitCh <- payload
+ }
+
+ writeCtx, cancelWrite := context.WithTimeout(r.Context(), 3*time.Second)
+ _ = conn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.completed","response":{"id":"resp_ws_failover_ok","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`))
+ cancelWrite()
+ _ = conn.Close(coderws.StatusNormalClosure, "done")
+ }))
+ defer secondUpstream.Close()
+
+ groupID := int64(4202)
+ accounts := []service.Account{
+ {
+ ID: 9902,
+ Name: "openai-ws-rate-limited",
+ Platform: service.PlatformOpenAI,
+ Type: service.AccountTypeAPIKey,
+ Status: service.StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-first",
+ "base_url": firstUpstream.URL,
+ },
+ Extra: map[string]any{
+ "openai_apikey_responses_websockets_v2_enabled": true,
+ "openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough,
+ },
+ },
+ {
+ ID: 9903,
+ Name: "openai-ws-healthy",
+ Platform: service.PlatformOpenAI,
+ Type: service.AccountTypeAPIKey,
+ Status: service.StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 2,
+ Credentials: map[string]any{
+ "api_key": "sk-second",
+ "base_url": secondUpstream.URL,
+ },
+ Extra: map[string]any{
+ "openai_apikey_responses_websockets_v2_enabled": true,
+ "openai_apikey_responses_websockets_v2_mode": service.OpenAIWSIngressModePassthrough,
+ },
+ },
+ }
+
+ cfg := &config.Config{}
+ cfg.RunMode = config.RunModeSimple
+ cfg.Default.RateMultiplier = 1
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
+ cfg.Gateway.MaxAccountSwitches = 3
+
+ accountRepo := &openAIWSFailoverHandlerAccountRepoStub{accounts: accounts}
+ rateLimitSvc := service.NewRateLimitService(accountRepo, nil, cfg, nil, nil)
+ billingCacheSvc := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil)
+ gatewaySvc := service.NewOpenAIGatewayService(
+ accountRepo,
+ nil,
+ nil,
+ nil,
+ nil,
+ nil,
+ nil,
+ cfg,
+ nil,
+ nil,
+ service.NewBillingService(cfg, nil),
+ rateLimitSvc,
+ billingCacheSvc,
+ nil,
+ &service.DeferredService{},
+ nil,
+ nil,
+ nil,
+ nil,
+ nil,
+ nil,
+ )
+
+ cache := &concurrencyCacheMock{
+ acquireUserSlotFn: func(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error) {
+ return true, nil
+ },
+ acquireAccountSlotFn: func(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) {
+ return true, nil
+ },
+ }
+ h := &OpenAIGatewayHandler{
+ gatewayService: gatewaySvc,
+ billingCacheService: billingCacheSvc,
+ apiKeyService: &service.APIKeyService{},
+ concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(cache), SSEPingFormatNone, time.Second),
+ maxAccountSwitches: 3,
+ }
+
+ apiKey := &service.APIKey{
+ ID: 1802,
+ GroupID: &groupID,
+ User: &service.User{ID: 1702, Status: service.StatusActive},
+ Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI, Status: service.StatusActive},
+ }
+ router := gin.New()
+ router.Use(func(c *gin.Context) {
+ c.Set(string(middleware.ContextKeyAPIKey), apiKey)
+ c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: apiKey.User.ID, Concurrency: 1})
+ c.Next()
+ })
+ router.GET("/openai/v1/responses", h.ResponsesWebSocket)
+ handlerServer := httptest.NewServer(router)
+ defer handlerServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(
+ dialCtx,
+ "ws"+strings.TrimPrefix(handlerServer.URL, "http")+"/openai/v1/responses",
+ &coderws.DialOptions{CompressionMode: coderws.CompressionContextTakeover},
+ )
+ cancelDial()
+ require.NoError(t, err)
+ defer func() { _ = clientConn.CloseNow() }()
+
+ writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
+ err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false}`))
+ cancelWrite()
+ require.NoError(t, err)
+
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 5*time.Second)
+ _, event, err := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, err)
+ require.Equal(t, "response.completed", gjson.GetBytes(event, "type").String())
+ require.Equal(t, "resp_ws_failover_ok", gjson.GetBytes(event, "response.id").String())
+
+ select {
+ case <-firstHitCh:
+ case <-time.After(3 * time.Second):
+ t.Fatal("等待第一个上游收到首帧超时")
+ }
+ select {
+ case <-secondHitCh:
+ case <-time.After(3 * time.Second):
+ t.Fatal("等待第二个上游收到重放首帧超时")
+ }
+ require.Equal(t, []int64{int64(9902)}, accountRepo.rateLimitedIDs)
+}
+
func runOpenAIResponsesWebSocketUsageLogCase(t *testing.T, tc openAIResponsesWSUsageLogCase) openAIResponsesWSUsageLogResult {
t.Helper()
gin.SetMode(gin.TestMode)
diff --git a/backend/internal/handler/openai_gateway_usage_context_test.go b/backend/internal/handler/openai_gateway_usage_context_test.go
new file mode 100644
index 00000000..7091c9c0
--- /dev/null
+++ b/backend/internal/handler/openai_gateway_usage_context_test.go
@@ -0,0 +1,41 @@
+package handler
+
+import (
+ "context"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
+ "github.com/stretchr/testify/require"
+)
+
+func TestSubmitUsageRecordTaskCopiesRequestContext(t *testing.T) {
+ parent := context.WithValue(context.Background(), ctxkey.ClientRequestID, "client-request-123")
+ parent = context.WithValue(parent, ctxkey.RequestID, "request-456")
+
+ var gotClientRequestID string
+ var gotRequestID string
+ h := &GatewayHandler{}
+ h.submitUsageRecordTask(parent, func(ctx context.Context) {
+ gotClientRequestID, _ = ctx.Value(ctxkey.ClientRequestID).(string)
+ gotRequestID, _ = ctx.Value(ctxkey.RequestID).(string)
+ })
+
+ require.Equal(t, "client-request-123", gotClientRequestID)
+ require.Equal(t, "request-456", gotRequestID)
+}
+
+func TestOpenAISubmitUsageRecordTaskCopiesRequestContext(t *testing.T) {
+ parent := context.WithValue(context.Background(), ctxkey.ClientRequestID, "openai-client-request-123")
+ parent = context.WithValue(parent, ctxkey.RequestID, "openai-request-456")
+
+ var gotClientRequestID string
+ var gotRequestID string
+ h := &OpenAIGatewayHandler{}
+ h.submitUsageRecordTask(parent, func(ctx context.Context) {
+ gotClientRequestID, _ = ctx.Value(ctxkey.ClientRequestID).(string)
+ gotRequestID, _ = ctx.Value(ctxkey.RequestID).(string)
+ })
+
+ require.Equal(t, "openai-client-request-123", gotClientRequestID)
+ require.Equal(t, "openai-request-456", gotRequestID)
+}
diff --git a/backend/internal/handler/openai_images.go b/backend/internal/handler/openai_images.go
index bbb08014..1e3b5306 100644
--- a/backend/internal/handler/openai_images.go
+++ b/backend/internal/handler/openai_images.go
@@ -73,9 +73,10 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error())
return
}
+ requestModel := parsed.Model
reqLog = reqLog.With(
- zap.String("model", parsed.Model),
+ zap.String("model", requestModel),
zap.Bool("stream", parsed.Stream),
zap.Bool("multipart", parsed.Multipart),
zap.String("capability", string(parsed.RequiredCapability)),
@@ -85,7 +86,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
h.errorResponse(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
return
}
- if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, parsed.Model, parsed.ModerationBody()); decision != nil && decision.Blocked {
+ if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, parsed.ModerationBody()); decision != nil && decision.Blocked {
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
return
}
@@ -98,13 +99,13 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
}
if parsed.Multipart {
- setOpsRequestContext(c, parsed.Model, parsed.Stream)
+ setOpsRequestContext(c, requestModel, parsed.Stream)
} else {
- setOpsRequestContext(c, parsed.Model, parsed.Stream)
+ setOpsRequestContext(c, requestModel, parsed.Stream)
}
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(parsed.Stream, false)))
- channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, parsed.Model)
+ channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, requestModel)
if h.errorPassthroughService != nil {
service.BindErrorPassthroughService(c, h.errorPassthroughService)
@@ -147,7 +148,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
c.Request.Context(),
apiKey.GroupID,
sessionHash,
- parsed.Model,
+ requestModel,
failedAccountIDs,
parsed.RequiredCapability,
)
@@ -311,7 +312,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
if result != nil {
upstreamModel = result.UpstreamModel
}
- h.submitMandatoryUsageRecordTask(func(ctx context.Context) {
+ h.submitMandatoryUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
Result: result,
APIKey: apiKey,
@@ -324,14 +325,14 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
- ChannelUsageFields: channelMapping.ToUsageFields(parsed.Model, upstreamModel),
+ ChannelUsageFields: channelMapping.ToUsageFields(requestModel, upstreamModel),
}); err != nil {
logger.L().With(
zap.String("component", "handler.openai_gateway.images"),
zap.Int64("user_id", subject.UserID),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
- zap.String("model", parsed.Model),
+ zap.String("model", requestModel),
zap.Int64("account_id", account.ID),
).Error("openai.images.record_usage_failed", zap.Error(err))
}
diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go
index 168fc271..b86c7f69 100644
--- a/backend/internal/handler/ops_error_logger.go
+++ b/backend/internal/handler/ops_error_logger.go
@@ -71,6 +71,49 @@ const (
opsErrorLogBatchSize = 32
)
+// looksLikeSystemKey 粗筛"形似本系统 key"的输入:长度 16-128 且仅含 [a-zA-Z0-9_-]。
+// 不用前缀匹配(APIKeyPrefix 可配置)。用于反查审计表前挡掉随机扫描的乱码输入。
+func looksLikeSystemKey(key string) bool {
+ if len(key) < 16 || len(key) > 128 {
+ return false
+ }
+ for _, c := range key {
+ allowed := (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
+ (c >= '0' && c <= '9') || c == '_' || c == '-'
+ if !allowed {
+ return false
+ }
+ }
+ return true
+}
+
+// keyPrefix 返回脱敏前缀(前 n 个字符);不足 n 则原样返回。
+func keyPrefix(key string, n int) string {
+ if len(key) <= n {
+ return key
+ }
+ return key[:n]
+}
+
+// extractAttemptedKey 按认证中间件同样的顺序从请求头提取提交的 key 明文。
+// 与 api_key_auth.go:43-59 一致:Authorization 仅取 Bearer scheme,非 Bearer 则忽略并继续 x-api-key → x-goog-api-key。
+func extractAttemptedKey(c *gin.Context) string {
+ if h := c.GetHeader("Authorization"); h != "" {
+ parts := strings.SplitN(h, " ", 2)
+ if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") {
+ return strings.TrimSpace(parts[1])
+ }
+ // 非 Bearer:与中间件一致,忽略 Authorization,继续尝试其它 header(不在此 return)。
+ }
+ if k := c.GetHeader("x-api-key"); k != "" {
+ return strings.TrimSpace(k)
+ }
+ if k := c.GetHeader("x-goog-api-key"); k != "" {
+ return strings.TrimSpace(k)
+ }
+ return ""
+}
+
type opsErrorLogJob struct {
ops *service.OpsService
entry *service.OpsInsertErrorLogInput
@@ -546,7 +589,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
return
}
- apiKey, _ := middleware2.GetAPIKeyFromContext(c)
+ apiKey := getOpsAPIKey(c)
clientRequestID, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
model, _ := c.Get(opsModelKey)
@@ -721,6 +764,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
if apiKey != nil {
entry.APIKeyID = &apiKey.ID
+ entry.APIKeyPrefix = keyPrefix(apiKey.Key, 8)
if apiKey.User != nil {
entry.UserID = &apiKey.User.ID
}
@@ -765,7 +809,7 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
return
}
- apiKey, _ := middleware2.GetAPIKeyFromContext(c)
+ apiKey := getOpsAPIKey(c)
clientRequestID, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
@@ -911,6 +955,8 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
if apiKey != nil {
entry.APIKeyID = &apiKey.ID
+ // 有效(未删除)key 报错时快照前缀,key 之后被删也保留;与 INVALID_API_KEY 的 attempted_key_prefix 互斥。
+ entry.APIKeyPrefix = keyPrefix(apiKey.Key, 8)
if apiKey.User != nil {
entry.UserID = &apiKey.User.ID
}
@@ -929,6 +975,22 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc {
entry.ClientIP = &clientIP
}
+ // 已删除 key 归因:仅 INVALID_API_KEY 才尝试。响应已写出,此处不阻塞客户端。
+ if parsed.Code == opsCodeInvalidAPIKey {
+ if attemptedKey := extractAttemptedKey(c); attemptedKey != "" {
+ entry.AttemptedKeyPrefix = keyPrefix(attemptedKey, 8)
+ if looksLikeSystemKey(attemptedKey) {
+ if res, lookupErr := ops.LookupDeletedKeyAudit(c.Request.Context(), attemptedKey); lookupErr != nil {
+ log.Printf("[OpsErrorLogger] LookupDeletedKeyAudit failed: %v", lookupErr)
+ } else if res != nil {
+ owner := res.UserID
+ entry.DeletedKeyOwnerUserID = &owner
+ entry.DeletedKeyName = res.KeyName
+ }
+ }
+ }
+ }
+
enqueueOpsErrorLog(ops, entry)
}
}
@@ -1035,6 +1097,20 @@ func parseOpsErrorResponse(body []byte) parsedOpsError {
return parsedOpsError{Message: truncateString(string(body), 1024)}
}
+// getOpsAPIKey 返回用于 Ops 错误日志的 API Key:优先取已鉴权写入的正式 key;
+// 鉴权早退(分组停用/删除、Key 停用/过期/额度、用户停用、IP 限制等)时,
+// 正式 key 尚未写入,回退到 middleware 写入的 ops fallback key
+// (含 User/Group/Platform),从而让日志能展示 用户/分组/平台。
+func getOpsAPIKey(c *gin.Context) *service.APIKey {
+ if apiKey, ok := middleware2.GetAPIKeyFromContext(c); ok && apiKey != nil {
+ return apiKey
+ }
+ if apiKey, ok := middleware2.GetOpsFallbackAPIKey(c); ok && apiKey != nil {
+ return apiKey
+ }
+ return nil
+}
+
func resolveOpsPlatform(apiKey *service.APIKey, fallback string) string {
if apiKey != nil && apiKey.Group != nil && apiKey.Group.Platform != "" {
return apiKey.Group.Platform
diff --git a/backend/internal/handler/ops_error_logger_attribution_test.go b/backend/internal/handler/ops_error_logger_attribution_test.go
new file mode 100644
index 00000000..9c68d845
--- /dev/null
+++ b/backend/internal/handler/ops_error_logger_attribution_test.go
@@ -0,0 +1,118 @@
+package handler
+
+import (
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+)
+
+func TestLooksLikeSystemKey(t *testing.T) {
+ cases := []struct {
+ in string
+ want bool
+ }{
+ {"sk-abcdef0123456789", true},
+ {"ABCdef_-0123456789", true},
+ {"short", false},
+ {"with space xxxxxxxxxx", false},
+ {"汉字key1234567890", false},
+ {"", false},
+ }
+ for _, c := range cases {
+ if got := looksLikeSystemKey(c.in); got != c.want {
+ t.Errorf("looksLikeSystemKey(%q)=%v want %v", c.in, got, c.want)
+ }
+ }
+ long := make([]byte, 129)
+ for i := range long {
+ long[i] = 'a'
+ }
+ if looksLikeSystemKey(string(long)) {
+ t.Errorf("129-char key should be rejected")
+ }
+}
+
+func TestKeyPrefix(t *testing.T) {
+ if got := keyPrefix("sk-3f2a9c7e", 8); got != "sk-3f2a9" {
+ t.Errorf("keyPrefix=%q want %q", got, "sk-3f2a9")
+ }
+ if got := keyPrefix("abc", 8); got != "abc" {
+ t.Errorf("short key should be returned as-is, got %q", got)
+ }
+}
+
+func TestExtractAttemptedKey(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ cases := []struct {
+ name string
+ headers map[string]string
+ want string
+ }{
+ {
+ name: "Bearer in Authorization",
+ headers: map[string]string{"Authorization": "Bearer sk-testkey0123456789"},
+ want: "sk-testkey0123456789",
+ },
+ {
+ name: "Bearer case-insensitive",
+ headers: map[string]string{"Authorization": "BEARER sk-testkey0123456789"},
+ want: "sk-testkey0123456789",
+ },
+ {
+ name: "x-api-key header",
+ headers: map[string]string{"x-api-key": "sk-xapikey0123456789"},
+ want: "sk-xapikey0123456789",
+ },
+ {
+ name: "x-goog-api-key header",
+ headers: map[string]string{"x-goog-api-key": "sk-goog0123456789"},
+ want: "sk-goog0123456789",
+ },
+ {
+ name: "Authorization takes priority over x-api-key",
+ headers: map[string]string{"Authorization": "Bearer sk-auth0123456789", "x-api-key": "sk-xapi0123456789"},
+ want: "sk-auth0123456789",
+ },
+ {
+ name: "x-api-key takes priority over x-goog-api-key",
+ headers: map[string]string{"x-api-key": "sk-xapi0123456789", "x-goog-api-key": "sk-goog0123456789"},
+ want: "sk-xapi0123456789",
+ },
+ {
+ name: "no key headers",
+ headers: map[string]string{},
+ want: "",
+ },
+ {
+ name: "Bearer with leading/trailing spaces trimmed",
+ headers: map[string]string{"Authorization": "Bearer sk-trimmed0123456789 "},
+ want: "sk-trimmed0123456789",
+ },
+ {
+ // 非 Bearer Authorization 应被忽略,继续 fall-through 到 x-api-key(与认证中间件一致)
+ name: "non-Bearer Authorization falls through to x-api-key",
+ headers: map[string]string{"Authorization": "junk-not-bearer", "x-api-key": "sk-realkey1234567"},
+ want: "sk-realkey1234567",
+ },
+ }
+
+ for _, tc := range cases {
+ t.Run(tc.name, func(t *testing.T) {
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
+ for k, v := range tc.headers {
+ req.Header.Set(k, v)
+ }
+ c.Request = req
+
+ got := extractAttemptedKey(c)
+ if got != tc.want {
+ t.Errorf("extractAttemptedKey(%v) = %q, want %q", tc.headers, got, tc.want)
+ }
+ })
+ }
+}
diff --git a/backend/internal/handler/ops_error_logger_test.go b/backend/internal/handler/ops_error_logger_test.go
index d4e1177e..cf1685f2 100644
--- a/backend/internal/handler/ops_error_logger_test.go
+++ b/backend/internal/handler/ops_error_logger_test.go
@@ -931,3 +931,45 @@ func TestSetOpsEndpointContext_NilContext(t *testing.T) {
setOpsEndpointContext(nil, "model", int16(1))
})
}
+
+func TestGetOpsAPIKeyFallsBackToOpsFallbackKey(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+
+ // 主 key 缺席(鉴权早退场景):返回 nil。
+ require.Nil(t, getOpsAPIKey(c))
+
+ // 写入 ops 专用 fallback key 后应能取到,且带齐 user/group。
+ groupID := int64(55)
+ apiKey := &service.APIKey{
+ ID: 100,
+ GroupID: &groupID,
+ User: &service.User{ID: 7},
+ Group: &service.Group{ID: groupID, Platform: service.PlatformAnthropic},
+ }
+ c.Set(string(middleware2.ContextKeyOpsFallbackAPIKey), apiKey)
+
+ got := getOpsAPIKey(c)
+ require.NotNil(t, got)
+ require.Equal(t, int64(100), got.ID)
+ require.NotNil(t, got.User)
+ require.Equal(t, int64(7), got.User.ID)
+ require.NotNil(t, got.Group)
+ require.Equal(t, service.PlatformAnthropic, got.Group.Platform)
+}
+
+func TestGetOpsAPIKeyPrefersPrimaryContextKey(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+
+ primary := &service.APIKey{ID: 1}
+ fallback := &service.APIKey{ID: 2}
+ c.Set(string(middleware2.ContextKeyAPIKey), primary)
+ c.Set(string(middleware2.ContextKeyOpsFallbackAPIKey), fallback)
+
+ got := getOpsAPIKey(c)
+ require.NotNil(t, got)
+ require.Equal(t, int64(1), got.ID, "已鉴权请求应优先使用正式 api key")
+}
diff --git a/backend/internal/handler/setting_handler.go b/backend/internal/handler/setting_handler.go
index 9f7350f2..5cd27f5b 100644
--- a/backend/internal/handler/setting_handler.go
+++ b/backend/internal/handler/setting_handler.go
@@ -98,6 +98,8 @@ func (h *SettingHandler) GetPublicSettings(c *gin.Context) {
AffiliateEnabled: settings.AffiliateEnabled,
RiskControlEnabled: settings.RiskControlEnabled,
+
+ AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests,
})
}
diff --git a/backend/internal/handler/usage_handler.go b/backend/internal/handler/usage_handler.go
index daa5695d..23bb62dd 100644
--- a/backend/internal/handler/usage_handler.go
+++ b/backend/internal/handler/usage_handler.go
@@ -1,6 +1,7 @@
package handler
import (
+ "net/http"
"strconv"
"strings"
"time"
@@ -18,15 +19,24 @@ import (
// UsageHandler handles usage-related requests
type UsageHandler struct {
- usageService *service.UsageService
- apiKeyService *service.APIKeyService
+ usageService *service.UsageService
+ apiKeyService *service.APIKeyService
+ opsService *service.OpsService
+ settingService *service.SettingService
}
// NewUsageHandler creates a new UsageHandler
-func NewUsageHandler(usageService *service.UsageService, apiKeyService *service.APIKeyService) *UsageHandler {
+func NewUsageHandler(
+ usageService *service.UsageService,
+ apiKeyService *service.APIKeyService,
+ opsService *service.OpsService,
+ settingService *service.SettingService,
+) *UsageHandler {
return &UsageHandler{
- usageService: usageService,
- apiKeyService: apiKeyService,
+ usageService: usageService,
+ apiKeyService: apiKeyService,
+ opsService: opsService,
+ settingService: settingService,
}
}
@@ -149,6 +159,117 @@ func (h *UsageHandler) List(c *gin.Context) {
response.Paginated(c, out, result.Total, page, pageSize)
}
+// ListErrors handles listing the current user's failed requests (redacted).
+// GET /api/v1/usage/errors
+func (h *UsageHandler) ListErrors(c *gin.Context) {
+ subject, ok := middleware2.GetAuthSubjectFromContext(c)
+ if !ok {
+ response.Unauthorized(c, "User not authenticated")
+ return
+ }
+
+ // Visibility switch (fail-closed). Defense-in-depth: frontend also hides the tab.
+ if h.settingService == nil || !h.settingService.IsUserErrorViewAllowed(c.Request.Context()) {
+ response.Forbidden(c, "Error requests view is disabled")
+ return
+ }
+ if h.opsService == nil {
+ response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
+ return
+ }
+
+ page, pageSize := response.ParsePagination(c)
+ if pageSize > 100 {
+ pageSize = 100
+ }
+
+ filter := &service.OpsErrorLogFilter{Page: page, PageSize: pageSize}
+
+ // Date range (half-open [start, end)), reuse usage-list semantics.
+ userTZ := c.Query("timezone")
+ if startDateStr := c.Query("start_date"); startDateStr != "" {
+ t, err := timezone.ParseInUserLocation("2006-01-02", startDateStr, userTZ)
+ if err != nil {
+ response.BadRequest(c, "Invalid start_date format, use YYYY-MM-DD")
+ return
+ }
+ filter.StartTime = &t
+ }
+ if endDateStr := c.Query("end_date"); endDateStr != "" {
+ t, err := timezone.ParseInUserLocation("2006-01-02", endDateStr, userTZ)
+ if err != nil {
+ response.BadRequest(c, "Invalid end_date format, use YYYY-MM-DD")
+ return
+ }
+ t = t.AddDate(0, 0, 1)
+ filter.EndTime = &t
+ }
+
+ filter.Model = strings.TrimSpace(c.Query("model"))
+
+ if k := strings.TrimSpace(c.Query("api_key_id")); k != "" {
+ n, err := strconv.ParseInt(k, 10, 64)
+ if err != nil || n < 0 {
+ response.BadRequest(c, "Invalid api_key_id")
+ return
+ }
+ if n > 0 {
+ filter.APIKeyID = &n
+ }
+ }
+
+ if sc := strings.TrimSpace(c.Query("status_code")); sc != "" {
+ n, err := strconv.Atoi(sc)
+ if err != nil || n < 0 {
+ response.BadRequest(c, "Invalid status_code")
+ return
+ }
+ filter.StatusCodes = []int{n}
+ }
+
+ if cat := strings.TrimSpace(c.Query("category")); cat != "" {
+ phases, types := service.CategoryToFilter(cat)
+ filter.ErrorPhasesAny = phases
+ filter.ErrorTypesAny = types
+ }
+
+ result, err := h.opsService.ListUserErrorRequests(c.Request.Context(), subject.UserID, filter)
+ if err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+ response.Paginated(c, result.Items, int64(result.Total), result.Page, result.PageSize)
+}
+
+// GetErrorDetail handles fetching one of the current user's failed-request details (redacted).
+// GET /api/v1/usage/errors/:id
+func (h *UsageHandler) GetErrorDetail(c *gin.Context) {
+ subject, ok := middleware2.GetAuthSubjectFromContext(c)
+ if !ok {
+ response.Unauthorized(c, "User not authenticated")
+ return
+ }
+ if h.settingService == nil || !h.settingService.IsUserErrorViewAllowed(c.Request.Context()) {
+ response.Forbidden(c, "Error requests view is disabled")
+ return
+ }
+ if h.opsService == nil {
+ response.Error(c, http.StatusServiceUnavailable, "Ops service not available")
+ return
+ }
+ id, err := strconv.ParseInt(strings.TrimSpace(c.Param("id")), 10, 64)
+ if err != nil || id <= 0 {
+ response.BadRequest(c, "Invalid id")
+ return
+ }
+ detail, err := h.opsService.GetUserErrorRequestDetail(c.Request.Context(), subject.UserID, id)
+ if err != nil {
+ response.ErrorFrom(c, err)
+ return
+ }
+ response.Success(c, detail)
+}
+
// GetByID handles getting a single usage record
// GET /api/v1/usage/:id
func (h *UsageHandler) GetByID(c *gin.Context) {
diff --git a/backend/internal/handler/usage_handler_daily_test.go b/backend/internal/handler/usage_handler_daily_test.go
index 36311fac..2a9186cf 100644
--- a/backend/internal/handler/usage_handler_daily_test.go
+++ b/backend/internal/handler/usage_handler_daily_test.go
@@ -64,7 +64,7 @@ func newDailyUsageTestRouter(usageRepo *dailyUsageRepoStub, apiKeyRepo *dailyUsa
gin.SetMode(gin.TestMode)
usageSvc := service.NewUsageService(usageRepo, nil, nil, nil)
apiKeySvc := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, nil)
- handler := NewUsageHandler(usageSvc, apiKeySvc)
+ handler := NewUsageHandler(usageSvc, apiKeySvc, nil, nil)
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: userID})
diff --git a/backend/internal/handler/usage_handler_request_type_test.go b/backend/internal/handler/usage_handler_request_type_test.go
index b49ed59b..ed08c5a8 100644
--- a/backend/internal/handler/usage_handler_request_type_test.go
+++ b/backend/internal/handler/usage_handler_request_type_test.go
@@ -34,7 +34,7 @@ func (s *userUsageRepoCapture) ListWithFilters(ctx context.Context, params pagin
func newUserUsageRequestTypeTestRouter(repo *userUsageRepoCapture) *gin.Engine {
gin.SetMode(gin.TestMode)
usageSvc := service.NewUsageService(repo, nil, nil, nil)
- handler := NewUsageHandler(usageSvc, nil)
+ handler := NewUsageHandler(usageSvc, nil, nil, nil)
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: 42})
diff --git a/backend/internal/handler/usage_record_submit_task_test.go b/backend/internal/handler/usage_record_submit_task_test.go
index e4c2837a..ebe5c3df 100644
--- a/backend/internal/handler/usage_record_submit_task_test.go
+++ b/backend/internal/handler/usage_record_submit_task_test.go
@@ -29,7 +29,7 @@ func TestGatewayHandlerSubmitUsageRecordTask_WithPool(t *testing.T) {
h := &GatewayHandler{usageRecordWorkerPool: pool}
done := make(chan struct{})
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
close(done)
})
@@ -44,7 +44,7 @@ func TestGatewayHandlerSubmitUsageRecordTask_WithoutPoolSyncFallback(t *testing.
h := &GatewayHandler{}
var called atomic.Bool
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
if _, ok := ctx.Deadline(); !ok {
t.Fatal("expected deadline in fallback context")
}
@@ -57,7 +57,7 @@ func TestGatewayHandlerSubmitUsageRecordTask_WithoutPoolSyncFallback(t *testing.
func TestGatewayHandlerSubmitUsageRecordTask_NilTask(t *testing.T) {
h := &GatewayHandler{}
require.NotPanics(t, func() {
- h.submitUsageRecordTask(nil)
+ h.submitUsageRecordTask(context.Background(), nil)
})
}
@@ -66,12 +66,12 @@ func TestGatewayHandlerSubmitUsageRecordTask_WithoutPool_TaskPanicRecovered(t *t
var called atomic.Bool
require.NotPanics(t, func() {
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
panic("usage task panic")
})
})
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
called.Store(true)
})
require.True(t, called.Load(), "panic 后后续任务应仍可执行")
@@ -82,7 +82,7 @@ func TestOpenAIGatewayHandlerSubmitUsageRecordTask_WithPool(t *testing.T) {
h := &OpenAIGatewayHandler{usageRecordWorkerPool: pool}
done := make(chan struct{})
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
close(done)
})
@@ -97,7 +97,7 @@ func TestOpenAIGatewayHandlerSubmitUsageRecordTask_WithoutPoolSyncFallback(t *te
h := &OpenAIGatewayHandler{}
var called atomic.Bool
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
if _, ok := ctx.Deadline(); !ok {
t.Fatal("expected deadline in fallback context")
}
@@ -110,7 +110,7 @@ func TestOpenAIGatewayHandlerSubmitUsageRecordTask_WithoutPoolSyncFallback(t *te
func TestOpenAIGatewayHandlerSubmitUsageRecordTask_NilTask(t *testing.T) {
h := &OpenAIGatewayHandler{}
require.NotPanics(t, func() {
- h.submitUsageRecordTask(nil)
+ h.submitUsageRecordTask(context.Background(), nil)
})
}
@@ -119,12 +119,12 @@ func TestOpenAIGatewayHandlerSubmitUsageRecordTask_WithoutPool_TaskPanicRecovere
var called atomic.Bool
require.NotPanics(t, func() {
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
panic("usage task panic")
})
})
- h.submitUsageRecordTask(func(ctx context.Context) {
+ h.submitUsageRecordTask(context.Background(), func(ctx context.Context) {
called.Store(true)
})
require.True(t, called.Load(), "panic 后后续任务应仍可执行")
@@ -152,7 +152,7 @@ func TestOpenAIGatewayHandlerSubmitMandatoryUsageRecordTask_DroppedTaskSyncFallb
pool.Submit(func(ctx context.Context) {})
var called atomic.Bool
- h.submitMandatoryUsageRecordTask(func(ctx context.Context) {
+ h.submitMandatoryUsageRecordTask(context.Background(), func(ctx context.Context) {
called.Store(true)
})
close(release)
@@ -182,7 +182,7 @@ func TestOpenAIGatewayHandlerSubmitOpenAIUsageRecordTask_ImageResultUsesMandator
pool.Submit(func(ctx context.Context) {})
var called atomic.Bool
- h.submitOpenAIUsageRecordTask(&service.OpenAIForwardResult{ImageCount: 1}, func(ctx context.Context) {
+ h.submitOpenAIUsageRecordTask(context.Background(), &service.OpenAIForwardResult{ImageCount: 1}, func(ctx context.Context) {
called.Store(true)
})
close(release)
diff --git a/backend/internal/handler/user_handler_test.go b/backend/internal/handler/user_handler_test.go
index 41647802..2e366c23 100644
--- a/backend/internal/handler/user_handler_test.go
+++ b/backend/internal/handler/user_handler_test.go
@@ -118,6 +118,9 @@ func (s *userHandlerRepoStub) RemoveGroupFromUserAllowedGroups(context.Context,
func (s *userHandlerRepoStub) UpdateTotpSecret(context.Context, int64, *string) error { return nil }
func (s *userHandlerRepoStub) EnableTotp(context.Context, int64) error { return nil }
func (s *userHandlerRepoStub) DisableTotp(context.Context, int64) error { return nil }
+func (s *userHandlerRepoStub) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ return s.GetByID(ctx, id)
+}
func (s *userHandlerRepoStub) ListUserAuthIdentities(context.Context, int64) ([]service.UserAuthIdentityRecord, error) {
out := make([]service.UserAuthIdentityRecord, len(s.identities))
copy(out, s.identities)
diff --git a/backend/internal/payment/provider/easypay.go b/backend/internal/payment/provider/easypay.go
index e7d8aab9..32d6b7be 100644
--- a/backend/internal/payment/provider/easypay.go
+++ b/backend/internal/payment/provider/easypay.go
@@ -213,22 +213,59 @@ func (e *EasyPay) QueryOrder(ctx context.Context, tradeNo string) (*payment.Quer
if err != nil {
return nil, fmt.Errorf("easypay query: %w", err)
}
+ type easyPayQueryData struct {
+ TradeStatus *string `json:"trade_status"`
+ Status *int `json:"status"`
+ Money *string `json:"money"`
+ TradeNo *string `json:"trade_no"`
+ }
var resp struct {
- Code int `json:"code"`
- Msg string `json:"msg"`
- Status int `json:"status"`
- Money string `json:"money"`
+ Code int `json:"code"`
+ Msg string `json:"msg"`
+ TradeStatus *string `json:"trade_status"`
+ Status *int `json:"status"`
+ Money *string `json:"money"`
+ TradeNo *string `json:"trade_no"`
+ Data easyPayQueryData `json:"data"`
}
if err := json.Unmarshal(body, &resp); err != nil {
return nil, fmt.Errorf("easypay parse query: %w", err)
}
status := payment.ProviderStatusPending
- if resp.Status == easypayStatusPaid {
+ if resp.TradeStatus != nil {
+ if *resp.TradeStatus == tradeStatusSuccess {
+ status = payment.ProviderStatusPaid
+ }
+ } else if resp.Data.TradeStatus != nil {
+ if *resp.Data.TradeStatus == tradeStatusSuccess {
+ status = payment.ProviderStatusPaid
+ }
+ } else if resp.Status != nil {
+ if *resp.Status == easypayStatusPaid {
+ status = payment.ProviderStatusPaid
+ }
+ } else if resp.Data.Status != nil && *resp.Data.Status == easypayStatusPaid {
status = payment.ProviderStatusPaid
}
- amount, _ := strconv.ParseFloat(resp.Money, 64)
+
+ money := ""
+ if resp.Money != nil {
+ money = *resp.Money
+ } else if resp.Data.Money != nil {
+ money = *resp.Data.Money
+ }
+ responseTradeNo := tradeNo
+ if resp.TradeNo != nil {
+ if *resp.TradeNo != "" {
+ responseTradeNo = *resp.TradeNo
+ }
+ } else if resp.Data.TradeNo != nil && *resp.Data.TradeNo != "" {
+ responseTradeNo = *resp.Data.TradeNo
+ }
+
+ amount, _ := strconv.ParseFloat(money, 64)
return &payment.QueryOrderResponse{
- TradeNo: tradeNo,
+ TradeNo: responseTradeNo,
Status: status,
Amount: amount,
Metadata: e.MerchantIdentityMetadata(),
diff --git a/backend/internal/payment/provider/easypay_query_test.go b/backend/internal/payment/provider/easypay_query_test.go
new file mode 100644
index 00000000..5042a94d
--- /dev/null
+++ b/backend/internal/payment/provider/easypay_query_test.go
@@ -0,0 +1,131 @@
+package provider
+
+import (
+ "context"
+ "net/http"
+ "net/http/httptest"
+ "net/url"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/payment"
+)
+
+func TestEasyPayQueryOrderStatusMapping(t *testing.T) {
+ t.Parallel()
+
+ const orderID = "order-123"
+ tests := []struct {
+ name string
+ body string
+ wantStatus string
+ wantTradeNo string
+ wantAmount float64
+ }{
+ {
+ name: "top level trade success is paid",
+ body: `{"code":1,"trade_status":"TRADE_SUCCESS","status":0,"money":"12.34","trade_no":"gateway-123"}`,
+ wantStatus: payment.ProviderStatusPaid,
+ wantTradeNo: "gateway-123",
+ wantAmount: 12.34,
+ },
+ {
+ name: "waiting trade status with paid numeric status stays pending",
+ body: `{"code":1,"trade_status":"WAITING","status":1,"money":"12.34","trade_no":"gateway-123"}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: "gateway-123",
+ wantAmount: 12.34,
+ },
+ {
+ name: "empty trade status with paid numeric status stays pending",
+ body: `{"code":1,"trade_status":"","status":1,"money":"12.34"}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: orderID,
+ wantAmount: 12.34,
+ },
+ {
+ name: "nested data trade success is paid",
+ body: `{"code":1,"data":{"trade_status":"TRADE_SUCCESS","status":0,"money":"9.99","trade_no":"data-456"}}`,
+ wantStatus: payment.ProviderStatusPaid,
+ wantTradeNo: "data-456",
+ wantAmount: 9.99,
+ },
+ {
+ name: "legacy numeric paid status remains compatible",
+ body: `{"code":1,"status":1,"money":"3.21"}`,
+ wantStatus: payment.ProviderStatusPaid,
+ wantTradeNo: orderID,
+ wantAmount: 3.21,
+ },
+ {
+ name: "legacy numeric non paid status is pending",
+ body: `{"code":1,"status":0,"money":"3.21"}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: orderID,
+ wantAmount: 3.21,
+ },
+ {
+ name: "query failure with missing status is pending",
+ body: `{"code":0,"msg":"订单不存在"}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: orderID,
+ },
+ {
+ name: "missing fields are pending",
+ body: `{}`,
+ wantStatus: payment.ProviderStatusPending,
+ wantTradeNo: orderID,
+ },
+ }
+
+ for _, tt := range tests {
+ tt := tt
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+
+ var gotForm url.Values
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ t.Errorf("method = %q, want %q", r.Method, http.MethodPost)
+ }
+ if r.URL.Path != "/api.php" {
+ t.Errorf("path = %q, want /api.php", r.URL.Path)
+ }
+ if err := r.ParseForm(); err != nil {
+ t.Errorf("ParseForm: %v", err)
+ }
+ gotForm = make(url.Values, len(r.PostForm))
+ for key, values := range r.PostForm {
+ gotForm[key] = append([]string(nil), values...)
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = w.Write([]byte(tt.body))
+ }))
+ defer server.Close()
+
+ provider := newTestEasyPay(t, server.URL)
+ resp, err := provider.QueryOrder(context.Background(), orderID)
+ if err != nil {
+ t.Fatalf("QueryOrder returned error: %v", err)
+ }
+ if resp.Status != tt.wantStatus {
+ t.Fatalf("status = %q, want %q (response=%+v)", resp.Status, tt.wantStatus, resp)
+ }
+ if resp.TradeNo != tt.wantTradeNo {
+ t.Fatalf("trade_no = %q, want %q", resp.TradeNo, tt.wantTradeNo)
+ }
+ if resp.Amount != tt.wantAmount {
+ t.Fatalf("amount = %v, want %v", resp.Amount, tt.wantAmount)
+ }
+ for key, want := range map[string]string{
+ "act": "order",
+ "pid": "pid-1",
+ "key": "pkey-1",
+ "out_trade_no": orderID,
+ } {
+ if got := gotForm.Get(key); got != want {
+ t.Fatalf("form[%s] = %q, want %q (form=%v)", key, got, want, gotForm)
+ }
+ }
+ })
+ }
+}
diff --git a/backend/internal/pkg/antigravity/claude_types.go b/backend/internal/pkg/antigravity/claude_types.go
index 0b8ae5f2..b651db94 100644
--- a/backend/internal/pkg/antigravity/claude_types.go
+++ b/backend/internal/pkg/antigravity/claude_types.go
@@ -155,6 +155,7 @@ var claudeModels = []modelDef{
{ID: "claude-opus-4-6", DisplayName: "Claude Opus 4.6", CreatedAt: "2026-02-05T00:00:00Z"},
{ID: "claude-opus-4-6-thinking", DisplayName: "Claude Opus 4.6 Thinking", CreatedAt: "2026-02-05T00:00:00Z"},
{ID: "claude-opus-4-7", DisplayName: "Claude Opus 4.7", CreatedAt: "2026-04-17T00:00:00Z"},
+ {ID: "claude-opus-4-8", DisplayName: "Claude Opus 4.8", CreatedAt: "2026-05-29T00:00:00Z"},
{ID: "claude-sonnet-4-6", DisplayName: "Claude Sonnet 4.6", CreatedAt: "2026-02-17T00:00:00Z"},
}
diff --git a/backend/internal/pkg/antigravity/claude_types_test.go b/backend/internal/pkg/antigravity/claude_types_test.go
index 9fc09b1b..fdf5c66d 100644
--- a/backend/internal/pkg/antigravity/claude_types_test.go
+++ b/backend/internal/pkg/antigravity/claude_types_test.go
@@ -12,6 +12,7 @@ func TestDefaultModels_ContainsNewAndLegacyImageModels(t *testing.T) {
}
requiredIDs := []string{
+ "claude-opus-4-8",
"claude-opus-4-6-thinking",
"gemini-2.5-flash-image",
"gemini-2.5-flash-image-preview",
diff --git a/backend/internal/pkg/antigravity/request_transformer.go b/backend/internal/pkg/antigravity/request_transformer.go
index b5de8166..9068ad97 100644
--- a/backend/internal/pkg/antigravity/request_transformer.go
+++ b/backend/internal/pkg/antigravity/request_transformer.go
@@ -204,6 +204,8 @@ type modelInfo struct {
// 只有在此映射表中的模型才会注入身份提示词
// 注意:模型映射逻辑在网关层完成;这里仅用于按模型前缀判断是否注入身份提示词。
var modelInfoMap = map[string]modelInfo{
+ "claude-opus-4-8": {DisplayName: "Claude Opus 4.8", CanonicalID: "claude-opus-4-8"},
+ "claude-opus-4-7": {DisplayName: "Claude Opus 4.7", CanonicalID: "claude-opus-4-7"},
"claude-opus-4-5": {DisplayName: "Claude Opus 4.5", CanonicalID: "claude-opus-4-5-20250929"},
"claude-opus-4-6": {DisplayName: "Claude Opus 4.6", CanonicalID: "claude-opus-4-6"},
"claude-sonnet-4-6": {DisplayName: "Claude Sonnet 4.6", CanonicalID: "claude-sonnet-4-6"},
@@ -587,7 +589,8 @@ func maxOutputTokensLimit(model string) int {
func isAntigravityOpusHighTierModel(model string) bool {
lower := strings.ToLower(model)
return strings.HasPrefix(lower, "claude-opus-4-6") ||
- strings.HasPrefix(lower, "claude-opus-4-7")
+ strings.HasPrefix(lower, "claude-opus-4-7") ||
+ strings.HasPrefix(lower, "claude-opus-4-8")
}
func buildGenerationConfig(req *ClaudeRequest) *GeminiGenerationConfig {
diff --git a/backend/internal/pkg/apicompat/anthropic_responses_test.go b/backend/internal/pkg/apicompat/anthropic_responses_test.go
index bb566081..8997835c 100644
--- a/backend/internal/pkg/apicompat/anthropic_responses_test.go
+++ b/backend/internal/pkg/apicompat/anthropic_responses_test.go
@@ -1597,3 +1597,139 @@ func TestAnthropicToResponses_TemperatureStrippedForAllGpt5Variants(t *testing.T
})
}
}
+
+// ---------------------------------------------------------------------------
+// AnthropicToResponsesResponse: Anthropic input_tokens excludes cached tokens
+// while OpenAI Responses input_tokens is the total including cached tokens.
+// ---------------------------------------------------------------------------
+
+func TestAnthropicToResponsesResponse_CacheTokensUseOpenAIInputSemantics(t *testing.T) {
+ resp := &AnthropicResponse{
+ ID: "msg_cache",
+ Model: "claude-sonnet-4-5-20250929",
+ Content: []AnthropicContentBlock{
+ {Type: "text", Text: "ok"},
+ },
+ StopReason: "end_turn",
+ Usage: AnthropicUsage{
+ InputTokens: 3318,
+ OutputTokens: 123,
+ CacheReadInputTokens: 50688,
+ CacheCreationInputTokens: 200,
+ },
+ }
+
+ out := AnthropicToResponsesResponse(resp)
+ require.NotNil(t, out.Usage)
+ // 3318 (uncached) + 50688 (read) + 200 (creation) = 54206
+ assert.Equal(t, 54206, out.Usage.InputTokens)
+ assert.Equal(t, 123, out.Usage.OutputTokens)
+ assert.Equal(t, 54329, out.Usage.TotalTokens)
+ require.NotNil(t, out.Usage.InputTokensDetails)
+ assert.Equal(t, 50688, out.Usage.InputTokensDetails.CachedTokens)
+}
+
+func TestAnthropicToResponsesResponse_NoCacheTokens(t *testing.T) {
+ resp := &AnthropicResponse{
+ ID: "msg_nocache",
+ Model: "claude-sonnet-4-5-20250929",
+ Content: []AnthropicContentBlock{
+ {Type: "text", Text: "ok"},
+ },
+ StopReason: "end_turn",
+ Usage: AnthropicUsage{
+ InputTokens: 100,
+ OutputTokens: 50,
+ },
+ }
+
+ out := AnthropicToResponsesResponse(resp)
+ require.NotNil(t, out.Usage)
+ assert.Equal(t, 100, out.Usage.InputTokens)
+ assert.Equal(t, 50, out.Usage.OutputTokens)
+ assert.Equal(t, 150, out.Usage.TotalTokens)
+ assert.Nil(t, out.Usage.InputTokensDetails)
+}
+
+func TestAnthropicEventToResponses_CacheTokensRoundTripFromMessageStart(t *testing.T) {
+ state := NewAnthropicEventToResponsesState()
+
+ // message_start carries cache fields on the initial Usage object.
+ AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
+ Type: "message_start",
+ Message: &AnthropicResponse{
+ ID: "msg_stream_cache",
+ Model: "claude-sonnet-4-5-20250929",
+ Usage: AnthropicUsage{
+ InputTokens: 12,
+ CacheReadInputTokens: 9,
+ CacheCreationInputTokens: 3,
+ },
+ },
+ }, state)
+
+ AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
+ Type: "message_delta",
+ Usage: &AnthropicUsage{
+ OutputTokens: 7,
+ },
+ }, state)
+
+ events := AnthropicEventToResponsesEvents(&AnthropicStreamEvent{Type: "message_stop"}, state)
+
+ // The terminal response.completed event must include OpenAI-semantic usage.
+ var completed *ResponsesStreamEvent
+ for i := range events {
+ if events[i].Type == "response.completed" {
+ completed = &events[i]
+ }
+ }
+ require.NotNil(t, completed, "response.completed event must be emitted")
+ require.NotNil(t, completed.Response)
+ require.NotNil(t, completed.Response.Usage)
+ // 12 (uncached) + 9 (read) + 3 (creation) = 24
+ assert.Equal(t, 24, completed.Response.Usage.InputTokens)
+ assert.Equal(t, 7, completed.Response.Usage.OutputTokens)
+ assert.Equal(t, 31, completed.Response.Usage.TotalTokens)
+ require.NotNil(t, completed.Response.Usage.InputTokensDetails)
+ assert.Equal(t, 9, completed.Response.Usage.InputTokensDetails.CachedTokens)
+}
+
+func TestAnthropicEventToResponses_CacheTokensFromMessageDelta(t *testing.T) {
+ state := NewAnthropicEventToResponsesState()
+
+ AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
+ Type: "message_start",
+ Message: &AnthropicResponse{
+ ID: "msg_delta_cache",
+ Model: "claude-sonnet-4-5-20250929",
+ Usage: AnthropicUsage{InputTokens: 20},
+ },
+ }, state)
+
+ // Some upstreams only emit cache fields on the final message_delta.
+ AnthropicEventToResponsesEvents(&AnthropicStreamEvent{
+ Type: "message_delta",
+ Usage: &AnthropicUsage{
+ OutputTokens: 8,
+ CacheReadInputTokens: 11,
+ CacheCreationInputTokens: 4,
+ },
+ }, state)
+
+ events := AnthropicEventToResponsesEvents(&AnthropicStreamEvent{Type: "message_stop"}, state)
+
+ var completed *ResponsesStreamEvent
+ for i := range events {
+ if events[i].Type == "response.completed" {
+ completed = &events[i]
+ }
+ }
+ require.NotNil(t, completed)
+ require.NotNil(t, completed.Response.Usage)
+ // 20 (uncached) + 11 (read) + 4 (creation) = 35
+ assert.Equal(t, 35, completed.Response.Usage.InputTokens)
+ assert.Equal(t, 8, completed.Response.Usage.OutputTokens)
+ require.NotNil(t, completed.Response.Usage.InputTokensDetails)
+ assert.Equal(t, 11, completed.Response.Usage.InputTokensDetails.CachedTokens)
+}
diff --git a/backend/internal/pkg/apicompat/anthropic_to_responses_response.go b/backend/internal/pkg/apicompat/anthropic_to_responses_response.go
index 9290e399..de8ab78d 100644
--- a/backend/internal/pkg/apicompat/anthropic_to_responses_response.go
+++ b/backend/internal/pkg/apicompat/anthropic_to_responses_response.go
@@ -95,10 +95,16 @@ func AnthropicToResponsesResponse(resp *AnthropicResponse) *ResponsesResponse {
}
// Usage
+ // Anthropic's input_tokens excludes cache_read/cache_creation, while OpenAI
+ // Responses' input_tokens is the total including cached tokens. Add them back
+ // when converting so downstream consumers see OpenAI semantics.
+ totalInputTokens := resp.Usage.InputTokens +
+ resp.Usage.CacheReadInputTokens +
+ resp.Usage.CacheCreationInputTokens
out.Usage = &ResponsesUsage{
- InputTokens: resp.Usage.InputTokens,
+ InputTokens: totalInputTokens,
OutputTokens: resp.Usage.OutputTokens,
- TotalTokens: resp.Usage.InputTokens + resp.Usage.OutputTokens,
+ TotalTokens: totalInputTokens + resp.Usage.OutputTokens,
}
if resp.Usage.CacheReadInputTokens > 0 {
out.Usage.InputTokensDetails = &ResponsesInputTokensDetails{
@@ -150,10 +156,13 @@ type AnthropicEventToResponsesState struct {
CurrentCallID string
CurrentName string
- // Usage from message_delta
- InputTokens int
- OutputTokens int
- CacheReadInputTokens int
+ // Usage from message_start / message_delta. InputTokens here follows
+ // Anthropic semantics (excludes cached tokens); they are added back when
+ // emitting the OpenAI Responses usage.
+ InputTokens int
+ OutputTokens int
+ CacheReadInputTokens int
+ CacheCreationInputTokens int
}
// NewAnthropicEventToResponsesState returns an initialised stream state.
@@ -225,6 +234,12 @@ func anthToResHandleMessageStart(evt *AnthropicStreamEvent, state *AnthropicEven
if evt.Message.Usage.InputTokens > 0 {
state.InputTokens = evt.Message.Usage.InputTokens
}
+ if evt.Message.Usage.CacheReadInputTokens > 0 {
+ state.CacheReadInputTokens = evt.Message.Usage.CacheReadInputTokens
+ }
+ if evt.Message.Usage.CacheCreationInputTokens > 0 {
+ state.CacheCreationInputTokens = evt.Message.Usage.CacheCreationInputTokens
+ }
}
if state.CreatedSent {
@@ -392,9 +407,15 @@ func anthToResHandleMessageDelta(evt *AnthropicStreamEvent, state *AnthropicEven
// Update usage
if evt.Usage != nil {
state.OutputTokens = evt.Usage.OutputTokens
+ if evt.Usage.InputTokens > 0 {
+ state.InputTokens = evt.Usage.InputTokens
+ }
if evt.Usage.CacheReadInputTokens > 0 {
state.CacheReadInputTokens = evt.Usage.CacheReadInputTokens
}
+ if evt.Usage.CacheCreationInputTokens > 0 {
+ state.CacheCreationInputTokens = evt.Usage.CacheCreationInputTokens
+ }
}
return nil
@@ -472,10 +493,13 @@ func makeResponsesCompletedEvent(
seq := state.SequenceNumber
state.SequenceNumber++
+ // Anthropic's input_tokens excludes cache_read/cache_creation; add them
+ // back to match OpenAI Responses semantics where input_tokens is the total.
+ totalInputTokens := state.InputTokens + state.CacheReadInputTokens + state.CacheCreationInputTokens
usage := &ResponsesUsage{
- InputTokens: state.InputTokens,
+ InputTokens: totalInputTokens,
OutputTokens: state.OutputTokens,
- TotalTokens: state.InputTokens + state.OutputTokens,
+ TotalTokens: totalInputTokens + state.OutputTokens,
}
if state.CacheReadInputTokens > 0 {
usage.InputTokensDetails = &ResponsesInputTokensDetails{
diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go
index 09b680c7..cc51cba2 100644
--- a/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go
+++ b/backend/internal/pkg/apicompat/chatcompletions_responses_bridge.go
@@ -42,14 +42,24 @@ func ResponsesToChatCompletionsRequest(req *ResponsesRequest) (*ChatCompletionsR
return out, nil
}
+// responsesInputToChatMessages converts a Responses request's instructions +
+// input[] into Chat Completions messages. It is a three-stage pipeline:
+//
+// parse — instructions become a system message; input[] is split into items
+// build — buildChatMessagesFromItems walks items, attaching reasoning to the
+// assistant message that produced a tool call, merging parallel tool
+// calls into one assistant message, and skipping item types that have
+// no Chat equivalent
+// normalize — normalizeChatMessages enforces the invariants DeepSeek requires
+//
+// The build + normalize split keeps every protocol rule in one place rather than
+// scattered across per-item cases, and makes unknown future codex item types
+// fail safe instead of leaking into the upstream request.
func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage) ([]ChatMessage, error) {
var messages []ChatMessage
if strings.TrimSpace(instructions) != "" {
content, _ := json.Marshal(instructions)
- messages = append(messages, ChatMessage{
- Role: "system",
- Content: content,
- })
+ messages = append(messages, ChatMessage{Role: "system", Content: content})
}
inputRaw = bytesTrimSpace(inputRaw)
@@ -57,13 +67,11 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
return messages, nil
}
+ // Bare string input is a single user turn.
var inputText string
if err := json.Unmarshal(inputRaw, &inputText); err == nil {
content, _ := json.Marshal(inputText)
- messages = append(messages, ChatMessage{
- Role: "user",
- Content: content,
- })
+ messages = append(messages, ChatMessage{Role: "user", Content: content})
return messages, nil
}
@@ -72,6 +80,24 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
return nil, fmt.Errorf("parse responses input: %w", err)
}
+ built, err := buildChatMessagesFromItems(messages, rawItems)
+ if err != nil {
+ return nil, err
+ }
+ return normalizeChatMessages(built), nil
+}
+
+// buildChatMessagesFromItems walks the Responses input items and appends the
+// corresponding Chat messages.
+func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessage) ([]ChatMessage, error) {
+ // pendingReasoning holds the reasoning text from a reasoning item until the
+ // assistant message it belongs to is emitted. DeepSeek's thinking mode
+ // requires the reasoning_content that produced a tool call to be passed back
+ // on that assistant message; dropping it yields a 400. It only survives
+ // across an assistant message (so a following tool call in the same turn
+ // still receives it); any other role ends the thinking span.
+ var pendingReasoning string
+
for _, raw := range rawItems {
raw = bytesTrimSpace(raw)
if len(raw) == 0 || string(raw) == "null" {
@@ -84,6 +110,7 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
if textErr := json.Unmarshal(raw, &text); textErr == nil {
content, _ := json.Marshal(text)
messages = append(messages, ChatMessage{Role: "user", Content: content})
+ pendingReasoning = ""
continue
}
return nil, fmt.Errorf("parse responses input item: %w", err)
@@ -92,22 +119,40 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
role := chatCompletionsBridgeRole(rawString(item["role"]))
itemType := rawString(item["type"])
switch itemType {
+ case "reasoning":
+ if txt := extractResponsesReasoningText(item); txt != "" {
+ pendingReasoning = txt
+ }
+ continue
case "function_call":
arguments := rawString(item["arguments"])
if strings.TrimSpace(arguments) == "" {
arguments = "{}"
}
- messages = append(messages, ChatMessage{
- Role: "assistant",
- ToolCalls: []ChatToolCall{{
- ID: rawString(item["call_id"]),
- Type: "function",
- Function: ChatFunctionCall{
- Name: rawString(item["name"]),
- Arguments: arguments,
- },
- }},
- })
+ toolCall := ChatToolCall{
+ ID: rawString(item["call_id"]),
+ Type: "function",
+ Function: ChatFunctionCall{
+ Name: rawString(item["name"]),
+ Arguments: arguments,
+ },
+ }
+ // Parallel tool calls arrive as consecutive function_call items and
+ // must share one assistant message; the matching tool replies then
+ // follow it. Merge into the immediately preceding assistant message.
+ if n := len(messages); n > 0 && messages[n-1].Role == "assistant" {
+ messages[n-1].ToolCalls = append(messages[n-1].ToolCalls, toolCall)
+ if messages[n-1].ReasoningContent == "" {
+ messages[n-1].ReasoningContent = pendingReasoning
+ }
+ } else {
+ messages = append(messages, ChatMessage{
+ Role: "assistant",
+ ToolCalls: []ChatToolCall{toolCall},
+ ReasoningContent: pendingReasoning,
+ })
+ }
+ pendingReasoning = ""
continue
case "function_call_output":
content, _ := json.Marshal(rawString(item["output"]))
@@ -116,10 +161,12 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
ToolCallID: rawString(item["call_id"]),
Content: content,
})
+ pendingReasoning = ""
continue
case "input_text", "text":
content, _ := json.Marshal(rawString(item["text"]))
messages = append(messages, ChatMessage{Role: "user", Content: content})
+ pendingReasoning = ""
continue
case "input_image":
content, err := chatContentFromSingleResponsesPart(itemType, item)
@@ -127,6 +174,18 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
return nil, err
}
messages = append(messages, ChatMessage{Role: "user", Content: content})
+ pendingReasoning = ""
+ continue
+ }
+
+ // Only genuine message items become chat messages. Codex emits other
+ // Responses item types with no Chat equivalent (web_search_call,
+ // local_shell_call, custom tool calls, file_search_call, ...). Converting
+ // them via the generic path would insert a spurious message between an
+ // assistant tool_calls message and its tool reply, which DeepSeek rejects
+ // ("insufficient tool messages following tool_calls message"). Skip them.
+ if itemType != "" && itemType != "message" {
+ pendingReasoning = ""
continue
}
@@ -140,15 +199,128 @@ func responsesInputToChatMessages(instructions string, inputRaw json.RawMessage)
if err != nil {
return nil, err
}
- messages = append(messages, ChatMessage{
- Role: role,
- Content: chatContent,
- })
+ messages = append(messages, ChatMessage{Role: role, Content: chatContent})
+ // Reasoning only survives across an assistant text message.
+ if role != "assistant" {
+ pendingReasoning = ""
+ }
}
return messages, nil
}
+// normalizeChatMessages is the single place that enforces the tool-call
+// invariant the DeepSeek / OpenAI Chat Completions schema requires: an assistant
+// message with tool_calls must be immediately followed by one tool message per
+// tool_call_id, in order, with nothing in between.
+//
+// Codex histories violate this in several ways that the builder alone can't fix:
+// - a non-tool message lands between an assistant tool_calls message and its
+// tool replies (e.g. an "Approved command prefix saved" system notice codex
+// injects mid tool-execution);
+// - a parallel tool_call's sibling output never arrives, or a call is left
+// dangling by a mid-execution reconnect (unanswered tool_call);
+// - a tool reply has no announcing assistant tool_call (orphan).
+//
+// It rebuilds the sequence so each assistant's answered tool_calls are followed
+// directly by their replies (in call order); unanswered tool_calls are dropped
+// (and an assistant left with neither tool_calls nor content is dropped); orphan
+// tool replies and intervening messages are emitted in their natural position
+// but never between an assistant tool_calls message and its replies.
+func normalizeChatMessages(messages []ChatMessage) []ChatMessage {
+ // Index every tool reply by its tool_call_id (last wins on duplicates).
+ replies := make(map[string]ChatMessage)
+ for _, m := range messages {
+ if m.Role == "tool" && m.ToolCallID != "" {
+ replies[m.ToolCallID] = m
+ }
+ }
+
+ out := make([]ChatMessage, 0, len(messages))
+ for _, m := range messages {
+ switch {
+ case m.Role == "tool":
+ // A bare tool message with no tool_call_id is a direct Chat
+ // Completions passthrough; keep it in place. A tool reply whose id is
+ // announced by an assistant is emitted right after that assistant
+ // (skip the standalone occurrence). Any other tool reply is an orphan
+ // and is dropped.
+ if m.ToolCallID == "" {
+ out = append(out, m)
+ }
+ continue
+ case len(m.ToolCalls) > 0:
+ kept := make([]ChatToolCall, 0, len(m.ToolCalls))
+ for _, tc := range m.ToolCalls {
+ if tc.ID == "" {
+ continue
+ }
+ if _, ok := replies[tc.ID]; ok {
+ kept = append(kept, tc)
+ }
+ }
+ if len(kept) == 0 {
+ // No answered tool_calls left: keep as a plain message if it has
+ // content, otherwise drop it entirely.
+ if isBlankChatContent(m.Content) {
+ continue
+ }
+ m.ToolCalls = nil
+ out = append(out, m)
+ continue
+ }
+ m.ToolCalls = kept
+ out = append(out, m)
+ for _, tc := range kept {
+ out = append(out, replies[tc.ID])
+ }
+ default:
+ out = append(out, m)
+ }
+ }
+ return out
+}
+
+// isBlankChatContent reports whether a chat message content holds no usable text.
+func isBlankChatContent(raw json.RawMessage) bool {
+ raw = bytesTrimSpace(raw)
+ if len(raw) == 0 || string(raw) == "null" || string(raw) == `""` {
+ return true
+ }
+ return chatMessageContentText(raw) == ""
+}
+
+// extractResponsesReasoningText pulls the reasoning text out of a Responses
+// reasoning item. The Chat→Responses bridge writes the upstream reasoning_content
+// verbatim into the summary_text parts (see closeChatReasoningItem), so codex
+// round-trips it there; prefer summary[].text and fall back to content.
+func extractResponsesReasoningText(item map[string]json.RawMessage) string {
+ var parts []string
+ collect := func(raw json.RawMessage) {
+ raw = bytesTrimSpace(raw)
+ if len(raw) == 0 || string(raw) == "null" {
+ return
+ }
+ var arr []map[string]json.RawMessage
+ if err := json.Unmarshal(raw, &arr); err == nil {
+ for _, p := range arr {
+ if t := rawString(p["text"]); t != "" {
+ parts = append(parts, t)
+ }
+ }
+ return
+ }
+ if t := rawString(raw); t != "" {
+ parts = append(parts, t)
+ }
+ }
+ collect(item["summary"])
+ if len(parts) == 0 {
+ collect(item["content"])
+ }
+ return strings.Join(parts, "\n")
+}
+
func chatCompletionsBridgeRole(role string) string {
trimmed := strings.TrimSpace(role)
if trimmed == "" {
@@ -448,10 +620,32 @@ type ChatCompletionsToResponsesStreamState struct {
CreatedSent bool
CompletedSent bool
+ // nextOutputIndex assigns sequential output_index values to items as they
+ // are opened (reasoning, message, tool calls), so the streamed indices match
+ // the order of items in the final response.output array.
+ nextOutputIndex int
+
+ // Reasoning item lifecycle. DeepSeek-style upstreams stream all
+ // reasoning_content before any content, so reasoning is modeled as its own
+ // "reasoning" output item that must be opened (output_item.added) before any
+ // reasoning delta and closed before the message/tool items open.
+ ReasoningItemID string
+ ReasoningIndex int
+ ReasoningOpen bool
+ ReasoningDone bool
+
+ // Message item + output_text content-part lifecycle.
MessageItemID string
- Text strings.Builder
- Reasoning strings.Builder
- ToolCalls map[int]*ChatToolCall
+ MessageIndex int
+ TextPartOpen bool
+
+ Text strings.Builder
+ Reasoning strings.Builder
+
+ // Tool-call lifecycle, keyed by the upstream tool_call index.
+ ToolCalls map[int]*ChatToolCall
+ ToolItemIDs map[int]string
+ ToolOutputIndex map[int]int
FinishReason string
Usage *ResponsesUsage
@@ -460,13 +654,21 @@ type ChatCompletionsToResponsesStreamState struct {
// NewChatCompletionsToResponsesStreamState returns an initialized stream state.
func NewChatCompletionsToResponsesStreamState(model string) *ChatCompletionsToResponsesStreamState {
return &ChatCompletionsToResponsesStreamState{
- ResponseID: generateResponsesID(),
- Model: model,
- Created: time.Now().Unix(),
- ToolCalls: make(map[int]*ChatToolCall),
+ ResponseID: generateResponsesID(),
+ Model: model,
+ Created: time.Now().Unix(),
+ ToolCalls: make(map[int]*ChatToolCall),
+ ToolItemIDs: make(map[int]string),
+ ToolOutputIndex: make(map[int]int),
}
}
+func (state *ChatCompletionsToResponsesStreamState) allocOutputIndex() int {
+ idx := state.nextOutputIndex
+ state.nextOutputIndex++
+ return idx
+}
+
// ChatCompletionsChunkToResponsesEvents converts one Chat Completions stream
// chunk into zero or more Responses stream events.
func ChatCompletionsChunkToResponsesEvents(
@@ -490,24 +692,34 @@ func ChatCompletionsChunkToResponsesEvents(
events = append(events, ensureChatToResponsesCreated(state)...)
for _, choice := range chunk.Choices {
- if choice.Delta.Content != nil {
+ // Reasoning is emitted as its own output item and must be opened
+ // (output_item.added + reasoning_summary_part.added) before the first
+ // delta, otherwise a strict client discards the delta. The leading
+ // empty-string reasoning delta upstreams send is filtered out.
+ if choice.Delta.ReasoningContent != nil && *choice.Delta.ReasoningContent != "" {
+ events = append(events, ensureChatReasoningItem(state)...)
+ _, _ = state.Reasoning.WriteString(*choice.Delta.ReasoningContent)
+ events = append(events, chatToResponsesEvent(state, "response.reasoning_summary_text.delta", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ SummaryIndex: 0,
+ Delta: *choice.Delta.ReasoningContent,
+ ItemID: state.ReasoningItemID,
+ }))
+ }
+ if choice.Delta.Content != nil && *choice.Delta.Content != "" {
+ // First real content closes the reasoning item, then opens the
+ // message item and its output_text content part.
+ events = append(events, closeChatReasoningItem(state)...)
events = append(events, ensureChatToResponsesMessageItem(state)...)
+ events = append(events, ensureChatToResponsesTextPart(state)...)
_, _ = state.Text.WriteString(*choice.Delta.Content)
events = append(events, chatToResponsesEvent(state, "response.output_text.delta", &ResponsesStreamEvent{
- OutputIndex: 0,
+ OutputIndex: state.MessageIndex,
ContentIndex: 0,
Delta: *choice.Delta.Content,
ItemID: state.MessageItemID,
}))
}
- if choice.Delta.ReasoningContent != nil {
- _, _ = state.Reasoning.WriteString(*choice.Delta.ReasoningContent)
- events = append(events, chatToResponsesEvent(state, "response.reasoning_summary_text.delta", &ResponsesStreamEvent{
- OutputIndex: 0,
- SummaryIndex: 0,
- Delta: *choice.Delta.ReasoningContent,
- }))
- }
for _, toolCall := range choice.Delta.ToolCalls {
idx := 0
if toolCall.Index != nil {
@@ -515,6 +727,8 @@ func ChatCompletionsChunkToResponsesEvents(
}
stored, ok := state.ToolCalls[idx]
if !ok {
+ // A tool call closes any open reasoning item first.
+ events = append(events, closeChatReasoningItem(state)...)
copyCall := toolCall
if copyCall.ID == "" {
copyCall.ID = generateItemID()
@@ -522,11 +736,14 @@ func ChatCompletionsChunkToResponsesEvents(
copyCall.Type = "function"
state.ToolCalls[idx] = ©Call
stored = ©Call
+ itemID := generateItemID()
+ state.ToolItemIDs[idx] = itemID
+ state.ToolOutputIndex[idx] = state.allocOutputIndex()
events = append(events, chatToResponsesEvent(state, "response.output_item.added", &ResponsesStreamEvent{
- OutputIndex: idx + 1,
+ OutputIndex: state.ToolOutputIndex[idx],
Item: &ResponsesOutput{
Type: "function_call",
- ID: generateItemID(),
+ ID: itemID,
CallID: stored.ID,
Name: stored.Function.Name,
Status: "in_progress",
@@ -543,7 +760,8 @@ func ChatCompletionsChunkToResponsesEvents(
if toolCall.Function.Arguments != "" {
stored.Function.Arguments += toolCall.Function.Arguments
events = append(events, chatToResponsesEvent(state, "response.function_call_arguments.delta", &ResponsesStreamEvent{
- OutputIndex: idx + 1,
+ OutputIndex: state.ToolOutputIndex[idx],
+ ItemID: state.ToolItemIDs[idx],
Delta: toolCall.Function.Arguments,
CallID: stored.ID,
Name: stored.Function.Name,
@@ -565,24 +783,44 @@ func FinalizeChatCompletionsResponsesStream(state *ChatCompletionsToResponsesStr
}
var events []ResponsesStreamEvent
events = append(events, ensureChatToResponsesCreated(state)...)
+
+ // Close a reasoning item that never transitioned to content (reasoning-only
+ // or empty completion).
+ events = append(events, closeChatReasoningItem(state)...)
+
if state.MessageItemID != "" {
- events = append(events, chatToResponsesEvent(state, "response.output_text.done", &ResponsesStreamEvent{
- OutputIndex: 0,
- ContentIndex: 0,
- Text: state.Text.String(),
- ItemID: state.MessageItemID,
- }))
+ if state.TextPartOpen {
+ events = append(events, chatToResponsesEvent(state, "response.output_text.done", &ResponsesStreamEvent{
+ OutputIndex: state.MessageIndex,
+ ContentIndex: 0,
+ Text: state.Text.String(),
+ ItemID: state.MessageItemID,
+ }))
+ events = append(events, chatToResponsesEvent(state, "response.content_part.done", &ResponsesStreamEvent{
+ OutputIndex: state.MessageIndex,
+ ContentIndex: 0,
+ ItemID: state.MessageItemID,
+ Part: &ResponsesContentPart{Type: "output_text", Text: state.Text.String()},
+ }))
+ }
events = append(events, chatToResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{
- OutputIndex: 0,
+ OutputIndex: state.MessageIndex,
Item: &ResponsesOutput{
- Type: "message",
- ID: state.MessageItemID,
- Role: "assistant",
- Status: "completed",
+ Type: "message",
+ ID: state.MessageItemID,
+ Role: "assistant",
+ Content: []ResponsesContentPart{{Type: "output_text", Text: state.Text.String()}},
+ Status: "completed",
},
}))
}
+ // Close every function_call item opened during the stream. Codex finalizes a
+ // tool call only after function_call_arguments.done + output_item.done for
+ // that item; without them the call never completes and the session wedges.
+ // Mirrors cc-switch's finalize_tools.
+ events = append(events, closeChatToolItems(state)...)
+
status := "completed"
var incompleteDetails *ResponsesIncompleteDetails
if state.FinishReason == "length" {
@@ -621,22 +859,142 @@ func ensureChatToResponsesCreated(state *ChatCompletionsToResponsesStreamState)
})}
}
+// ensureChatReasoningItem opens the reasoning output item (output_item.added +
+// reasoning_summary_part.added) before the first reasoning delta. Codex renders
+// streaming reasoning only when this summary-part lifecycle is present.
+func ensureChatReasoningItem(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
+ if state.ReasoningOpen || state.ReasoningDone {
+ return nil
+ }
+ state.ReasoningOpen = true
+ state.ReasoningItemID = generateItemID()
+ state.ReasoningIndex = state.allocOutputIndex()
+ return []ResponsesStreamEvent{
+ chatToResponsesEvent(state, "response.output_item.added", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ Item: &ResponsesOutput{Type: "reasoning", ID: state.ReasoningItemID, Status: "in_progress"},
+ }),
+ chatToResponsesEvent(state, "response.reasoning_summary_part.added", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ SummaryIndex: 0,
+ ItemID: state.ReasoningItemID,
+ Part: &ResponsesContentPart{Type: "summary_text"},
+ }),
+ }
+}
+
+// closeChatReasoningItem emits the reasoning item's terminal events
+// (reasoning_summary_text.done + reasoning_summary_part.done + output_item.done).
+func closeChatReasoningItem(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
+ if !state.ReasoningOpen {
+ return nil
+ }
+ state.ReasoningOpen = false
+ state.ReasoningDone = true
+ reasoning := state.Reasoning.String()
+ return []ResponsesStreamEvent{
+ chatToResponsesEvent(state, "response.reasoning_summary_text.done", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ SummaryIndex: 0,
+ Text: reasoning,
+ ItemID: state.ReasoningItemID,
+ }),
+ chatToResponsesEvent(state, "response.reasoning_summary_part.done", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ SummaryIndex: 0,
+ ItemID: state.ReasoningItemID,
+ Part: &ResponsesContentPart{Type: "summary_text", Text: reasoning},
+ }),
+ chatToResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{
+ OutputIndex: state.ReasoningIndex,
+ Item: &ResponsesOutput{
+ Type: "reasoning",
+ ID: state.ReasoningItemID,
+ Status: "completed",
+ Summary: []ResponsesSummary{{Type: "summary_text", Text: reasoning}},
+ },
+ }),
+ }
+}
+
func ensureChatToResponsesMessageItem(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
if state.MessageItemID != "" {
return nil
}
state.MessageItemID = generateItemID()
+ state.MessageIndex = state.allocOutputIndex()
return []ResponsesStreamEvent{chatToResponsesEvent(state, "response.output_item.added", &ResponsesStreamEvent{
- OutputIndex: 0,
+ OutputIndex: state.MessageIndex,
Item: &ResponsesOutput{
- Type: "message",
- ID: state.MessageItemID,
- Role: "assistant",
- Status: "in_progress",
+ Type: "message",
+ ID: state.MessageItemID,
+ Role: "assistant",
+ Status: "in_progress",
+ Content: []ResponsesContentPart{{Type: "output_text"}},
},
})}
}
+func ensureChatToResponsesTextPart(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
+ if state.TextPartOpen {
+ return nil
+ }
+ state.TextPartOpen = true
+ return []ResponsesStreamEvent{chatToResponsesEvent(state, "response.content_part.added", &ResponsesStreamEvent{
+ OutputIndex: state.MessageIndex,
+ ContentIndex: 0,
+ ItemID: state.MessageItemID,
+ Part: &ResponsesContentPart{Type: "output_text", Text: ""},
+ })}
+}
+
+// closeChatToolItems emits function_call_arguments.done + output_item.done for
+// every tool call opened during the stream, carrying the full call_id/name/
+// arguments so codex can deserialize and execute the call. Mirrors cc-switch's
+// finalize_tools.
+func closeChatToolItems(state *ChatCompletionsToResponsesStreamState) []ResponsesStreamEvent {
+ if len(state.ToolCalls) == 0 {
+ return nil
+ }
+ var events []ResponsesStreamEvent
+ for i := 0; i < len(state.ToolCalls); i++ {
+ toolCall, ok := state.ToolCalls[i]
+ if !ok || toolCall == nil {
+ continue
+ }
+ itemID, opened := state.ToolItemIDs[i]
+ if !opened {
+ continue
+ }
+ arguments := toolCall.Function.Arguments
+ if strings.TrimSpace(arguments) == "" {
+ arguments = "{}"
+ }
+ outputIndex := state.ToolOutputIndex[i]
+ events = append(events,
+ chatToResponsesEvent(state, "response.function_call_arguments.done", &ResponsesStreamEvent{
+ OutputIndex: outputIndex,
+ ItemID: itemID,
+ CallID: toolCall.ID,
+ Name: toolCall.Function.Name,
+ Arguments: arguments,
+ }),
+ chatToResponsesEvent(state, "response.output_item.done", &ResponsesStreamEvent{
+ OutputIndex: outputIndex,
+ Item: &ResponsesOutput{
+ Type: "function_call",
+ ID: itemID,
+ CallID: toolCall.ID,
+ Name: toolCall.Function.Name,
+ Arguments: arguments,
+ Status: "completed",
+ },
+ }),
+ )
+ }
+ return events
+}
+
func (state *ChatCompletionsToResponsesStreamState) chatOutput() []ResponsesOutput {
var outputs []ResponsesOutput
if state.Reasoning.Len() > 0 {
diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_request_invariants_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_request_invariants_test.go
new file mode 100644
index 00000000..e54a4532
--- /dev/null
+++ b/backend/internal/pkg/apicompat/chatcompletions_responses_request_invariants_test.go
@@ -0,0 +1,187 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+// assertChatInvariants enforces the DeepSeek / OpenAI Chat Completions message
+// invariants that, when violated, surface as upstream 400s. Used to validate the
+// request-direction converter against golden codex request shapes.
+func assertChatInvariants(t *testing.T, messages []ChatMessage) {
+ t.Helper()
+ for i, m := range messages {
+ // Every assistant tool_calls message must be immediately followed by one
+ // tool message per tool_call_id, in order.
+ if len(m.ToolCalls) > 0 {
+ for j, tc := range m.ToolCalls {
+ k := i + 1 + j
+ require.Lessf(t, k, len(messages), "tool_call %s has no following tool message", tc.ID)
+ require.Equalf(t, "tool", messages[k].Role, "tool_call %s not followed by a tool message", tc.ID)
+ require.Equalf(t, tc.ID, messages[k].ToolCallID, "tool reply order mismatch for %s", tc.ID)
+ }
+ }
+ // No two consecutive assistant messages.
+ if i > 0 && m.Role == "assistant" && messages[i-1].Role == "assistant" {
+ t.Fatalf("consecutive assistant messages at %d", i)
+ }
+ // No orphan tool replies.
+ if m.Role == "tool" {
+ require.NotEmptyf(t, m.ToolCallID, "tool message without tool_call_id at %d", i)
+ }
+ }
+}
+
+func convertGolden(t *testing.T, input string) []ChatMessage {
+ t.Helper()
+ msgs, err := responsesInputToChatMessages("You are a helpful assistant.", json.RawMessage(input))
+ require.NoError(t, err)
+ return msgs
+}
+
+// Golden sample: a single tool-call turn (codex runs one shell/curl command),
+// the shape that produced the original "no response" / 400.
+func TestGolden_SingleToolCall(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"latest sha?"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"need to run curl"}]},
+ {"type":"function_call","call_id":"call_a","name":"exec_command","arguments":"{\"cmd\":\"curl x\"}"},
+ {"type":"function_call_output","call_id":"call_a","output":"deadbeef"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ // reasoning_content must ride on the assistant tool-call message.
+ var asst *ChatMessage
+ for i := range msgs {
+ if len(msgs[i].ToolCalls) > 0 {
+ asst = &msgs[i]
+ }
+ }
+ require.NotNil(t, asst)
+ require.Equal(t, "need to run curl", asst.ReasoningContent)
+}
+
+// Golden sample: parallel tool calls (codex runs git log + git tag at once).
+func TestGolden_ParallelToolCalls(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"features?"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"inspect repo"}]},
+ {"type":"function_call","call_id":"c0","name":"exec_command","arguments":"{\"cmd\":\"git log\"}"},
+ {"type":"function_call","call_id":"c1","name":"exec_command","arguments":"{\"cmd\":\"git tag\"}"},
+ {"type":"function_call_output","call_id":"c0","output":"log"},
+ {"type":"function_call_output","call_id":"c1","output":"tags"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ // Both parallel calls share ONE assistant message.
+ var toolMsgs int
+ for _, m := range msgs {
+ if len(m.ToolCalls) == 2 {
+ require.Equal(t, "c0", m.ToolCalls[0].ID)
+ require.Equal(t, "c1", m.ToolCalls[1].ID)
+ }
+ if m.Role == "tool" {
+ toolMsgs++
+ }
+ }
+ require.Equal(t, 2, toolMsgs)
+}
+
+// Golden sample: an unknown item type (web_search_call from a 联网查询) sitting
+// between a function_call and its output must not break tool↔reply adjacency.
+func TestGolden_UnknownItemBetweenToolCallAndOutput(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"search"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"let me search"}]},
+ {"type":"function_call","call_id":"c0","name":"exec_command","arguments":"{}"},
+ {"type":"web_search_call","id":"ws_1","status":"completed","action":{"type":"search","query":"x"}},
+ {"type":"function_call_output","call_id":"c0","output":"result"}
+ ]`)
+ assertChatInvariants(t, msgs)
+}
+
+// Sequential tool calls (a tool reply between two calls) must stay in distinct
+// assistant messages.
+func TestRequest_SequentialToolCallsStaySeparate(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"function_call","call_id":"c1","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"c1","output":"r1"},
+ {"type":"function_call","call_id":"c2","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"c2","output":"r2"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ assistants := 0
+ for _, m := range msgs {
+ if len(m.ToolCalls) == 1 {
+ assistants++
+ }
+ }
+ require.Equal(t, 2, assistants)
+}
+
+// Golden sample: codex injects a message (e.g. an "Approved command prefix
+// saved" notice) between a function_call and its output. The intervening message
+// must be moved after the tool reply so the assistant tool_calls is immediately
+// followed by its reply.
+func TestGolden_MessageBetweenToolCallAndOutput(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"do it"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"run cmd"}]},
+ {"type":"function_call","call_id":"A","name":"exec","arguments":"{}"},
+ {"type":"message","role":"developer","content":[{"type":"input_text","text":"Approved command prefix saved"}]},
+ {"type":"function_call_output","call_id":"A","output":"ok"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ // The assistant tool_calls message is immediately followed by its tool reply.
+ for i, m := range msgs {
+ if len(m.ToolCalls) > 0 {
+ require.Equal(t, "tool", msgs[i+1].Role)
+ require.Equal(t, "A", msgs[i+1].ToolCallID)
+ }
+ }
+}
+
+// Golden sample: a parallel tool call where one sibling's output is missing
+// (codex interrupted/reconnected mid-execution). The unanswered tool_call must
+// be dropped so the remaining assistant tool_calls are all answered.
+func TestGolden_PartialParallelDropsUnansweredCall(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"r"}]},
+ {"type":"function_call","call_id":"A","name":"exec","arguments":"{}"},
+ {"type":"function_call","call_id":"B","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"A","output":"oa"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ for _, m := range msgs {
+ for _, tc := range m.ToolCalls {
+ require.NotEqual(t, "B", tc.ID, "unanswered tool_call B should have been dropped")
+ }
+ }
+}
+
+// Golden sample: a dangling tool_call at the end of the history (no output yet).
+// The assistant message holding only that call must be dropped entirely.
+func TestGolden_DanglingToolCallDropped(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"reasoning","summary":[{"type":"summary_text","text":"r"}]},
+ {"type":"function_call","call_id":"A","name":"exec","arguments":"{}"}
+ ]`)
+ assertChatInvariants(t, msgs)
+ for _, m := range msgs {
+ require.Empty(t, m.ToolCalls, "dangling unanswered tool_call should have been dropped")
+ }
+}
+
+// normalizeChatMessages drops an orphan tool reply whose tool_call was never
+// announced.
+func TestNormalize_DropsOrphanToolReply(t *testing.T) {
+ msgs := convertGolden(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"function_call_output","call_id":"ghost","output":"orphan"}
+ ]`)
+ for _, m := range msgs {
+ require.NotEqualf(t, "tool", m.Role, "orphan tool reply should have been dropped")
+ }
+}
diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_stream_lifecycle_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_stream_lifecycle_test.go
new file mode 100644
index 00000000..beb47303
--- /dev/null
+++ b/backend/internal/pkg/apicompat/chatcompletions_responses_stream_lifecycle_test.go
@@ -0,0 +1,103 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "strings"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func collectStreamEvents(t *testing.T, chunks []string) []ResponsesStreamEvent {
+ t.Helper()
+ state := NewChatCompletionsToResponsesStreamState("deepseek-v4-pro")
+ var events []ResponsesStreamEvent
+ for _, payload := range chunks {
+ var chunk ChatCompletionsChunk
+ require.NoError(t, json.Unmarshal([]byte(payload), &chunk))
+ events = append(events, ChatCompletionsChunkToResponsesEvents(&chunk, state)...)
+ }
+ events = append(events, FinalizeChatCompletionsResponsesStream(state)...)
+ return events
+}
+
+// TestStream_ReasoningOpensItemBeforeDelta guards the bug where a strict client
+// (Codex) drops reasoning deltas that reference an item not yet opened.
+func TestStream_ReasoningOpensItemBeforeDelta(t *testing.T) {
+ events := collectStreamEvents(t, []string{
+ `{"choices":[{"index":0,"delta":{"role":"assistant","content":null,"reasoning_content":""}}]}`,
+ `{"choices":[{"index":0,"delta":{"reasoning_content":"think"}}]}`,
+ `{"choices":[{"index":0,"delta":{"content":"hello"}}]}`,
+ `{"choices":[{"index":0,"delta":{"content":""},"finish_reason":"stop"}],"usage":{"prompt_tokens":1,"completion_tokens":2,"total_tokens":3}}`,
+ })
+
+ open := map[int]string{} // output_index -> item type
+ for _, e := range events {
+ switch e.Type {
+ case "response.output_item.added":
+ require.NotNil(t, e.Item)
+ open[e.OutputIndex] = e.Item.Type
+ case "response.reasoning_summary_text.delta":
+ require.Equalf(t, "reasoning", open[e.OutputIndex], "reasoning delta before its item was opened")
+ case "response.output_text.delta":
+ require.Equalf(t, "message", open[e.OutputIndex], "text delta before its item was opened")
+ }
+ }
+}
+
+// TestStream_ToolCallLifecycleComplete guards that a tool call is fully closed
+// (function_call_arguments.done + output_item.done with full arguments), which
+// codex needs to execute the call.
+func TestStream_ToolCallLifecycleComplete(t *testing.T) {
+ events := collectStreamEvents(t, []string{
+ `{"choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"plan"}}]}`,
+ `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_a","type":"function","function":{"name":"exec","arguments":""}}]}}]}`,
+ `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"cmd\":\"ls\"}"}}]}}]}`,
+ `{"choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}],"usage":{"prompt_tokens":1,"completion_tokens":2,"total_tokens":3}}`,
+ })
+
+ var sawAdded, sawArgsDone, sawItemDone bool
+ for _, e := range events {
+ switch e.Type {
+ case "response.output_item.added":
+ if e.Item != nil && e.Item.Type == "function_call" {
+ sawAdded = true
+ }
+ case "response.function_call_arguments.done":
+ sawArgsDone = true
+ require.Equal(t, `{"cmd":"ls"}`, e.Arguments)
+ case "response.output_item.done":
+ if e.Item != nil && e.Item.Type == "function_call" {
+ sawItemDone = true
+ require.Equal(t, `{"cmd":"ls"}`, e.Item.Arguments)
+ require.Equal(t, "call_a", e.Item.CallID)
+ }
+ }
+ }
+ require.True(t, sawAdded, "function_call output_item.added missing")
+ require.True(t, sawArgsDone, "function_call_arguments.done missing")
+ require.True(t, sawItemDone, "function_call output_item.done missing")
+}
+
+// TestStream_SSEWireComplete drives the full stream through SSE encoding and
+// asserts the function_call events carry complete fields on the wire.
+func TestStream_SSEWireComplete(t *testing.T) {
+ events := collectStreamEvents(t, []string{
+ `{"choices":[{"index":0,"delta":{"role":"assistant","reasoning_content":"plan"}}]}`,
+ `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_a","type":"function","function":{"name":"exec","arguments":"{}"}}]}}]}`,
+ `{"choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}`,
+ })
+
+ var addedLine string
+ for _, e := range events {
+ sse, err := ResponsesEventToSSE(e)
+ require.NoError(t, err)
+ if e.Type == "response.output_item.added" && e.Item != nil && e.Item.Type == "function_call" {
+ addedLine = sse
+ }
+ }
+ require.NotEmpty(t, addedLine)
+ // The function_call added event must carry arguments:"" on the wire.
+ require.True(t, strings.Contains(addedLine, `"arguments":""`), "added line missing arguments: %s", addedLine)
+ require.Contains(t, addedLine, `"call_id":"call_a"`)
+}
diff --git a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go
index 016c2415..b03b012f 100644
--- a/backend/internal/pkg/apicompat/chatcompletions_responses_test.go
+++ b/backend/internal/pkg/apicompat/chatcompletions_responses_test.go
@@ -663,6 +663,115 @@ func TestResponsesToChatCompletions_CachedTokens(t *testing.T) {
assert.Equal(t, 80, chat.Usage.PromptTokensDetails.CachedTokens)
}
+func TestResponsesToChatCompletions_ReasoningTokens(t *testing.T) {
+ resp := &ResponsesResponse{
+ ID: "resp_reasoning",
+ Status: "completed",
+ Output: []ResponsesOutput{
+ {
+ Type: "message",
+ Content: []ResponsesContentPart{{Type: "output_text", Text: "ping"}},
+ },
+ },
+ Usage: &ResponsesUsage{
+ InputTokens: 24,
+ OutputTokens: 33,
+ TotalTokens: 57,
+ OutputTokensDetails: &ResponsesOutputTokensDetails{
+ ReasoningTokens: 32,
+ },
+ },
+ }
+
+ chat := ResponsesToChatCompletions(resp, "gpt-5.5")
+ require.NotNil(t, chat.Usage)
+ assert.Equal(t, 33, chat.Usage.CompletionTokens)
+ require.NotNil(t, chat.Usage.CompletionTokensDetails)
+ assert.Equal(t, 32, chat.Usage.CompletionTokensDetails.ReasoningTokens)
+}
+
+func TestResponsesToChatCompletions_AllTokenDetailsPassThrough(t *testing.T) {
+ // Covers the full OpenAI CompletionUsage detail field set so future audio
+ // and prediction-outputs responses propagate without further changes.
+ resp := &ResponsesResponse{
+ ID: "resp_full_details",
+ Status: "completed",
+ Output: []ResponsesOutput{
+ {
+ Type: "message",
+ Content: []ResponsesContentPart{{Type: "output_text", Text: "x"}},
+ },
+ },
+ Usage: &ResponsesUsage{
+ InputTokens: 100,
+ OutputTokens: 50,
+ TotalTokens: 150,
+ InputTokensDetails: &ResponsesInputTokensDetails{
+ CachedTokens: 60,
+ AudioTokens: 4,
+ },
+ OutputTokensDetails: &ResponsesOutputTokensDetails{
+ ReasoningTokens: 30,
+ AudioTokens: 2,
+ AcceptedPredictionTokens: 10,
+ RejectedPredictionTokens: 3,
+ },
+ },
+ }
+
+ chat := ResponsesToChatCompletions(resp, "gpt-5.5")
+ require.NotNil(t, chat.Usage)
+ require.NotNil(t, chat.Usage.PromptTokensDetails)
+ assert.Equal(t, 60, chat.Usage.PromptTokensDetails.CachedTokens)
+ assert.Equal(t, 4, chat.Usage.PromptTokensDetails.AudioTokens)
+
+ require.NotNil(t, chat.Usage.CompletionTokensDetails)
+ assert.Equal(t, 30, chat.Usage.CompletionTokensDetails.ReasoningTokens)
+ assert.Equal(t, 2, chat.Usage.CompletionTokensDetails.AudioTokens)
+ assert.Equal(t, 10, chat.Usage.CompletionTokensDetails.AcceptedPredictionTokens)
+ assert.Equal(t, 3, chat.Usage.CompletionTokensDetails.RejectedPredictionTokens)
+
+ raw, err := json.Marshal(chat.Usage)
+ require.NoError(t, err)
+ assert.Contains(t, string(raw), `"prompt_tokens_details"`)
+ assert.Contains(t, string(raw), `"completion_tokens_details"`)
+ assert.Contains(t, string(raw), `"reasoning_tokens":30`)
+ assert.Contains(t, string(raw), `"accepted_prediction_tokens":10`)
+}
+
+func TestResponsesToChatCompletions_NoReasoningTokensWhenZero(t *testing.T) {
+ // Non-reasoning models do not return reasoning_tokens. The mapping must
+ // omit completion_tokens_details entirely rather than emitting a zero-valued
+ // field, so non-reasoning responses stay clean.
+ resp := &ResponsesResponse{
+ ID: "resp_no_reasoning",
+ Status: "completed",
+ Output: []ResponsesOutput{
+ {
+ Type: "message",
+ Content: []ResponsesContentPart{{Type: "output_text", Text: "hi"}},
+ },
+ },
+ Usage: &ResponsesUsage{
+ InputTokens: 10,
+ OutputTokens: 5,
+ TotalTokens: 15,
+ OutputTokensDetails: &ResponsesOutputTokensDetails{
+ ReasoningTokens: 0,
+ },
+ },
+ }
+
+ chat := ResponsesToChatCompletions(resp, "gpt-4o")
+ require.NotNil(t, chat.Usage)
+ assert.Nil(t, chat.Usage.CompletionTokensDetails)
+
+ raw, err := json.Marshal(chat.Usage)
+ require.NoError(t, err)
+ assert.NotContains(t, string(raw), "completion_tokens_details")
+ assert.NotContains(t, string(raw), "reasoning_tokens")
+}
+
func TestResponsesToChatCompletions_WebSearch(t *testing.T) {
resp := &ResponsesResponse{
ID: "resp_ws",
@@ -825,6 +934,32 @@ func TestResponsesEventToChatChunks_Completed(t *testing.T) {
assert.Equal(t, 30, chunks[1].Usage.PromptTokensDetails.CachedTokens)
}
+func TestResponsesEventToChatChunks_CompletedWithReasoningTokens(t *testing.T) {
+ state := NewResponsesEventToChatState()
+ state.Model = "gpt-5.5"
+ state.IncludeUsage = true
+
+ chunks := ResponsesEventToChatChunks(&ResponsesStreamEvent{
+ Type: "response.completed",
+ Response: &ResponsesResponse{
+ Status: "completed",
+ Usage: &ResponsesUsage{
+ InputTokens: 24,
+ OutputTokens: 33,
+ TotalTokens: 57,
+ OutputTokensDetails: &ResponsesOutputTokensDetails{
+ ReasoningTokens: 32,
+ },
+ },
+ },
+ }, state)
+ require.Len(t, chunks, 2)
+
+ require.NotNil(t, chunks[1].Usage)
+ require.NotNil(t, chunks[1].Usage.CompletionTokensDetails)
+ assert.Equal(t, 32, chunks[1].Usage.CompletionTokensDetails.ReasoningTokens)
+}
+
func TestResponsesEventToChatChunks_ResponseDone(t *testing.T) {
state := NewResponsesEventToChatState()
state.Model = "gpt-4o"
diff --git a/backend/internal/pkg/apicompat/responses_stream_event_wire.go b/backend/internal/pkg/apicompat/responses_stream_event_wire.go
new file mode 100644
index 00000000..df7a82e3
--- /dev/null
+++ b/backend/internal/pkg/apicompat/responses_stream_event_wire.go
@@ -0,0 +1,199 @@
+package apicompat
+
+import "encoding/json"
+
+// MarshalJSON renders a ResponsesStreamEvent into its wire form.
+//
+// The OpenAI Responses streaming protocol requires several fields to be present
+// even when they hold a zero value: output_index/content_index/summary_index are
+// meaningful at 0, a function_call item must always carry call_id/name/arguments
+// (arguments may be ""), a message item must carry content:[] and an output_text
+// part must carry text/annotations/logprobs. Go's `omitempty` drops exactly those
+// zero values, and strict clients (Codex CLI) reject items/deltas whose required
+// fields are missing.
+//
+// Rather than marshalling with omitempty and patching the JSON afterwards, every
+// streamed event type is constructed explicitly here — the Go analogue of the
+// reference gateways' (cc-switch, CCX) per-event object construction. This is the
+// single source of truth for Responses SSE field presence and applies uniformly
+// to every emitter (Chat→Responses bridge and Anthropic→Responses converter).
+//
+// Event types not listed fall back to the default struct marshalling, which
+// bounds the blast radius of this method to the streamed item/part/text/tool
+// events.
+func (e ResponsesStreamEvent) MarshalJSON() ([]byte, error) {
+ switch e.Type {
+ case "response.output_text.delta", "response.output_text.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ m["content_index"] = e.ContentIndex
+ if e.Type == "response.output_text.done" {
+ m["text"] = e.Text
+ } else {
+ m["delta"] = e.Delta
+ }
+ return json.Marshal(m)
+
+ case "response.content_part.added", "response.content_part.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ m["content_index"] = e.ContentIndex
+ m["part"] = outputTextPartWire(e.Part)
+ return json.Marshal(m)
+
+ case "response.reasoning_summary_text.delta", "response.reasoning_summary_text.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ m["summary_index"] = e.SummaryIndex
+ if e.Type == "response.reasoning_summary_text.done" {
+ m["text"] = e.Text
+ } else {
+ m["delta"] = e.Delta
+ }
+ return json.Marshal(m)
+
+ case "response.reasoning_summary_part.added", "response.reasoning_summary_part.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ m["summary_index"] = e.SummaryIndex
+ m["part"] = summaryTextPartWire(e.Part)
+ return json.Marshal(m)
+
+ case "response.output_item.added", "response.output_item.done":
+ m := e.wireBase()
+ m["output_index"] = e.OutputIndex
+ m["item"] = responsesItemWire(e.Item)
+ return json.Marshal(m)
+
+ case "response.function_call_arguments.delta", "response.function_call_arguments.done":
+ m := e.wireBase()
+ e.putItemID(m)
+ m["output_index"] = e.OutputIndex
+ if e.CallID != "" {
+ m["call_id"] = e.CallID
+ }
+ if e.Name != "" {
+ m["name"] = e.Name
+ }
+ if e.Type == "response.function_call_arguments.done" {
+ m["arguments"] = e.Arguments
+ } else {
+ m["delta"] = e.Delta
+ }
+ return json.Marshal(m)
+
+ default:
+ // response.created / completed / done / failed / incomplete and any
+ // event type not shaped above keep the default struct marshalling.
+ type alias ResponsesStreamEvent
+ return json.Marshal(alias(e))
+ }
+}
+
+func (e ResponsesStreamEvent) wireBase() map[string]any {
+ m := map[string]any{
+ "type": e.Type,
+ "sequence_number": e.SequenceNumber,
+ }
+ return m
+}
+
+func (e ResponsesStreamEvent) putItemID(m map[string]any) {
+ if e.ItemID != "" {
+ m["item_id"] = e.ItemID
+ }
+}
+
+// outputTextPartWire renders a content part for a message's output_text, always
+// carrying text/annotations/logprobs (matching cc-switch's push_text_delta).
+func outputTextPartWire(part *ResponsesContentPart) map[string]any {
+ text := ""
+ if part != nil {
+ text = part.Text
+ }
+ return map[string]any{
+ "type": "output_text",
+ "text": text,
+ "annotations": []any{},
+ "logprobs": []any{},
+ }
+}
+
+// summaryTextPartWire renders a reasoning summary part.
+func summaryTextPartWire(part *ResponsesContentPart) map[string]any {
+ text := ""
+ if part != nil {
+ text = part.Text
+ }
+ return map[string]any{
+ "type": "summary_text",
+ "text": text,
+ }
+}
+
+// responsesItemWire renders an output_item with every field the item's type
+// requires to be present, including the empty arrays/strings that omitempty
+// would otherwise drop. Mirrors cc-switch's response_function_call_item and the
+// message/reasoning item shapes codex expects.
+func responsesItemWire(item *ResponsesOutput) map[string]any {
+ if item == nil {
+ return map[string]any{}
+ }
+ m := map[string]any{
+ "type": item.Type,
+ "id": item.ID,
+ }
+ if item.Status != "" {
+ m["status"] = item.Status
+ }
+ switch item.Type {
+ case "message":
+ role := item.Role
+ if role == "" {
+ role = "assistant"
+ }
+ m["role"] = role
+ m["content"] = messageContentWire(item.Content)
+ case "reasoning":
+ m["summary"] = reasoningSummaryWire(item.Summary)
+ if item.EncryptedContent != "" {
+ m["encrypted_content"] = item.EncryptedContent
+ }
+ case "function_call":
+ m["call_id"] = item.CallID
+ m["name"] = item.Name
+ m["arguments"] = item.Arguments
+ }
+ return m
+}
+
+// messageContentWire renders a message item's content array; always an array
+// (never null), with each output_text part carrying its text.
+func messageContentWire(parts []ResponsesContentPart) []map[string]any {
+ out := make([]map[string]any, 0, len(parts))
+ for _, p := range parts {
+ typ := p.Type
+ if typ == "" {
+ typ = "output_text"
+ }
+ out = append(out, map[string]any{"type": typ, "text": p.Text})
+ }
+ return out
+}
+
+// reasoningSummaryWire renders a reasoning item's summary array; always an array.
+func reasoningSummaryWire(summary []ResponsesSummary) []map[string]any {
+ out := make([]map[string]any, 0, len(summary))
+ for _, s := range summary {
+ typ := s.Type
+ if typ == "" {
+ typ = "summary_text"
+ }
+ out = append(out, map[string]any{"type": typ, "text": s.Text})
+ }
+ return out
+}
diff --git a/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go b/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go
new file mode 100644
index 00000000..b4f6871d
--- /dev/null
+++ b/backend/internal/pkg/apicompat/responses_stream_event_wire_test.go
@@ -0,0 +1,113 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+// marshalEvent marshals through the custom MarshalJSON and returns the decoded
+// object plus the set of top-level keys.
+func marshalEvent(t *testing.T, e ResponsesStreamEvent) map[string]any {
+ t.Helper()
+ b, err := json.Marshal(e)
+ require.NoError(t, err)
+ var m map[string]any
+ require.NoError(t, json.Unmarshal(b, &m))
+ return m
+}
+
+// TestWire_IndexFieldsPresentAtZero guards the omitempty trap: output_index/
+// content_index/summary_index must serialize even when 0.
+func TestWire_IndexFieldsPresentAtZero(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.output_text.delta", OutputIndex: 0, ContentIndex: 0, ItemID: "msg_1", Delta: "hi",
+ })
+ require.Contains(t, m, "output_index")
+ require.Contains(t, m, "content_index")
+ require.EqualValues(t, 0, m["output_index"])
+
+ r := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.reasoning_summary_text.delta", OutputIndex: 0, SummaryIndex: 0, ItemID: "rs_1", Delta: "think",
+ })
+ require.Contains(t, r, "output_index")
+ require.Contains(t, r, "summary_index")
+}
+
+// TestWire_FunctionCallItemAlwaysComplete guards that a function_call item
+// always carries call_id/name/arguments, including arguments:"" on .added.
+func TestWire_FunctionCallItemAlwaysComplete(t *testing.T) {
+ added := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.output_item.added",
+ OutputIndex: 1,
+ Item: &ResponsesOutput{Type: "function_call", ID: "fc_1", CallID: "call_a", Name: "exec", Status: "in_progress"},
+ })
+ item, ok := added["item"].(map[string]any)
+ require.True(t, ok, "item must be an object")
+ for _, k := range []string{"call_id", "name", "arguments"} {
+ require.Containsf(t, item, k, "function_call item missing %q", k)
+ }
+ require.Equal(t, "", item["arguments"])
+}
+
+// TestWire_MessageItemContentAlwaysArray guards content:[] presence.
+func TestWire_MessageItemContentAlwaysArray(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.output_item.added",
+ OutputIndex: 0,
+ Item: &ResponsesOutput{Type: "message", ID: "msg_1", Role: "assistant", Status: "in_progress"},
+ })
+ item, ok := m["item"].(map[string]any)
+ require.True(t, ok, "item must be an object")
+ require.Contains(t, item, "content")
+ _, ok = item["content"].([]any)
+ require.True(t, ok, "content must be an array")
+}
+
+// TestWire_ReasoningItemSummaryAlwaysArray guards summary:[] presence.
+func TestWire_ReasoningItemSummaryAlwaysArray(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.output_item.added",
+ OutputIndex: 0,
+ Item: &ResponsesOutput{Type: "reasoning", ID: "rs_1", Status: "in_progress"},
+ })
+ item, ok := m["item"].(map[string]any)
+ require.True(t, ok, "item must be an object")
+ require.Contains(t, item, "summary")
+ _, ok = item["summary"].([]any)
+ require.True(t, ok, "summary must be an array")
+}
+
+// TestWire_ContentPartCarriesAnnotationsLogprobs guards the output_text part shape.
+func TestWire_ContentPartCarriesAnnotationsLogprobs(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.content_part.added", OutputIndex: 0, ContentIndex: 0, ItemID: "msg_1",
+ Part: &ResponsesContentPart{Type: "output_text", Text: ""},
+ })
+ part, ok := m["part"].(map[string]any)
+ require.True(t, ok, "part must be an object")
+ require.Equal(t, "output_text", part["type"])
+ require.Contains(t, part, "text")
+ require.Contains(t, part, "annotations")
+ require.Contains(t, part, "logprobs")
+}
+
+// TestWire_ArgumentsDonePresentEvenEmpty guards arguments presence on done.
+func TestWire_ArgumentsDonePresentEvenEmpty(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.function_call_arguments.done", OutputIndex: 1, ItemID: "fc_1", CallID: "call_a", Name: "exec", Arguments: "",
+ })
+ require.Contains(t, m, "arguments")
+ require.Equal(t, "", m["arguments"])
+}
+
+// TestWire_UnknownEventFallsBackToDefault ensures non-streamed event types keep
+// default marshalling (the response object is preserved).
+func TestWire_UnknownEventFallsBackToDefault(t *testing.T) {
+ m := marshalEvent(t, ResponsesStreamEvent{
+ Type: "response.completed",
+ Response: &ResponsesResponse{ID: "resp_1", Object: "response", Status: "completed"},
+ })
+ require.Contains(t, m, "response")
+}
diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic_cc_chain_test.go b/backend/internal/pkg/apicompat/responses_to_anthropic_cc_chain_test.go
new file mode 100644
index 00000000..d64680f4
--- /dev/null
+++ b/backend/internal/pkg/apicompat/responses_to_anthropic_cc_chain_test.go
@@ -0,0 +1,102 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+// These tests drive the exact production path for Chat Completions clients on an
+// Anthropic-platform group: ForwardAsChatCompletions runs
+// ChatCompletionsToResponses → ResponsesToAnthropicRequest
+// (gateway_forward_as_chat_completions.go), then forwards the Anthropic body
+// upstream. They assert the tool-pairing repair holds through that full chain,
+// not only for codex-style Responses input.
+func ccChainToAnthropic(t *testing.T, ccReq *ChatCompletionsRequest) []AnthropicMessage {
+ t.Helper()
+ respReq, err := ChatCompletionsToResponses(ccReq)
+ require.NoError(t, err)
+ anthReq, err := ResponsesToAnthropicRequest(respReq)
+ require.NoError(t, err)
+ assertAnthropicPairing(t, anthReq.Messages)
+ return anthReq.Messages
+}
+
+// Reproduces the production 400:
+//
+// unexpected ...content.0: tool_use_id found in tool_result blocks:
+// call_00_TgfbRvKlnD7oK6Dg00sL1661. Each tool_result block must have a
+// corresponding tool_use block in the previous message.
+//
+// A Chat Completions client trimmed its history and kept a tool result whose
+// announcing assistant tool_calls message was dropped (sliding-window context
+// management). The orphan tool_result has no matching tool_use → upstream 400.
+// The repair drops the orphan so the request is valid.
+func TestCCChain_OrphanToolResultFromTrimmedHistory(t *testing.T) {
+ orphanID := "call_00_TgfbRvKlnD7oK6Dg00sL1661"
+ msgs := ccChainToAnthropic(t, &ChatCompletionsRequest{
+ Model: "deepseek-v4-pro",
+ Messages: []ChatMessage{
+ {Role: "user", Content: json.RawMessage(`"search the web for X"`)},
+ // The assistant tool_calls message that announced orphanID was trimmed.
+ {Role: "tool", ToolCallID: orphanID, Content: json.RawMessage(`"stale search results"`)},
+ {Role: "assistant", Content: json.RawMessage(`"Here is what I found."`)},
+ {Role: "user", Content: json.RawMessage(`"thanks, now do Y"`)},
+ },
+ })
+ for _, m := range msgs {
+ require.Falsef(t, hasToolResult(parseContentBlocks(m.Content), orphanID),
+ "orphan tool_result %s should have been dropped", orphanID)
+ }
+}
+
+// A parallel web_search where one sibling's result never came back (the tool
+// failed/was skipped). The unanswered tool_use would otherwise trip Anthropic's
+// "tool_use without tool_result" check; the repair drops it.
+func TestCCChain_ParallelToolOneResultMissing(t *testing.T) {
+ msgs := ccChainToAnthropic(t, &ChatCompletionsRequest{
+ Model: "deepseek-v4-pro",
+ Messages: []ChatMessage{
+ {Role: "user", Content: json.RawMessage(`"search A and B"`)},
+ {Role: "assistant", Content: json.RawMessage(`"searching both"`), ToolCalls: []ChatToolCall{
+ {ID: "call_a", Type: "function", Function: ChatFunctionCall{Name: "web_search", Arguments: `{"q":"A"}`}},
+ {ID: "call_b", Type: "function", Function: ChatFunctionCall{Name: "web_search", Arguments: `{"q":"B"}`}},
+ }},
+ {Role: "tool", ToolCallID: "call_a", Content: json.RawMessage(`"result A"`)},
+ // call_b's result is missing.
+ },
+ })
+ for _, m := range msgs {
+ require.Falsef(t, hasToolUse(parseContentBlocks(m.Content), "call_b"),
+ "unanswered tool_use call_b should have been dropped")
+ }
+}
+
+// Baseline: a well-formed multi-round tool history (text + tool_calls per
+// assistant turn) converts and pairs correctly through the full chain.
+func TestCCChain_WellFormedMultiRound(t *testing.T) {
+ msgs := ccChainToAnthropic(t, &ChatCompletionsRequest{
+ Model: "deepseek-v4-pro",
+ Messages: []ChatMessage{
+ {Role: "user", Content: json.RawMessage(`"do A then B"`)},
+ {Role: "assistant", Content: json.RawMessage(`"running A"`), ToolCalls: []ChatToolCall{
+ {ID: "call_a", Type: "function", Function: ChatFunctionCall{Name: "exec", Arguments: `{"cmd":"A"}`}},
+ }},
+ {Role: "tool", ToolCallID: "call_a", Content: json.RawMessage(`"A ok"`)},
+ {Role: "assistant", Content: json.RawMessage(`"A done, running B"`), ToolCalls: []ChatToolCall{
+ {ID: "call_b", Type: "function", Function: ChatFunctionCall{Name: "exec", Arguments: `{"cmd":"B"}`}},
+ }},
+ {Role: "tool", ToolCallID: "call_b", Content: json.RawMessage(`"B ok"`)},
+ {Role: "assistant", Content: json.RawMessage(`"all done"`)},
+ },
+ })
+ // Both calls survive and stay paired (assertAnthropicPairing already checks).
+ var sawA, sawB bool
+ for _, m := range msgs {
+ blocks := parseContentBlocks(m.Content)
+ sawA = sawA || hasToolUse(blocks, "call_a")
+ sawB = sawB || hasToolUse(blocks, "call_b")
+ }
+ require.True(t, sawA && sawB, "both well-formed calls should be preserved")
+}
diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic_request.go b/backend/internal/pkg/apicompat/responses_to_anthropic_request.go
index 8fa652f2..672ad80c 100644
--- a/backend/internal/pkg/apicompat/responses_to_anthropic_request.go
+++ b/backend/internal/pkg/apicompat/responses_to_anthropic_request.go
@@ -192,12 +192,135 @@ func convertResponsesInputToAnthropic(inputRaw json.RawMessage) (json.RawMessage
}
}
- // Merge consecutive same-role messages (Anthropic requires alternating roles)
+ // Repair tool_use/tool_result pairing, then merge consecutive same-role
+ // messages (Anthropic requires alternating roles). The first merge groups
+ // parallel calls (and their results) so the pairing pass sees them together;
+ // the pairing pass may re-split a user turn (e.g. when an injected message
+ // sat between a call and its output), so a second merge restores alternation.
+ messages = mergeConsecutiveMessages(messages)
+ messages = normalizeAnthropicToolPairing(messages)
messages = mergeConsecutiveMessages(messages)
return system, messages, nil
}
+// normalizeAnthropicToolPairing rebuilds the message sequence so it satisfies
+// Anthropic's tool_use/tool_result invariants, which the naive item-by-item
+// conversion violates whenever the Responses history interleaves anything
+// between a function_call and its function_call_output:
+//
+// - every tool_result block must have a matching tool_use in the immediately
+// preceding assistant message ("tool_result ... must have a corresponding
+// tool_use block in the previous message");
+// - every tool_use block must be answered by a tool_result in the immediately
+// following user message (Anthropic rejects unanswered tool_use ids);
+// - user/assistant turns must alternate.
+//
+// codex (Responses, store:false) re-sends the whole history each turn and
+// frequently injects items between a call and its output — a developer/approval
+// notice, or a sibling parallel call whose output never arrived. The unrepaired
+// converter emits each function_call as its own assistant message and each
+// output as its own user message, so any such interleaving breaks
+// tool_use↔tool_result adjacency and yields an upstream 400.
+//
+// The repair indexes every tool_result by its tool_use id, then for each
+// assistant message carrying tool_use blocks keeps only the answered ones
+// (dropping unanswered/dangling calls — and the assistant message entirely if it
+// has no other content) and emits the matching tool_result blocks, in call
+// order, as the very next user message. Standalone tool_result blocks are
+// dropped from their original position (re-emitted adjacent to their call);
+// orphan tool_results with no announcing tool_use are dropped. Non-tool content
+// passes through in place. This mirrors normalizeChatMessages on the
+// Responses→Chat path.
+func normalizeAnthropicToolPairing(messages []AnthropicMessage) []AnthropicMessage {
+ // Index every tool_result block by its tool_use id (last wins on dup).
+ results := make(map[string]AnthropicContentBlock)
+ for _, m := range messages {
+ if m.Role != "user" {
+ continue
+ }
+ for _, b := range parseContentBlocks(m.Content) {
+ if b.Type == "tool_result" && b.ToolUseID != "" {
+ results[b.ToolUseID] = b
+ }
+ }
+ }
+
+ out := make([]AnthropicMessage, 0, len(messages))
+ for _, m := range messages {
+ blocks := parseContentBlocks(m.Content)
+ switch m.Role {
+ case "assistant":
+ var toolUses, others []AnthropicContentBlock
+ for _, b := range blocks {
+ if b.Type == "tool_use" {
+ toolUses = append(toolUses, b)
+ } else {
+ others = append(others, b)
+ }
+ }
+ if len(toolUses) == 0 {
+ out = append(out, m)
+ continue
+ }
+ kept := make([]AnthropicContentBlock, 0, len(toolUses))
+ for _, tu := range toolUses {
+ if _, ok := results[tu.ID]; ok {
+ kept = append(kept, tu)
+ }
+ }
+ if len(kept) == 0 {
+ // No answered calls: keep any non-tool content, else drop.
+ if len(others) > 0 {
+ out = append(out, anthropicMessageFromBlocks("assistant", others))
+ }
+ continue
+ }
+ asstBlocks := make([]AnthropicContentBlock, 0, len(others)+len(kept))
+ asstBlocks = append(asstBlocks, others...)
+ asstBlocks = append(asstBlocks, kept...)
+ out = append(out, anthropicMessageFromBlocks("assistant", asstBlocks))
+
+ resBlocks := make([]AnthropicContentBlock, 0, len(kept))
+ for _, tu := range kept {
+ resBlocks = append(resBlocks, results[tu.ID])
+ }
+ out = append(out, anthropicMessageFromBlocks("user", resBlocks))
+
+ case "user":
+ var nonResult []AnthropicContentBlock
+ hasResult := false
+ for _, b := range blocks {
+ if b.Type == "tool_result" {
+ hasResult = true
+ continue
+ }
+ nonResult = append(nonResult, b)
+ }
+ if !hasResult {
+ out = append(out, m)
+ continue
+ }
+ // The tool_result blocks are re-emitted next to their call; keep any
+ // other content of this user turn in place, drop it if there is none.
+ if len(nonResult) > 0 {
+ out = append(out, anthropicMessageFromBlocks("user", nonResult))
+ }
+
+ default:
+ out = append(out, m)
+ }
+ }
+ return out
+}
+
+// anthropicMessageFromBlocks builds an AnthropicMessage whose content is the
+// marshaled block array.
+func anthropicMessageFromBlocks(role string, blocks []AnthropicContentBlock) AnthropicMessage {
+ content, _ := json.Marshal(blocks)
+ return AnthropicMessage{Role: role, Content: content}
+}
+
// extractTextFromContent extracts text from a content field that may be a
// plain string or an array of content parts.
func extractTextFromContent(raw json.RawMessage) string {
diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic_tool_pairing_test.go b/backend/internal/pkg/apicompat/responses_to_anthropic_tool_pairing_test.go
new file mode 100644
index 00000000..b2522f27
--- /dev/null
+++ b/backend/internal/pkg/apicompat/responses_to_anthropic_tool_pairing_test.go
@@ -0,0 +1,165 @@
+package apicompat
+
+import (
+ "encoding/json"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+// assertAnthropicPairing enforces the Anthropic Messages tool-pairing invariants
+// that, when violated, surface as upstream 400s.
+func assertAnthropicPairing(t *testing.T, messages []AnthropicMessage) {
+ t.Helper()
+ for i, m := range messages {
+ blocks := parseContentBlocks(m.Content)
+
+ // No two consecutive same-role messages.
+ if i > 0 {
+ require.NotEqualf(t, messages[i-1].Role, m.Role, "consecutive %s messages at %d", m.Role, i)
+ }
+
+ for _, b := range blocks {
+ switch b.Type {
+ case "tool_result":
+ // Must have a matching tool_use in the immediately previous message.
+ require.Positivef(t, i, "tool_result %s has no previous message", b.ToolUseID)
+ prev := parseContentBlocks(messages[i-1].Content)
+ require.Truef(t, hasToolUse(prev, b.ToolUseID),
+ "tool_result %s has no corresponding tool_use in previous message", b.ToolUseID)
+ case "tool_use":
+ // Must be answered by a tool_result in the immediately next message.
+ require.Lessf(t, i+1, len(messages), "tool_use %s has no following message", b.ID)
+ next := parseContentBlocks(messages[i+1].Content)
+ require.Truef(t, hasToolResult(next, b.ID),
+ "tool_use %s is not answered in the next message", b.ID)
+ }
+ }
+ }
+}
+
+func hasToolUse(blocks []AnthropicContentBlock, id string) bool {
+ for _, b := range blocks {
+ if b.Type == "tool_use" && b.ID == id {
+ return true
+ }
+ }
+ return false
+}
+
+func hasToolResult(blocks []AnthropicContentBlock, toolUseID string) bool {
+ for _, b := range blocks {
+ if b.Type == "tool_result" && b.ToolUseID == toolUseID {
+ return true
+ }
+ }
+ return false
+}
+
+func convertAnthropic(t *testing.T, input string) []AnthropicMessage {
+ t.Helper()
+ _, messages, err := convertResponsesInputToAnthropic(json.RawMessage(input))
+ require.NoError(t, err)
+ assertAnthropicPairing(t, messages)
+ return messages
+}
+
+// Tests use call_-prefixed ids because fromResponsesCallIDToAnthropic passes
+// those through unchanged (matching codex's real call_00_... ids); bare ids
+// would be rewritten to toolu_.
+
+// A developer/approval message injected between a function_call and its output
+// must be moved out of the tool_use→tool_result adjacency. This is the shape
+// that produced the production 400 "tool_result ... must have a corresponding
+// tool_use block in the previous message".
+func TestAnthropicPairing_DeveloperMessageBetween(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"do it"}]},
+ {"type":"function_call","call_id":"call_A","name":"exec","arguments":"{}"},
+ {"type":"message","role":"developer","content":[{"type":"input_text","text":"Approved command prefix saved"}]},
+ {"type":"function_call_output","call_id":"call_A","output":"ok"}
+ ]`)
+ // The assistant tool_use message is immediately followed by its tool_result.
+ for i, m := range msgs {
+ if hasToolUse(parseContentBlocks(m.Content), "call_A") {
+ require.Equal(t, "user", msgs[i+1].Role)
+ require.True(t, hasToolResult(parseContentBlocks(msgs[i+1].Content), "call_A"))
+ }
+ }
+}
+
+// Parallel tool calls where both outputs arrive stay grouped: one assistant
+// message with both tool_use blocks, the next user message with both results.
+func TestAnthropicPairing_ParallelBothAnswered(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"features?"}]},
+ {"type":"function_call","call_id":"call_c0","name":"exec","arguments":"{}"},
+ {"type":"function_call","call_id":"call_c1","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"call_c0","output":"log"},
+ {"type":"function_call_output","call_id":"call_c1","output":"tags"}
+ ]`)
+ var sawGrouped bool
+ for _, m := range msgs {
+ blocks := parseContentBlocks(m.Content)
+ if hasToolUse(blocks, "call_c0") && hasToolUse(blocks, "call_c1") {
+ sawGrouped = true
+ }
+ }
+ require.True(t, sawGrouped, "parallel tool_use blocks should share one assistant message")
+}
+
+// A parallel call whose sibling output never arrived must be dropped so every
+// remaining tool_use is answered.
+func TestAnthropicPairing_ParallelOneUnanswered(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"function_call","call_id":"call_A","name":"exec","arguments":"{}"},
+ {"type":"function_call","call_id":"call_B","name":"exec","arguments":"{}"},
+ {"type":"function_call_output","call_id":"call_A","output":"oa"}
+ ]`)
+ for _, m := range msgs {
+ require.Falsef(t, hasToolUse(parseContentBlocks(m.Content), "call_B"),
+ "unanswered tool_use call_B should have been dropped")
+ }
+}
+
+// An orphan tool_result whose tool_use was never announced must be dropped.
+func TestAnthropicPairing_OrphanToolResultDropped(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"function_call_output","call_id":"call_ghost","output":"orphan"}
+ ]`)
+ for _, m := range msgs {
+ require.Falsef(t, hasToolResult(parseContentBlocks(m.Content), "call_ghost"),
+ "orphan tool_result should have been dropped")
+ }
+}
+
+// A dangling tool_call at the end of the history (no output yet) drops the
+// assistant message holding only that call, leaving no tool_use behind.
+func TestAnthropicPairing_DanglingCallDropped(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"q"}]},
+ {"type":"function_call","call_id":"call_A","name":"exec","arguments":"{}"}
+ ]`)
+ for _, m := range msgs {
+ require.Falsef(t, hasToolUse(parseContentBlocks(m.Content), "call_A"),
+ "dangling tool_use call_A should have been dropped")
+ }
+}
+
+// Baseline: a single answered call pairs correctly and preserves the surrounding
+// turns.
+func TestAnthropicPairing_SingleCall(t *testing.T) {
+ msgs := convertAnthropic(t, `[
+ {"type":"message","role":"user","content":[{"type":"input_text","text":"latest sha?"}]},
+ {"type":"function_call","call_id":"call_A","name":"exec","arguments":"{\"cmd\":\"git rev-parse HEAD\"}"},
+ {"type":"function_call_output","call_id":"call_A","output":"deadbeef"},
+ {"type":"message","role":"assistant","content":[{"type":"output_text","text":"It is deadbeef."}]}
+ ]`)
+ // user, assistant(tool_use), user(tool_result), assistant(text)
+ require.GreaterOrEqual(t, len(msgs), 4)
+ require.Equal(t, "user", msgs[0].Role)
+ require.True(t, hasToolUse(parseContentBlocks(msgs[1].Content), "call_A"))
+ require.True(t, hasToolResult(parseContentBlocks(msgs[2].Content), "call_A"))
+}
diff --git a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go
index 7e8354ee..8809b4fc 100644
--- a/backend/internal/pkg/apicompat/responses_to_chatcompletions.go
+++ b/backend/internal/pkg/apicompat/responses_to_chatcompletions.go
@@ -81,19 +81,7 @@ func ResponsesToChatCompletions(resp *ResponsesResponse, model string) *ChatComp
FinishReason: finishReason,
}}
- if resp.Usage != nil {
- usage := &ChatUsage{
- PromptTokens: resp.Usage.InputTokens,
- CompletionTokens: resp.Usage.OutputTokens,
- TotalTokens: resp.Usage.InputTokens + resp.Usage.OutputTokens,
- }
- if resp.Usage.InputTokensDetails != nil && resp.Usage.InputTokensDetails.CachedTokens > 0 {
- usage.PromptTokensDetails = &ChatTokenDetails{
- CachedTokens: resp.Usage.InputTokensDetails.CachedTokens,
- }
- }
- out.Usage = usage
- }
+ out.Usage = chatUsageFromResponsesUsage(resp.Usage)
return out
}
@@ -341,14 +329,48 @@ func chatUsageFromResponsesUsage(u *ResponsesUsage) *ChatUsage {
CompletionTokens: u.OutputTokens,
TotalTokens: u.InputTokens + u.OutputTokens,
}
- if u.InputTokensDetails != nil && u.InputTokensDetails.CachedTokens > 0 {
- usage.PromptTokensDetails = &ChatTokenDetails{
- CachedTokens: u.InputTokensDetails.CachedTokens,
- }
- }
+ usage.PromptTokensDetails = promptDetailsFromResponses(u.InputTokensDetails)
+ usage.CompletionTokensDetails = completionDetailsFromResponses(u.OutputTokensDetails)
return usage
}
+// promptDetailsFromResponses maps Responses-API input_tokens_details into a
+// Chat-Completions prompt_tokens_details. Returns nil when nothing would be
+// emitted, so upstreams that do not break down prompt usage stay clean.
+func promptDetailsFromResponses(src *ResponsesInputTokensDetails) *ChatTokenDetails {
+ if src == nil {
+ return nil
+ }
+ if src.CachedTokens == 0 && src.AudioTokens == 0 {
+ return nil
+ }
+ return &ChatTokenDetails{
+ CachedTokens: src.CachedTokens,
+ AudioTokens: src.AudioTokens,
+ }
+}
+
+// completionDetailsFromResponses maps Responses-API output_tokens_details
+// into a Chat-Completions completion_tokens_details. Mirrors the OpenAI
+// official CompletionUsage schema: reasoning_tokens, audio_tokens, and
+// the predicted-outputs accepted/rejected counts. Returns nil when nothing
+// would be emitted so non-reasoning, non-audio responses stay clean.
+func completionDetailsFromResponses(src *ResponsesOutputTokensDetails) *ChatTokenDetails {
+ if src == nil {
+ return nil
+ }
+ if src.ReasoningTokens == 0 && src.AudioTokens == 0 &&
+ src.AcceptedPredictionTokens == 0 && src.RejectedPredictionTokens == 0 {
+ return nil
+ }
+ return &ChatTokenDetails{
+ ReasoningTokens: src.ReasoningTokens,
+ AudioTokens: src.AudioTokens,
+ AcceptedPredictionTokens: src.AcceptedPredictionTokens,
+ RejectedPredictionTokens: src.RejectedPredictionTokens,
+ }
+}
+
func makeChatDeltaChunk(state *ResponsesEventToChatState, delta ChatDelta) ChatCompletionsChunk {
return ChatCompletionsChunk{
ID: state.ID,
diff --git a/backend/internal/pkg/apicompat/types.go b/backend/internal/pkg/apicompat/types.go
index 8b576647..d2937802 100644
--- a/backend/internal/pkg/apicompat/types.go
+++ b/backend/internal/pkg/apicompat/types.go
@@ -362,11 +362,15 @@ func (u *ResponsesUsage) UnmarshalJSON(data []byte) error {
// ResponsesInputTokensDetails breaks down input token usage.
type ResponsesInputTokensDetails struct {
CachedTokens int `json:"cached_tokens,omitempty"`
+ AudioTokens int `json:"audio_tokens,omitempty"`
}
// ResponsesOutputTokensDetails breaks down output token usage.
type ResponsesOutputTokensDetails struct {
- ReasoningTokens int `json:"reasoning_tokens,omitempty"`
+ ReasoningTokens int `json:"reasoning_tokens,omitempty"`
+ AudioTokens int `json:"audio_tokens,omitempty"`
+ AcceptedPredictionTokens int `json:"accepted_prediction_tokens,omitempty"`
+ RejectedPredictionTokens int `json:"rejected_prediction_tokens,omitempty"`
}
// ---------------------------------------------------------------------------
@@ -402,6 +406,10 @@ type ResponsesStreamEvent struct {
// Reuses Text/Delta fields above, SummaryIndex identifies which summary part
SummaryIndex int `json:"summary_index,omitempty"`
+ // response.content_part.added / done and
+ // response.reasoning_summary_part.added / done
+ Part *ResponsesContentPart `json:"part,omitempty"`
+
// error event fields
Code string `json:"code,omitempty"`
Param string `json:"param,omitempty"`
@@ -517,15 +525,27 @@ type ChatChoice struct {
// ChatUsage holds token counts in Chat Completions format.
type ChatUsage struct {
- PromptTokens int `json:"prompt_tokens"`
- CompletionTokens int `json:"completion_tokens"`
- TotalTokens int `json:"total_tokens"`
- PromptTokensDetails *ChatTokenDetails `json:"prompt_tokens_details,omitempty"`
+ PromptTokens int `json:"prompt_tokens"`
+ CompletionTokens int `json:"completion_tokens"`
+ TotalTokens int `json:"total_tokens"`
+ PromptTokensDetails *ChatTokenDetails `json:"prompt_tokens_details,omitempty"`
+ CompletionTokensDetails *ChatTokenDetails `json:"completion_tokens_details,omitempty"`
}
-// ChatTokenDetails provides a breakdown of token usage.
+// ChatTokenDetails provides a breakdown of token usage. The same type is
+// reused for both prompt_tokens_details and completion_tokens_details;
+// unset fields are omitted so each side only emits the fields that apply.
+//
+// Field set mirrors OpenAI's official CompletionUsage schema:
+// - prompt_tokens_details: cached_tokens, audio_tokens
+// - completion_tokens_details: reasoning_tokens, audio_tokens,
+// accepted_prediction_tokens, rejected_prediction_tokens
type ChatTokenDetails struct {
- CachedTokens int `json:"cached_tokens,omitempty"`
+ CachedTokens int `json:"cached_tokens,omitempty"`
+ AudioTokens int `json:"audio_tokens,omitempty"`
+ ReasoningTokens int `json:"reasoning_tokens,omitempty"`
+ AcceptedPredictionTokens int `json:"accepted_prediction_tokens,omitempty"`
+ RejectedPredictionTokens int `json:"rejected_prediction_tokens,omitempty"`
}
// ChatCompletionsChunk is a single streaming chunk from POST /v1/chat/completions.
diff --git a/backend/internal/pkg/claude/constants.go b/backend/internal/pkg/claude/constants.go
index 351f2f8b..dde84724 100644
--- a/backend/internal/pkg/claude/constants.go
+++ b/backend/internal/pkg/claude/constants.go
@@ -134,6 +134,12 @@ var DefaultModels = []Model{
DisplayName: "Claude Opus 4.7",
CreatedAt: "2026-04-17T00:00:00Z",
},
+ {
+ ID: "claude-opus-4-8",
+ Type: "model",
+ DisplayName: "Claude Opus 4.8",
+ CreatedAt: "2026-05-29T00:00:00Z",
+ },
{
ID: "claude-sonnet-4-6",
Type: "model",
diff --git a/backend/internal/pkg/openai/allowed_client.go b/backend/internal/pkg/openai/allowed_client.go
new file mode 100644
index 00000000..d4ca14ee
--- /dev/null
+++ b/backend/internal/pkg/openai/allowed_client.go
@@ -0,0 +1,78 @@
+package openai
+
+import "strings"
+
+// 命名预设 ID。账号侧 codex_cli_only_allowed_clients 只能引用这些预设键,
+// 具体匹配规则固化在下方 registry 中,配置只能「选择启用哪些预设」、不能自定义规则,
+// 以防该白名单退化为可任意放宽的后门。
+const (
+ // AllowedClientClaudeCode 对应 Claude Code CLI 的 codex 插件。
+ AllowedClientClaudeCode = "claude_code"
+)
+
+// AllowedClientEntry 描述一个被额外放行的非官方 Codex 客户端签名。
+// Originator 必须精确等值匹配(归一化后)。
+// UAContains 为必填字段:列表为空,或列表中存在任何空白 marker,均视为非法配置,
+// 整体安全失败(return false);每一项都必须出现在 User-Agent 中。
+// 这确保双因子匹配不会因缺失 UA 声明而退化为仅凭可伪造的 originator 单因子放行。
+type AllowedClientEntry struct {
+ Originator string
+ UAContains []string
+}
+
+// allowedClientRegistry 固化各命名预设的签名规则。
+//
+// Claude Code codex 插件签名来源:插件以 clientInfo.name="Claude Code" 完成 app-server
+// initialize 握手,codex 据此把 originator 设为 "Claude Code",User-Agent 前缀同样为
+// "Claude Code/"(两者同源)。若上游 Claude Code 插件更改 clientInfo.name,此处需同步更新。
+var allowedClientRegistry = map[string]AllowedClientEntry{
+ AllowedClientClaudeCode: {
+ Originator: "Claude Code",
+ UAContains: []string{"Claude Code/"},
+ },
+}
+
+// IsAllowedClientMatch 判断请求头是否命中给定的额外客户端签名。
+// originator 必须精确等值(归一化后);UAContains 中每一项都必须出现在 UA 中。
+// UAContains 为必填:列表为空或含任何空白 marker 均视为非法配置,整体安全失败。
+func IsAllowedClientMatch(userAgent, originator string, entry AllowedClientEntry) bool {
+ wantOriginator := normalizeCodexClientHeader(entry.Originator)
+ if wantOriginator == "" {
+ return false
+ }
+ if normalizeCodexClientHeader(originator) != wantOriginator {
+ return false
+ }
+ // 预设必须声明 UA 特征:否则将退化为仅凭可伪造的 originator 单因子匹配。
+ if len(entry.UAContains) == 0 {
+ return false
+ }
+ ua := normalizeCodexClientHeader(userAgent)
+ for _, marker := range entry.UAContains {
+ normalizedMarker := normalizeCodexClientHeader(marker)
+ if normalizedMarker == "" {
+ // 空白 marker 让该项失去校验能力,会让双因子退化为仅 originator
+ // 单因子;视为非法配置,安全失败。
+ return false
+ }
+ if !strings.Contains(ua, normalizedMarker) {
+ return false
+ }
+ }
+ return true
+}
+
+// MatchAllowedClients 判断请求头是否命中 clientIDs 引用的任一预设签名。
+// 未知预设 ID 会被忽略;空列表恒不放行(默认拒绝)。
+func MatchAllowedClients(userAgent, originator string, clientIDs []string) bool {
+ for _, id := range clientIDs {
+ entry, ok := allowedClientRegistry[normalizeCodexClientHeader(id)]
+ if !ok {
+ continue
+ }
+ if IsAllowedClientMatch(userAgent, originator, entry) {
+ return true
+ }
+ }
+ return false
+}
diff --git a/backend/internal/pkg/openai/allowed_client_test.go b/backend/internal/pkg/openai/allowed_client_test.go
new file mode 100644
index 00000000..c42aa4d5
--- /dev/null
+++ b/backend/internal/pkg/openai/allowed_client_test.go
@@ -0,0 +1,95 @@
+package openai
+
+import "testing"
+
+// 真实的 Claude Code codex 插件请求头:originator 与 UA 前缀同源于 clientInfo.name="Claude Code"。
+const (
+ testClaudeCodeOriginator = "Claude Code"
+ testClaudeCodeUserAgent = "Claude Code/0.5.0 (Macos 15.5; arm64) iTerm2.app (Claude Code; 1.0.4)"
+)
+
+func TestIsAllowedClientMatch(t *testing.T) {
+ entry := AllowedClientEntry{Originator: "Claude Code", UAContains: []string{"Claude Code/"}}
+
+ tests := []struct {
+ name string
+ ua string
+ originator string
+ want bool
+ }{
+ {name: "真实签名命中", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, want: true},
+ {name: "大小写不敏感", ua: "claude code/0.5.0 (macos)", originator: "claude code", want: true},
+ {name: "originator 两侧空白被裁剪", ua: testClaudeCodeUserAgent, originator: " Claude Code ", want: true},
+ {name: "originator 非精确(带后缀)不命中", ua: testClaudeCodeUserAgent, originator: "Claude Code Extra", want: false},
+ {name: "originator 为空不命中", ua: testClaudeCodeUserAgent, originator: "", want: false},
+ {name: "originator 是官方 codex 不命中", ua: testClaudeCodeUserAgent, originator: "codex_cli_rs", want: false},
+ {name: "UA 缺少 Claude Code/ 标记不命中", ua: "curl/8.0", originator: testClaudeCodeOriginator, want: false},
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ if got := IsAllowedClientMatch(tt.ua, tt.originator, entry); got != tt.want {
+ t.Fatalf("IsAllowedClientMatch(%q, %q) = %v, want %v", tt.ua, tt.originator, got, tt.want)
+ }
+ })
+ }
+}
+
+func TestIsAllowedClientMatch_EmptyOriginatorEntryNeverMatches(t *testing.T) {
+ // registry 条目若没有配置 Originator,绝不放行,避免成为宽松后门。
+ entry := AllowedClientEntry{Originator: "", UAContains: []string{"Claude Code/"}}
+ if IsAllowedClientMatch(testClaudeCodeUserAgent, "", entry) {
+ t.Fatal("空 Originator 的条目不应匹配任何请求")
+ }
+}
+
+func TestIsAllowedClientMatch_EmptyUAContainsNeverMatches(t *testing.T) {
+ // 预设必须声明 UA 特征,否则退化为仅凭可伪造的 originator 单因子匹配,绝不放行。
+ entry := AllowedClientEntry{Originator: "Claude Code", UAContains: nil}
+ if IsAllowedClientMatch(testClaudeCodeUserAgent, testClaudeCodeOriginator, entry) {
+ t.Fatal("未声明 UA 特征的预设不应匹配,避免退化为单因子 originator 匹配")
+ }
+}
+
+func TestIsAllowedClientMatch_WhitespaceUAMarkerNeverMatches(t *testing.T) {
+ // 全空白 marker 归一化后为空,若被跳过则退化为仅 originator 单因子;
+ // 任何空白 marker 视为非法预设配置,必须安全失败。
+ entry := AllowedClientEntry{Originator: "Claude Code", UAContains: []string{" "}}
+ if IsAllowedClientMatch(testClaudeCodeUserAgent, testClaudeCodeOriginator, entry) {
+ t.Fatal("UAContains 含全空白 marker 不应匹配,避免退化为单因子 originator 匹配")
+ }
+}
+
+func TestIsAllowedClientMatch_MixedEmptyUAMarkerNeverMatches(t *testing.T) {
+ // 即便 UAContains 含一个真实 marker,只要其中混入任何空白 marker 也视为非法配置;
+ // 防止维护者只为对齐凑数而插入空字符串。
+ entry := AllowedClientEntry{Originator: "Claude Code", UAContains: []string{"", "Claude Code/"}}
+ if IsAllowedClientMatch(testClaudeCodeUserAgent, testClaudeCodeOriginator, entry) {
+ t.Fatal("UAContains 混入空白 marker 不应匹配")
+ }
+}
+
+func TestMatchAllowedClients(t *testing.T) {
+ tests := []struct {
+ name string
+ ua string
+ originator string
+ clientIDs []string
+ want bool
+ }{
+ {name: "claude_code 预设命中真实签名", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: []string{AllowedClientClaudeCode}, want: true},
+ {name: "claude_code 预设 + 伪造 originator 不命中", ua: testClaudeCodeUserAgent, originator: "my_client", clientIDs: []string{AllowedClientClaudeCode}, want: false},
+ {name: "空列表不放行", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: nil, want: false},
+ {name: "未知预设 ID 不放行", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: []string{"unknown_client"}, want: false},
+ {name: "ID 大小写/空白容错", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: []string{" Claude_Code "}, want: true},
+ {name: "多预设任一命中即放行", ua: testClaudeCodeUserAgent, originator: testClaudeCodeOriginator, clientIDs: []string{"unknown_client", AllowedClientClaudeCode}, want: true},
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ if got := MatchAllowedClients(tt.ua, tt.originator, tt.clientIDs); got != tt.want {
+ t.Fatalf("MatchAllowedClients(%q, %q, %v) = %v, want %v", tt.ua, tt.originator, tt.clientIDs, got, tt.want)
+ }
+ })
+ }
+}
diff --git a/backend/internal/repository/account_repo.go b/backend/internal/repository/account_repo.go
index 525abf65..bc970f76 100644
--- a/backend/internal/repository/account_repo.go
+++ b/backend/internal/repository/account_repo.go
@@ -1077,7 +1077,7 @@ func (r *accountRepository) SetRateLimited(ctx context.Context, id int64, resetA
return nil
}
-func (r *accountRepository) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time) error {
+func (r *accountRepository) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error {
if scope == "" {
return nil
}
@@ -1086,6 +1086,11 @@ func (r *accountRepository) SetModelRateLimit(ctx context.Context, id int64, sco
"rate_limited_at": now.Format(time.RFC3339),
"rate_limit_reset_at": resetAt.UTC().Format(time.RFC3339),
}
+ if len(reason) > 0 {
+ if value := strings.TrimSpace(reason[0]); value != "" {
+ payload["reason"] = value
+ }
+ }
raw, err := json.Marshal(payload)
if err != nil {
return err
@@ -1121,6 +1126,7 @@ func (r *accountRepository) SetModelRateLimit(ctx context.Context, id int64, sco
if err := enqueueSchedulerOutbox(ctx, r.sql, service.SchedulerOutboxEventAccountChanged, &id, nil, nil); err != nil {
logger.LegacyPrintf("repository.account", "[SchedulerOutbox] enqueue model rate limit failed: account=%d err=%v", id, err)
}
+ r.syncSchedulerAccountSnapshot(ctx, id)
return nil
}
diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go
index 43b13937..18f6878b 100644
--- a/backend/internal/repository/api_key_repo.go
+++ b/backend/internal/repository/api_key_repo.go
@@ -3,6 +3,7 @@ package repository
import (
"context"
"database/sql"
+ "errors"
"fmt"
"strings"
"time"
@@ -183,6 +184,7 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se
group.FieldAllowMessagesDispatch,
group.FieldDefaultMappedModel,
group.FieldMessagesDispatchModelConfig,
+ group.FieldModelsListConfig,
group.FieldRpmLimit,
)
}).
@@ -303,6 +305,66 @@ func (r *apiKeyRepository) Delete(ctx context.Context, id int64) error {
return nil
}
+// DeleteWithAudit 在同一事务内:
+// 1. 把(明文 key、所有者、key 名称)写入 deleted_api_key_audits;
+// 2. 软删除该 key(tombstone 覆盖 key 列以释放唯一约束)。
+//
+// 保证"被删除的 key 一定能反查到所有者"。事务模式与 group_repo.DeleteCascade 一致。
+func (r *apiKeyRepository) DeleteWithAudit(ctx context.Context, id int64) error {
+ tombstoneKey := fmt.Sprintf("__deleted__%d__%d", id, time.Now().UnixNano())
+
+ tx, err := r.client.Tx(ctx)
+ if err != nil && !errors.Is(err, dbent.ErrTxStarted) {
+ return err
+ }
+ exec := r.client
+ if err == nil {
+ defer func() { _ = tx.Rollback() }()
+ exec = tx.Client()
+ }
+ // err == dbent.ErrTxStarted 时复用当前事务(exec = r.client)。
+
+ // 1. 审计:数据源即 api_keys 当前行;WHERE deleted_at IS NULL 保证只对未删除行写一次。
+ if _, err := exec.ExecContext(ctx, `
+ INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
+ SELECT key, id, user_id, name, NOW()
+ FROM api_keys
+ WHERE id = $1 AND deleted_at IS NULL`, id); err != nil {
+ return err
+ }
+
+ // 2. 软删除(tombstone 覆盖 key)。
+ res, err := exec.ExecContext(ctx, `
+ UPDATE api_keys
+ SET key = $1, deleted_at = NOW(), updated_at = NOW()
+ WHERE id = $2 AND deleted_at IS NULL`, tombstoneKey, id)
+ if err != nil {
+ return err
+ }
+ affected, err := res.RowsAffected()
+ if err != nil {
+ return err
+ }
+ if affected == 0 {
+ // 并发/重复删除:记录已存在(已软删)则幂等返回 nil(defer 回滚空事务),否则 NotFound。
+ exists, existErr := r.client.APIKey.Query().
+ Where(apikey.IDEQ(id)).
+ Exist(mixins.SkipSoftDelete(ctx))
+ if existErr != nil {
+ return existErr
+ }
+ if exists {
+ return nil
+ }
+ return service.ErrAPIKeyNotFound
+ }
+
+ if tx != nil {
+ return tx.Commit()
+ }
+ return nil
+}
+
func (r *apiKeyRepository) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, filters service.APIKeyListFilters) ([]service.APIKey, *pagination.PaginationResult, error) {
q := r.activeQuery().Where(apikey.UserIDEQ(userID))
@@ -678,6 +740,7 @@ func userEntityToService(u *dbent.User) *service.User {
RPMLimit: u.RpmLimit,
CreatedAt: u.CreatedAt,
UpdatedAt: u.UpdatedAt,
+ DeletedAt: u.DeletedAt,
}
// Parse extra emails JSON (supports both old []string and new []NotifyEmailEntry format)
if u.BalanceNotifyExtraEmails != "" && u.BalanceNotifyExtraEmails != "[]" {
@@ -723,6 +786,7 @@ func groupEntityToService(g *dbent.Group) *service.Group {
RequirePrivacySet: g.RequirePrivacySet,
DefaultMappedModel: g.DefaultMappedModel,
MessagesDispatchModelConfig: g.MessagesDispatchModelConfig,
+ ModelsListConfig: g.ModelsListConfig,
RPMLimit: g.RpmLimit,
CreatedAt: g.CreatedAt,
UpdatedAt: g.UpdatedAt,
diff --git a/backend/internal/repository/api_key_repo_integration_test.go b/backend/internal/repository/api_key_repo_integration_test.go
index e926ed86..fdf9bc83 100644
--- a/backend/internal/repository/api_key_repo_integration_test.go
+++ b/backend/internal/repository/api_key_repo_integration_test.go
@@ -555,3 +555,46 @@ func TestIncrementQuotaUsed_Concurrent(t *testing.T) {
require.Equal(t, float64(goroutines)*increment, got.QuotaUsed,
"并发递增后总和应为 %v,实际为 %v", float64(goroutines)*increment, got.QuotaUsed)
}
+
+func (s *APIKeyRepoSuite) TestDeleteWithAudit_WritesAuditAndSoftDeletes() {
+ user := s.mustCreateUser("delwithaudit@test.com")
+ key := &service.APIKey{
+ UserID: user.ID,
+ Key: "sk-del-audit-1",
+ Name: "Audit Me",
+ Status: service.StatusActive,
+ }
+ s.Require().NoError(s.repo.Create(s.ctx, key))
+
+ s.Require().NoError(s.repo.DeleteWithAudit(s.ctx, key.ID))
+
+ _, err := s.repo.GetByID(s.ctx, key.ID)
+ s.Require().Error(err)
+
+ rows, qErr := s.client.QueryContext(s.ctx,
+ `SELECT key, key_name, user_id, api_key_id FROM deleted_api_key_audits WHERE api_key_id = $1`, key.ID)
+ s.Require().NoError(qErr)
+ defer rows.Close()
+ s.Require().True(rows.Next(), "expected one audit row")
+ var auditKey, auditName string
+ var auditUserID, auditAPIKeyID int64
+ s.Require().NoError(rows.Scan(&auditKey, &auditName, &auditUserID, &auditAPIKeyID))
+ s.Require().Equal("sk-del-audit-1", auditKey)
+ s.Require().Equal("Audit Me", auditName)
+ s.Require().Equal(user.ID, auditUserID)
+ s.Require().Equal(key.ID, auditAPIKeyID)
+}
+
+func (s *APIKeyRepoSuite) TestDeleteWithAudit_RepeatIsIdempotent() {
+ user := s.mustCreateUser("delwithaudit-idem@test.com")
+ key := &service.APIKey{UserID: user.ID, Key: "sk-del-audit-2", Name: "K", Status: service.StatusActive}
+ s.Require().NoError(s.repo.Create(s.ctx, key))
+
+ s.Require().NoError(s.repo.DeleteWithAudit(s.ctx, key.ID))
+ s.Require().NoError(s.repo.DeleteWithAudit(s.ctx, key.ID))
+}
+
+func (s *APIKeyRepoSuite) TestDeleteWithAudit_NotFound() {
+ err := s.repo.DeleteWithAudit(s.ctx, 999999)
+ s.Require().ErrorIs(err, service.ErrAPIKeyNotFound)
+}
diff --git a/backend/internal/repository/billing_cache.go b/backend/internal/repository/billing_cache.go
index 60dae954..de229da9 100644
--- a/backend/internal/repository/billing_cache.go
+++ b/backend/internal/repository/billing_cache.go
@@ -7,6 +7,7 @@ import (
"log"
"math/rand/v2"
"strconv"
+ "strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
@@ -338,38 +339,26 @@ func userPlatformQuotaCacheKey(userID int64, platform string) string {
return fmt.Sprintf("billing:user_platform_quota:%d:%s", userID, platform)
}
-func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) (*service.UserPlatformQuotaCacheEntry, bool, error) {
- key := userPlatformQuotaCacheKey(userID, platform)
- fields := []string{
- "daily_usage", "weekly_usage", "monthly_usage", "version", "schema_version",
- "daily_limit", "weekly_limit", "monthly_limit",
- "daily_window_start", "weekly_window_start", "monthly_window_start",
+// parseUserPlatformQuotaHash 将 Redis HGETALL 返回的 map[string]string 反序列化为
+// *service.UserPlatformQuotaCacheEntry。空 map(key 不存在)返回 nil。
+// GetUserPlatformQuotaCache 和 BatchGetUserPlatformQuotaCache 共用此函数,确保解析逻辑一致。
+func parseUserPlatformQuotaHash(m map[string]string) *service.UserPlatformQuotaCacheEntry {
+ if len(m) == 0 {
+ return nil
}
- vals, err := c.rdb.HMGet(ctx, key, fields...).Result()
- if err != nil {
- return nil, false, err
- }
- // 前4个全为nil → key 不存在
- if vals[0] == nil && vals[1] == nil && vals[2] == nil && vals[3] == nil {
- return nil, false, nil
- }
- parseFloat := func(v any) float64 {
- if v == nil {
+ parseFloat := func(s string) float64 {
+ if s == "" {
return 0
}
- s, ok := v.(string)
- if !ok {
+ f, err := strconv.ParseFloat(s, 64)
+ if err != nil {
+ log.Printf("billing_cache: corrupt quota usage field %q (using 0): %v", s, err)
return 0
}
- f, _ := strconv.ParseFloat(s, 64)
return f
}
- parseFloatPtr := func(v any) *float64 {
- if v == nil {
- return nil
- }
- s, ok := v.(string)
- if !ok || s == "" {
+ parseFloatPtr := func(s string) *float64 {
+ if s == "" {
return nil
}
f, err := strconv.ParseFloat(s, 64)
@@ -378,12 +367,8 @@ func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int
}
return &f
}
- parseTimePtr := func(v any) *time.Time {
- if v == nil {
- return nil
- }
- s, ok := v.(string)
- if !ok || s == "" {
+ parseTimePtr := func(s string) *time.Time {
+ if s == "" {
return nil
}
n, err := strconv.ParseInt(s, 10, 64)
@@ -393,30 +378,37 @@ func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int
t := time.Unix(n, 0).UTC()
return &t
}
- parseInt64 := func(v any) int64 {
- if v == nil {
- return 0
- }
- s, ok := v.(string)
- if !ok {
- return 0
- }
+ parseInt64 := func(s string) int64 {
n, _ := strconv.ParseInt(s, 10, 64)
return n
}
return &service.UserPlatformQuotaCacheEntry{
- DailyUsageUSD: parseFloat(vals[0]),
- WeeklyUsageUSD: parseFloat(vals[1]),
- MonthlyUsageUSD: parseFloat(vals[2]),
- Version: parseInt64(vals[3]),
- SchemaVersion: parseInt64(vals[4]),
- DailyLimitUSD: parseFloatPtr(vals[5]),
- WeeklyLimitUSD: parseFloatPtr(vals[6]),
- MonthlyLimitUSD: parseFloatPtr(vals[7]),
- DailyWindowStart: parseTimePtr(vals[8]),
- WeeklyWindowStart: parseTimePtr(vals[9]),
- MonthlyWindowStart: parseTimePtr(vals[10]),
- }, true, nil
+ DailyUsageUSD: parseFloat(m["daily_usage"]),
+ WeeklyUsageUSD: parseFloat(m["weekly_usage"]),
+ MonthlyUsageUSD: parseFloat(m["monthly_usage"]),
+ Version: parseInt64(m["version"]),
+ SchemaVersion: parseInt64(m["schema_version"]),
+ DailyLimitUSD: parseFloatPtr(m["daily_limit"]),
+ WeeklyLimitUSD: parseFloatPtr(m["weekly_limit"]),
+ MonthlyLimitUSD: parseFloatPtr(m["monthly_limit"]),
+ DailyWindowStart: parseTimePtr(m["daily_window_start"]),
+ WeeklyWindowStart: parseTimePtr(m["weekly_window_start"]),
+ MonthlyWindowStart: parseTimePtr(m["monthly_window_start"]),
+ }
+}
+
+func (c *billingCache) GetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) (*service.UserPlatformQuotaCacheEntry, bool, error) {
+ key := userPlatformQuotaCacheKey(userID, platform)
+ m, err := c.rdb.HGetAll(ctx, key).Result()
+ if err != nil {
+ return nil, false, err
+ }
+ entry := parseUserPlatformQuotaHash(m)
+ if entry == nil {
+ // 空 map → key 不存在 → MISS
+ return nil, false, nil
+ }
+ return entry, true, nil
}
func (c *billingCache) SetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string, entry *service.UserPlatformQuotaCacheEntry, ttl time.Duration) error {
@@ -468,9 +460,12 @@ func (c *billingCache) DeleteUserPlatformQuotaCache(ctx context.Context, userID
// SetCache 重建为新版 entry —— 若此处仍累加,上层覆盖时会丢失这部分增量,导致 Redis usage 比真实偏小。
// key 不存在同样跳过(由下次 SetCache 重建)。
// KEYS[1] = hash key
+// KEYS[2] = 脏集 key(dirty set)
// ARGV[1] = cost (string float)
// ARGV[2] = ttl seconds
// ARGV[3] = expected schema_version (Go 侧 UserPlatformQuotaCacheSchemaV1)
+// ARGV[4] = dirty set member(空串则不 SADD)
+// ARGV[5] = 脏集兜底 TTL 秒
const updateUserPlatformQuotaUsageScript = `
if redis.call("EXISTS", KEYS[1]) == 0 then
return 0
@@ -484,18 +479,125 @@ redis.call("HINCRBYFLOAT", KEYS[1], "weekly_usage", ARGV[1])
redis.call("HINCRBYFLOAT", KEYS[1], "monthly_usage", ARGV[1])
redis.call("HINCRBY", KEYS[1], "version", 1)
redis.call("EXPIRE", KEYS[1], ARGV[2])
+if ARGV[4] ~= "" then
+ redis.call("SADD", KEYS[2], ARGV[4])
+ redis.call("EXPIRE", KEYS[2], ARGV[5])
+end
return 1
`
-func (c *billingCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error {
- key := userPlatformQuotaCacheKey(userID, platform)
- _, err := c.rdb.Eval(ctx, updateUserPlatformQuotaUsageScript, []string{key},
+// userPlatformQuotaDirtySetKey 返回脏集(dirty set)的 Redis key。
+// 使用与 userPlatformQuotaCacheKey 相同的前缀 "billing:"。
+func userPlatformQuotaDirtySetKey() string { return "billing:" + "upq:dirty" }
+
+// userPlatformQuotaDirtyTTLSeconds 脏集兜底 TTL(秒):初始 SADD(Lua)与 Readd 共用,
+// 确保 flusher 长期停摆时脏集最终过期;正常运行因持续 SADD 不断续期。
+const userPlatformQuotaDirtyTTLSeconds = 86400
+
+// userPlatformQuotaDirtyMember 构造脏集成员字符串 "userID:platform"。
+func userPlatformQuotaDirtyMember(userID int64, platform string) string {
+ return strconv.FormatInt(userID, 10) + ":" + platform
+}
+
+func (c *billingCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error {
+ member := ""
+ if markDirty {
+ member = userPlatformQuotaDirtyMember(userID, platform)
+ }
+ _, err := c.rdb.Eval(ctx, updateUserPlatformQuotaUsageScript,
+ []string{userPlatformQuotaCacheKey(userID, platform), userPlatformQuotaDirtySetKey()},
strconv.FormatFloat(cost, 'f', -1, 64),
int(ttl.Seconds()),
service.UserPlatformQuotaCacheSchemaV1,
+ member,
+ userPlatformQuotaDirtyTTLSeconds,
).Result()
if err != nil && !errors.Is(err, redis.Nil) {
return err
}
return nil
}
+
+// parseUserPlatformQuotaDirtyMember 将脏集成员字符串 "userID:platform" 解析为
+// service.UserPlatformQuotaKey。解析失败返回 ok=false。
+func parseUserPlatformQuotaDirtyMember(m string) (service.UserPlatformQuotaKey, bool) {
+ parts := strings.SplitN(m, ":", 2)
+ if len(parts) != 2 {
+ return service.UserPlatformQuotaKey{}, false
+ }
+ uid, err := strconv.ParseInt(parts[0], 10, 64)
+ if err != nil {
+ return service.UserPlatformQuotaKey{}, false
+ }
+ return service.UserPlatformQuotaKey{UserID: uid, Platform: parts[1]}, true
+}
+
+// PopDirtyUserPlatformQuotaKeys 从脏集随机弹出最多 n 个 key。
+// 脏集为空时返回 (nil, nil)。
+func (c *billingCache) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]service.UserPlatformQuotaKey, error) {
+ members, err := c.rdb.SPopN(ctx, userPlatformQuotaDirtySetKey(), int64(n)).Result()
+ if err != nil {
+ if errors.Is(err, redis.Nil) {
+ return nil, nil
+ }
+ return nil, err
+ }
+ keys := make([]service.UserPlatformQuotaKey, 0, len(members))
+ for _, m := range members {
+ k, ok := parseUserPlatformQuotaDirtyMember(m)
+ if !ok {
+ log.Printf("billing_cache: skipping invalid dirty member %q", m)
+ continue
+ }
+ keys = append(keys, k)
+ }
+ return keys, nil
+}
+
+// ReaddDirtyUserPlatformQuotaKeys 将 keys 重新加入脏集(flush 失败时回填)。
+// 通过 pipeline 同时执行 SAdd + Expire,确保 Readd 后脏集具有兜底 TTL。
+// 空切片时直接返回 nil。
+func (c *billingCache) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []service.UserPlatformQuotaKey) error {
+ if len(keys) == 0 {
+ return nil
+ }
+ dirtyKey := userPlatformQuotaDirtySetKey()
+ members := make([]any, len(keys))
+ for i, k := range keys {
+ members[i] = userPlatformQuotaDirtyMember(k.UserID, k.Platform)
+ }
+ pipe := c.rdb.Pipeline()
+ pipe.SAdd(ctx, dirtyKey, members...)
+ pipe.Expire(ctx, dirtyKey, userPlatformQuotaDirtyTTLSeconds*time.Second)
+ _, err := pipe.Exec(ctx)
+ return err
+}
+
+// BatchGetUserPlatformQuotaCache 通过 Pipeline 批量 HGETALL 获取多个 user×platform 的
+// quota cache。返回切片与 keys 顺序、长度对齐;MISS 或解析失败位置返回 nil。
+func (c *billingCache) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []service.UserPlatformQuotaKey) ([]*service.UserPlatformQuotaCacheEntry, error) {
+ if len(keys) == 0 {
+ return nil, nil
+ }
+ pipe := c.rdb.Pipeline()
+ cmds := make([]*redis.MapStringStringCmd, len(keys))
+ for i, k := range keys {
+ cmds[i] = pipe.HGetAll(ctx, userPlatformQuotaCacheKey(k.UserID, k.Platform))
+ }
+ if _, err := pipe.Exec(ctx); err != nil && !errors.Is(err, redis.Nil) {
+ return nil, err
+ }
+ results := make([]*service.UserPlatformQuotaCacheEntry, len(keys))
+ for i, cmd := range cmds {
+ m, err := cmd.Result()
+ if err != nil {
+ if !errors.Is(err, redis.Nil) {
+ log.Printf("billing_cache: BatchGet HGETALL cmd[%d] failed: %v (skip, self-heal)", i, err)
+ }
+ // 单个命令失败 → 对应位置 nil,继续
+ continue
+ }
+ results[i] = parseUserPlatformQuotaHash(m)
+ }
+ return results, nil
+}
diff --git a/backend/internal/repository/billing_cache_user_platform_quota_test.go b/backend/internal/repository/billing_cache_user_platform_quota_test.go
index 8d49fd31..15b185e7 100644
--- a/backend/internal/repository/billing_cache_user_platform_quota_test.go
+++ b/backend/internal/repository/billing_cache_user_platform_quota_test.go
@@ -88,7 +88,7 @@ func TestUserPlatformQuotaCache_NilLimitSetThenGet(t *testing.T) {
func TestUserPlatformQuotaCache_IncrMissIsNoop(t *testing.T) {
c, _ := newMiniRedisCache(t)
- if err := c.IncrUserPlatformQuotaUsageCache(context.Background(), 1, "openai", 0.5, time.Minute); err != nil {
+ if err := c.IncrUserPlatformQuotaUsageCache(context.Background(), 1, "openai", 0.5, time.Minute, false); err != nil {
t.Fatal(err)
}
_, ok, _ := c.GetUserPlatformQuotaCache(context.Background(), 1, "openai")
@@ -105,10 +105,10 @@ func TestUserPlatformQuotaCache_IncrHitAccumulates(t *testing.T) {
Version: 1,
SchemaVersion: service.UserPlatformQuotaCacheSchemaV1,
}, time.Minute)
- if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.5, time.Minute); err != nil {
+ if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.5, time.Minute, false); err != nil {
t.Fatal(err)
}
- if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.25, time.Minute); err != nil {
+ if err := c.IncrUserPlatformQuotaUsageCache(ctx, 1, "openai", 0.25, time.Minute, false); err != nil {
t.Fatal(err)
}
got, _, _ := c.GetUserPlatformQuotaCache(ctx, 1, "openai")
diff --git a/backend/internal/repository/content_moderation_repo.go b/backend/internal/repository/content_moderation_repo.go
index 6ada004a..9b19cce9 100644
--- a/backend/internal/repository/content_moderation_repo.go
+++ b/backend/internal/repository/content_moderation_repo.go
@@ -192,6 +192,7 @@ SELECT COUNT(*)
FROM content_moderation_logs
WHERE user_id = $1
AND flagged = TRUE
+ AND action <> 'hash_block'
AND created_at >= $2
AND created_at > COALESCE((SELECT at FROM last_auto_ban), '-infinity'::timestamptz)
`, userID, since).Scan(&count)
@@ -246,7 +247,7 @@ func buildContentModerationLogWhere(filter service.ContentModerationLogFilter) (
case "hit", "flagged":
where = append(where, "l.flagged = TRUE")
case "blocked", "block":
- where = append(where, "l.action = 'block'")
+ where = append(where, "l.action IN ('block', 'keyword_block', 'hash_block')")
case "pass", "allow":
where = append(where, "l.flagged = FALSE AND l.error = ''")
case "error":
diff --git a/backend/internal/repository/content_moderation_repo_test.go b/backend/internal/repository/content_moderation_repo_test.go
new file mode 100644
index 00000000..6d5faa12
--- /dev/null
+++ b/backend/internal/repository/content_moderation_repo_test.go
@@ -0,0 +1,40 @@
+package repository
+
+import (
+ "context"
+ "regexp"
+ "strings"
+ "testing"
+ "time"
+
+ sqlmock "github.com/DATA-DOG/go-sqlmock"
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+func TestBuildContentModerationLogWhere_BlockedIncludesAllBlockActions(t *testing.T) {
+ where, args := buildContentModerationLogWhere(service.ContentModerationLogFilter{Result: "blocked"})
+
+ require.Empty(t, args)
+ sql := strings.Join(where, " AND ")
+ require.Contains(t, sql, "l.action IN ('block', 'keyword_block', 'hash_block')")
+ require.NotContains(t, sql, "l.action = 'block'")
+}
+
+func TestContentModerationRepositoryCountFlaggedByUserSince_ExcludesHashBlock(t *testing.T) {
+ db, mock, err := sqlmock.New()
+ require.NoError(t, err)
+ defer func() { _ = db.Close() }()
+
+ repo := NewContentModerationRepository(db)
+ since := time.Now().Add(-time.Hour)
+ mock.ExpectQuery(regexp.QuoteMeta("AND action <> 'hash_block'")).
+ WithArgs(int64(1001), since).
+ WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(2))
+
+ count, err := repo.CountFlaggedByUserSince(context.Background(), 1001, since)
+
+ require.NoError(t, err)
+ require.Equal(t, 2, count)
+ require.NoError(t, mock.ExpectationsWereMet())
+}
diff --git a/backend/internal/repository/db_pool.go b/backend/internal/repository/db_pool.go
index d7116ab1..e110068c 100644
--- a/backend/internal/repository/db_pool.go
+++ b/backend/internal/repository/db_pool.go
@@ -1,12 +1,26 @@
+// Package repository contains persistence infrastructure helpers.
+//
+// DB pool lifetimes are clamped here because lib/pq starts watchCancel
+// goroutines for context-aware queries. If a cloud proxy silently drops idle
+// TCP without RST/FIN, those goroutines can block in Read until database/sql
+// retires the connection. This is a short-term mitigation; the long-term
+// follow-up is migrating PostgreSQL access to jackc/pgx/v5/stdlib.
package repository
import (
"database/sql"
+ "log/slog"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
)
+const (
+ defaultConnMaxLifetime = 30 * time.Minute
+ defaultConnMaxIdleTime = 5 * time.Minute
+ maxConfiguredConnAge = 24 * time.Hour
+)
+
type dbPoolSettings struct {
MaxOpenConns int
MaxIdleConns int
@@ -14,19 +28,41 @@ type dbPoolSettings struct {
ConnMaxIdleTime time.Duration
}
-func buildDBPoolSettings(cfg *config.Config) dbPoolSettings {
+func clampDBPoolSettings(cfg *config.Config) dbPoolSettings {
return dbPoolSettings{
MaxOpenConns: cfg.Database.MaxOpenConns,
MaxIdleConns: cfg.Database.MaxIdleConns,
- ConnMaxLifetime: time.Duration(cfg.Database.ConnMaxLifetimeMinutes) * time.Minute,
- ConnMaxIdleTime: time.Duration(cfg.Database.ConnMaxIdleTimeMinutes) * time.Minute,
+ ConnMaxLifetime: clampDBPoolDuration("database.conn_max_lifetime_minutes", cfg.Database.ConnMaxLifetimeMinutes, defaultConnMaxLifetime),
+ ConnMaxIdleTime: clampDBPoolDuration("database.conn_max_idle_time_minutes", cfg.Database.ConnMaxIdleTimeMinutes, defaultConnMaxIdleTime),
}
}
+func clampDBPoolDuration(key string, minutes int, fallback time.Duration) time.Duration {
+ if minutes <= 0 || minutes > int(maxConfiguredConnAge/time.Minute) {
+ slog.Warn("database connection pool duration clamped",
+ "key", key,
+ "before", minutes,
+ "after", int(fallback/time.Minute),
+ )
+ return fallback
+ }
+
+ return time.Duration(minutes) * time.Minute
+}
+
func applyDBPoolSettings(db *sql.DB, cfg *config.Config) {
- settings := buildDBPoolSettings(cfg)
+ settings := clampDBPoolSettings(cfg)
db.SetMaxOpenConns(settings.MaxOpenConns)
db.SetMaxIdleConns(settings.MaxIdleConns)
db.SetConnMaxLifetime(settings.ConnMaxLifetime)
db.SetConnMaxIdleTime(settings.ConnMaxIdleTime)
+
+ slog.Info("database connection pool configured",
+ slog.Group("effective",
+ slog.Int("max_open", settings.MaxOpenConns),
+ slog.Int("max_idle", settings.MaxIdleConns),
+ slog.Duration("max_lifetime", settings.ConnMaxLifetime),
+ slog.Duration("max_idle_time", settings.ConnMaxIdleTime),
+ ),
+ )
}
diff --git a/backend/internal/repository/db_pool_test.go b/backend/internal/repository/db_pool_test.go
index 3868106a..2757f97c 100644
--- a/backend/internal/repository/db_pool_test.go
+++ b/backend/internal/repository/db_pool_test.go
@@ -11,21 +11,62 @@ import (
_ "github.com/lib/pq"
)
-func TestBuildDBPoolSettings(t *testing.T) {
- cfg := &config.Config{
- Database: config.DatabaseConfig{
- MaxOpenConns: 50,
- MaxIdleConns: 10,
- ConnMaxLifetimeMinutes: 30,
- ConnMaxIdleTimeMinutes: 5,
+func TestClampDBPoolSettings(t *testing.T) {
+ tests := []struct {
+ name string
+ connMaxLifetime int
+ connMaxIdleTime int
+ wantMaxLifetime time.Duration
+ wantConnMaxIdleTime time.Duration
+ }{
+ {
+ name: "zero values fall back to safe defaults",
+ connMaxLifetime: 0,
+ connMaxIdleTime: 0,
+ wantMaxLifetime: 30 * time.Minute,
+ wantConnMaxIdleTime: 5 * time.Minute,
+ },
+ {
+ name: "negative values fall back to safe defaults",
+ connMaxLifetime: -1,
+ connMaxIdleTime: -5,
+ wantMaxLifetime: 30 * time.Minute,
+ wantConnMaxIdleTime: 5 * time.Minute,
+ },
+ {
+ name: "reasonable values pass through",
+ connMaxLifetime: 15,
+ connMaxIdleTime: 3,
+ wantMaxLifetime: 15 * time.Minute,
+ wantConnMaxIdleTime: 3 * time.Minute,
+ },
+ {
+ name: "values over twenty four hours fall back to safe defaults",
+ connMaxLifetime: 24*60 + 1,
+ connMaxIdleTime: 24*60 + 1,
+ wantMaxLifetime: 30 * time.Minute,
+ wantConnMaxIdleTime: 5 * time.Minute,
},
}
- settings := buildDBPoolSettings(cfg)
- require.Equal(t, 50, settings.MaxOpenConns)
- require.Equal(t, 10, settings.MaxIdleConns)
- require.Equal(t, 30*time.Minute, settings.ConnMaxLifetime)
- require.Equal(t, 5*time.Minute, settings.ConnMaxIdleTime)
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ cfg := &config.Config{
+ Database: config.DatabaseConfig{
+ MaxOpenConns: 50,
+ MaxIdleConns: 10,
+ ConnMaxLifetimeMinutes: tt.connMaxLifetime,
+ ConnMaxIdleTimeMinutes: tt.connMaxIdleTime,
+ },
+ }
+
+ settings := clampDBPoolSettings(cfg)
+ require.Equal(t, 50, settings.MaxOpenConns)
+ require.Equal(t, 10, settings.MaxIdleConns)
+ require.Equal(t, tt.wantMaxLifetime, settings.ConnMaxLifetime)
+ require.Equal(t, tt.wantConnMaxIdleTime, settings.ConnMaxIdleTime)
+ })
+ }
}
func TestApplyDBPoolSettings(t *testing.T) {
diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go
index 9c3b2010..ac8669ab 100644
--- a/backend/internal/repository/group_repo.go
+++ b/backend/internal/repository/group_repo.go
@@ -66,6 +66,7 @@ func (r *groupRepository) Create(ctx context.Context, groupIn *service.Group) er
SetRequirePrivacySet(groupIn.RequirePrivacySet).
SetDefaultMappedModel(groupIn.DefaultMappedModel).
SetMessagesDispatchModelConfig(groupIn.MessagesDispatchModelConfig).
+ SetModelsListConfig(groupIn.ModelsListConfig).
SetRpmLimit(groupIn.RPMLimit)
// 设置模型路由配置
@@ -141,6 +142,7 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er
SetRequirePrivacySet(groupIn.RequirePrivacySet).
SetDefaultMappedModel(groupIn.DefaultMappedModel).
SetMessagesDispatchModelConfig(groupIn.MessagesDispatchModelConfig).
+ SetModelsListConfig(groupIn.ModelsListConfig).
SetRpmLimit(groupIn.RPMLimit)
// 显式处理可空字段:nil 需要 clear,非 nil 需要 set。
diff --git a/backend/internal/repository/ops_error_where_test.go b/backend/internal/repository/ops_error_where_test.go
new file mode 100644
index 00000000..9bebb158
--- /dev/null
+++ b/backend/internal/repository/ops_error_where_test.go
@@ -0,0 +1,95 @@
+package repository
+
+import (
+ "strings"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+)
+
+func TestBuildOpsErrorLogsWhere_UserScopedFilters(t *testing.T) {
+ uid := int64(42)
+ kid := int64(7)
+ filter := &service.OpsErrorLogFilter{
+ UserID: &uid,
+ APIKeyID: &kid,
+ Model: "claude-sonnet-4-5",
+ ExcludeCountTokens: true,
+ ErrorPhasesAny: []string{"auth"},
+ ErrorTypesAny: []string{"rate_limit_error"},
+ View: "all",
+ }
+ where, args := buildOpsErrorLogsWhere(filter)
+
+ for _, want := range []string{
+ "e.user_id = $",
+ "e.api_key_id = $",
+ "COALESCE(e.requested_model, e.model, '') = $",
+ "COALESCE(e.is_count_tokens, false) = false",
+ "e.error_phase = ANY($",
+ "e.error_type = ANY($",
+ } {
+ if !strings.Contains(where, want) {
+ t.Fatalf("where missing %q\nfull: %s", want, where)
+ }
+ }
+ if len(args) != 5 {
+ t.Fatalf("expected 5 args, got %d", len(args))
+ }
+}
+
+func TestBuildOpsErrorLogsWhere_ModelFuzzy(t *testing.T) {
+ // 默认(ModelFuzzy=false)保持精确匹配
+ exact := &service.OpsErrorLogFilter{Model: "claude"}
+ whereExact, _ := buildOpsErrorLogsWhere(exact)
+ if !strings.Contains(whereExact, "COALESCE(e.requested_model, e.model, '') = $") {
+ t.Fatalf("default should be exact match, got: %s", whereExact)
+ }
+
+ // ModelFuzzy=true → ILIKE
+ fuzzy := &service.OpsErrorLogFilter{Model: "claude", ModelFuzzy: true}
+ whereFuzzy, args := buildOpsErrorLogsWhere(fuzzy)
+ if !strings.Contains(whereFuzzy, "COALESCE(e.requested_model, e.model, '') ILIKE $") {
+ t.Fatalf("ModelFuzzy should use ILIKE, got: %s", whereFuzzy)
+ }
+ if len(args) != 1 || args[0] != "%claude%" {
+ t.Fatalf("expected arg \"%%claude%%\", got %v", args)
+ }
+
+ // 通配符转义:输入含 % 应被转义为字面量
+ esc := &service.OpsErrorLogFilter{Model: "50%off", ModelFuzzy: true}
+ _, escArgs := buildOpsErrorLogsWhere(esc)
+ if len(escArgs) != 1 || escArgs[0] != `%50\%off%` {
+ t.Fatalf("expected escaped arg, got %v", escArgs)
+ }
+
+ esc2 := &service.OpsErrorLogFilter{Model: "gpt_4o", ModelFuzzy: true}
+ _, escArgs2 := buildOpsErrorLogsWhere(esc2)
+ if len(escArgs2) != 1 || escArgs2[0] != `%gpt\_4o%` {
+ t.Fatalf("underscore should be escaped, got %v", escArgs2)
+ }
+}
+
+func TestBuildOpsErrorLogsWhere_MatchDeletedKeyOwner(t *testing.T) {
+ uid := int64(42)
+
+ // 开关开启 → 归属放宽为 OR(user_id 或 deleted_key_owner_user_id),且共用同一占位符
+ on := &service.OpsErrorLogFilter{UserID: &uid, MatchDeletedKeyOwner: true}
+ whereOn, argsOn := buildOpsErrorLogsWhere(on)
+ if !strings.Contains(whereOn, "(e.user_id = $1 OR e.deleted_key_owner_user_id = $1)") {
+ t.Fatalf("MatchDeletedKeyOwner=true should widen to OR, got: %s", whereOn)
+ }
+ if len(argsOn) != 1 || argsOn[0] != uid {
+ t.Fatalf("expected single reused arg %d, got %v", uid, argsOn)
+ }
+
+ // 开关关闭(默认)→ 仅精确 user_id,绝不出现 deleted_key_owner_user_id(admin 回归)
+ off := &service.OpsErrorLogFilter{UserID: &uid}
+ whereOff, _ := buildOpsErrorLogsWhere(off)
+ if !strings.Contains(whereOff, "e.user_id = $1") {
+ t.Fatalf("default should match user_id exactly, got: %s", whereOff)
+ }
+ if strings.Contains(whereOff, "deleted_key_owner_user_id") {
+ t.Fatalf("default must NOT include deleted_key_owner_user_id, got: %s", whereOff)
+ }
+}
diff --git a/backend/internal/repository/ops_repo.go b/backend/internal/repository/ops_repo.go
index 4371b8a2..f300a171 100644
--- a/backend/internal/repository/ops_repo.go
+++ b/backend/internal/repository/ops_repo.go
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"encoding/json"
+ "errors"
"fmt"
"strings"
"time"
@@ -54,9 +55,13 @@ INSERT INTO ops_error_logs (
upstream_latency_ms,
response_latency_ms,
time_to_first_token_ms,
- created_at
+ created_at,
+ attempted_key_prefix,
+ deleted_key_owner_user_id,
+ deleted_key_name,
+ api_key_prefix
) VALUES (
- $1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37
+ $1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26,$27,$28,$29,$30,$31,$32,$33,$34,$35,$36,$37,$38,$39,$40,$41
)`
func NewOpsRepository(db *sql.DB) service.OpsRepository {
@@ -165,6 +170,10 @@ func opsInsertErrorLogArgs(input *service.OpsInsertErrorLogInput) []any {
opsNullInt64(input.ResponseLatencyMs),
opsNullInt64(input.TimeToFirstTokenMs),
input.CreatedAt,
+ opsNullString(input.AttemptedKeyPrefix),
+ opsNullInt64(input.DeletedKeyOwnerUserID),
+ opsNullString(input.DeletedKeyName),
+ opsNullString(input.APIKeyPrefix),
}
}
@@ -231,12 +240,16 @@ SELECT
COALESCE(e.upstream_endpoint, ''),
COALESCE(e.requested_model, ''),
COALESCE(e.upstream_model, ''),
- e.request_type
+ e.request_type,
+ COALESCE(ak.name, ''),
+ ak.deleted_at,
+ COALESCE(e.deleted_key_name, '')
FROM ops_error_logs e
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN users u2 ON e.resolved_by_user_id = u2.id
+LEFT JOIN api_keys ak ON ak.id = e.api_key_id
` + where + `
ORDER BY e.created_at DESC
LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
@@ -263,6 +276,9 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
var resolvedBy sql.NullInt64
var resolvedByName string
var requestType sql.NullInt64
+ var apiKeyName string
+ var apiKeyDeletedAt sql.NullTime
+ var deletedKeyName string
if err := rows.Scan(
&item.ID,
&item.CreatedAt,
@@ -296,6 +312,9 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
&item.RequestedModel,
&item.UpstreamModel,
&requestType,
+ &apiKeyName,
+ &apiKeyDeletedAt,
+ &deletedKeyName,
); err != nil {
return nil, err
}
@@ -336,6 +355,15 @@ LIMIT $` + itoa(len(args)+1) + ` OFFSET $` + itoa(len(args)+2)
v := int16(requestType.Int64)
item.RequestType = &v
}
+ // Key 名称:优先关联到的 ak.name(已软删的 key name 仍保留);
+ // 关联不到(api_key_id 为空 / 历史硬删)时回退错误记录里快照的 deleted_key_name。
+ if apiKeyName != "" {
+ item.APIKeyName = apiKeyName
+ } else {
+ item.APIKeyName = deletedKeyName
+ }
+ // 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
+ item.APIKeyDeleted = apiKeyDeletedAt.Valid || (apiKeyName == "" && deletedKeyName != "")
out = append(out, &item)
}
if err := rows.Err(); err != nil {
@@ -402,11 +430,20 @@ SELECT
e.routing_latency_ms,
e.upstream_latency_ms,
e.response_latency_ms,
- e.time_to_first_token_ms
+ e.time_to_first_token_ms,
+ COALESCE(e.attempted_key_prefix, ''),
+ e.deleted_key_owner_user_id,
+ COALESCE(du.email, ''),
+ COALESCE(e.deleted_key_name, ''),
+ COALESCE(e.api_key_prefix, ''),
+ COALESCE(ak.name, ''),
+ ak.deleted_at
FROM ops_error_logs e
LEFT JOIN users u ON e.user_id = u.id
LEFT JOIN accounts a ON e.account_id = a.id
LEFT JOIN groups g ON e.group_id = g.id
+LEFT JOIN users du ON e.deleted_key_owner_user_id = du.id
+LEFT JOIN api_keys ak ON ak.id = e.api_key_id
WHERE e.id = $1
LIMIT 1`
@@ -426,6 +463,9 @@ LIMIT 1`
var responseLatency sql.NullInt64
var ttft sql.NullInt64
var requestType sql.NullInt64
+ var deletedKeyOwnerUserID sql.NullInt64
+ var detailAPIKeyName string
+ var detailAPIKeyDeletedAt sql.NullTime
err := r.db.QueryRowContext(ctx, q, id).Scan(
&out.ID,
@@ -471,6 +511,13 @@ LIMIT 1`
&upstreamLatency,
&responseLatency,
&ttft,
+ &out.AttemptedKeyPrefix,
+ &deletedKeyOwnerUserID,
+ &out.DeletedKeyOwnerEmail,
+ &out.DeletedKeyName,
+ &out.APIKeyPrefix,
+ &detailAPIKeyName,
+ &detailAPIKeyDeletedAt,
)
if err != nil {
return nil, err
@@ -533,6 +580,18 @@ LIMIT 1`
v := int16(requestType.Int64)
out.RequestType = &v
}
+ if deletedKeyOwnerUserID.Valid {
+ v := deletedKeyOwnerUserID.Int64
+ out.DeletedKeyOwnerUserID = &v
+ }
+ // Key 名称:优先关联到的 ak.name;关联不到时回退快照的 deleted_key_name。
+ if detailAPIKeyName != "" {
+ out.APIKeyName = detailAPIKeyName
+ } else {
+ out.APIKeyName = out.DeletedKeyName
+ }
+ // 已删除:ak.deleted_at 非空(软删),或仅命中 deleted_key_name 兜底。
+ out.APIKeyDeleted = detailAPIKeyDeletedAt.Valid || (detailAPIKeyName == "" && out.DeletedKeyName != "")
// Normalize upstream_errors to empty string when stored as JSON null.
out.UpstreamErrors = strings.TrimSpace(out.UpstreamErrors)
@@ -543,6 +602,26 @@ LIMIT 1`
return &out, nil
}
+// LookupDeletedKeyAudit 按明文 key 反查最近一条已删除 key 审计。
+// 同一 key 可能有多条历史(反复创建/删除),取 deleted_at 最近一条(id 作同毫秒 tiebreaker)。
+// 未命中返回 (nil, nil)。
+func (r *opsRepository) LookupDeletedKeyAudit(ctx context.Context, key string) (*service.DeletedKeyAuditResult, error) {
+ var res service.DeletedKeyAuditResult
+ err := r.db.QueryRowContext(ctx, `
+ SELECT user_id, key_name
+ FROM deleted_api_key_audits
+ WHERE key = $1
+ ORDER BY deleted_at DESC, id DESC
+ LIMIT 1`, key).Scan(&res.UserID, &res.KeyName)
+ if err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return nil, nil
+ }
+ return nil, err
+ }
+ return &res, nil
+}
+
func (r *opsRepository) UpdateErrorResolution(ctx context.Context, errorID int64, resolved bool, resolvedByUserID *int64, resolvedAt *time.Time) error {
if r == nil || r.db == nil {
return fmt.Errorf("nil ops repository")
@@ -815,6 +894,14 @@ INSERT INTO ops_system_log_cleanup_audits (
return err
}
+var likePatternReplacer = strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`)
+
+// escapeLikePattern 转义 LIKE/ILIKE 通配符(\ % _),避免用户输入被当作通配符。
+// Postgres 默认以反斜杠为转义符,无需额外 ESCAPE 子句。
+func escapeLikePattern(s string) string {
+ return likePatternReplacer.Replace(s)
+}
+
func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) {
clauses := make([]string, 0, 12)
args := make([]any, 0, 12)
@@ -927,6 +1014,41 @@ func buildOpsErrorLogsWhere(filter *service.OpsErrorLogFilter) (string, []any) {
clauses = append(clauses, "EXISTS (SELECT 1 FROM users u WHERE u.id = e.user_id AND u.email ILIKE $"+n+")")
}
+ if filter.UserID != nil && *filter.UserID > 0 {
+ args = append(args, *filter.UserID)
+ n := itoa(len(args))
+ if filter.MatchDeletedKeyOwner {
+ // 用户侧:把「删 key 后认证失败」(user_id=NULL,靠 deleted_key_owner 归因)的记录也纳入。
+ clauses = append(clauses, "(e.user_id = $"+n+" OR e.deleted_key_owner_user_id = $"+n+")")
+ } else {
+ clauses = append(clauses, "e.user_id = $"+n)
+ }
+ }
+ if filter.APIKeyID != nil && *filter.APIKeyID > 0 {
+ args = append(args, *filter.APIKeyID)
+ clauses = append(clauses, "e.api_key_id = $"+itoa(len(args)))
+ }
+ if m := strings.TrimSpace(filter.Model); m != "" {
+ if filter.ModelFuzzy {
+ args = append(args, "%"+escapeLikePattern(m)+"%")
+ clauses = append(clauses, "COALESCE(e.requested_model, e.model, '') ILIKE $"+itoa(len(args)))
+ } else {
+ args = append(args, m)
+ clauses = append(clauses, "COALESCE(e.requested_model, e.model, '') = $"+itoa(len(args)))
+ }
+ }
+ if filter.ExcludeCountTokens {
+ clauses = append(clauses, "COALESCE(e.is_count_tokens, false) = false")
+ }
+ if len(filter.ErrorPhasesAny) > 0 {
+ args = append(args, pq.Array(filter.ErrorPhasesAny))
+ clauses = append(clauses, "e.error_phase = ANY($"+itoa(len(args))+")")
+ }
+ if len(filter.ErrorTypesAny) > 0 {
+ args = append(args, pq.Array(filter.ErrorTypesAny))
+ clauses = append(clauses, "e.error_type = ANY($"+itoa(len(args))+")")
+ }
+
return "WHERE " + strings.Join(clauses, " AND "), args
}
diff --git a/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go b/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go
new file mode 100644
index 00000000..470b1c0d
--- /dev/null
+++ b/backend/internal/repository/ops_repo_get_error_log_by_id_integration_test.go
@@ -0,0 +1,94 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+// TestGetErrorLogByID_DeletedKeyOwner 验证:
+// 1. 带 deleted_key_owner_user_id 的记录能正确 JOIN users 返回 DeletedKeyOwnerEmail
+// 2. 新列全为 NULL 的普通记录 Scan 不报错,这些字段为空/nil
+func TestGetErrorLogByID_DeletedKeyOwner(t *testing.T) {
+ ctx := context.Background()
+ _, _ = integrationDB.ExecContext(ctx, "TRUNCATE ops_error_logs RESTART IDENTITY CASCADE")
+
+ repo := NewOpsRepository(integrationDB).(*opsRepository)
+
+ // ── Case 1: 带 deleted_key_owner 信息的记录 ──────────────────────────────
+ owner := mustCreateUser(t, integrationEntClient, &service.User{
+ Email: "deleted-key-owner-" + time.Now().Format("150405.000000000") + "@example.com",
+ })
+
+ var insertedID int64
+ err := integrationDB.QueryRowContext(ctx, `
+ INSERT INTO ops_error_logs (
+ error_phase, error_type, severity, status_code, created_at,
+ attempted_key_prefix, deleted_key_owner_user_id, deleted_key_name
+ ) VALUES (
+ 'auth', 'INVALID_API_KEY', 'error', 401, NOW(),
+ 'sk-test-abc', $1, 'my-deleted-key'
+ ) RETURNING id`,
+ owner.ID,
+ ).Scan(&insertedID)
+ require.NoError(t, err)
+ require.Positive(t, insertedID)
+
+ detail, err := repo.GetErrorLogByID(ctx, insertedID)
+ require.NoError(t, err)
+ require.NotNil(t, detail)
+
+ require.Equal(t, "sk-test-abc", detail.AttemptedKeyPrefix)
+ require.NotNil(t, detail.DeletedKeyOwnerUserID)
+ require.Equal(t, owner.ID, *detail.DeletedKeyOwnerUserID)
+ require.Equal(t, owner.Email, detail.DeletedKeyOwnerEmail)
+ require.Equal(t, "my-deleted-key", detail.DeletedKeyName)
+
+ // ── Case 2: 新列全为 NULL 的普通错误记录 ──────────────────────────────────
+ var plainID int64
+ err = integrationDB.QueryRowContext(ctx, `
+ INSERT INTO ops_error_logs (
+ error_phase, error_type, severity, status_code, created_at
+ ) VALUES (
+ 'upstream', 'upstream_error', 'error', 500, NOW()
+ ) RETURNING id`,
+ ).Scan(&plainID)
+ require.NoError(t, err)
+ require.Positive(t, plainID)
+
+ plain, err := repo.GetErrorLogByID(ctx, plainID)
+ require.NoError(t, err)
+ require.NotNil(t, plain)
+
+ require.Empty(t, plain.AttemptedKeyPrefix, "no prefix for plain error")
+ require.Nil(t, plain.DeletedKeyOwnerUserID, "no owner for plain error")
+ require.Empty(t, plain.DeletedKeyOwnerEmail, "no owner email for plain error")
+ require.Empty(t, plain.DeletedKeyName, "no key name for plain error")
+ require.Empty(t, plain.APIKeyPrefix, "no api key prefix for plain error")
+
+ // ── Case 3: 有效(未删除)key 报错,经 InsertErrorLog 快照 api_key_prefix ──────
+ // 走真实 InsertErrorLog 写入路径(覆盖新列 + $41 占位符),再 GetErrorLogByID 读回。
+ validID, err := repo.InsertErrorLog(ctx, &service.OpsInsertErrorLogInput{
+ ErrorPhase: "request",
+ ErrorType: "api_error",
+ Severity: "error",
+ StatusCode: 402,
+ CreatedAt: time.Now(),
+ APIKeyPrefix: "sk-valid",
+ })
+ require.NoError(t, err)
+ require.Positive(t, validID)
+
+ valid, err := repo.GetErrorLogByID(ctx, validID)
+ require.NoError(t, err)
+ require.NotNil(t, valid)
+
+ require.Equal(t, "sk-valid", valid.APIKeyPrefix)
+ require.Empty(t, valid.AttemptedKeyPrefix, "attempted prefix and api key prefix are mutually exclusive")
+ require.Nil(t, valid.DeletedKeyOwnerUserID, "valid key error has no deleted owner")
+}
diff --git a/backend/internal/repository/ops_repo_lookup_deleted_key_audit_integration_test.go b/backend/internal/repository/ops_repo_lookup_deleted_key_audit_integration_test.go
new file mode 100644
index 00000000..c77aefb9
--- /dev/null
+++ b/backend/internal/repository/ops_repo_lookup_deleted_key_audit_integration_test.go
@@ -0,0 +1,36 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestOpsRepositoryLookupDeletedKeyAudit(t *testing.T) {
+ ctx := context.Background()
+ _, _ = integrationDB.ExecContext(ctx, "TRUNCATE deleted_api_key_audits RESTART IDENTITY")
+ repo := NewOpsRepository(integrationDB).(*opsRepository)
+
+ // 同一 key 两条审计,取最近一条(deleted_at DESC, id DESC)
+ _, err := integrationDB.ExecContext(ctx, `
+ INSERT INTO deleted_api_key_audits (key, api_key_id, user_id, key_name, deleted_at)
+ VALUES ('sk-lookup-1', 10, 100, 'old', $1),
+ ('sk-lookup-1', 11, 200, 'new', $2)`,
+ time.Now().Add(-time.Hour), time.Now())
+ require.NoError(t, err)
+
+ res, err := repo.LookupDeletedKeyAudit(ctx, "sk-lookup-1")
+ require.NoError(t, err)
+ require.NotNil(t, res)
+ require.Equal(t, int64(200), res.UserID)
+ require.Equal(t, "new", res.KeyName)
+
+ // 未命中返回 nil
+ miss, err := repo.LookupDeletedKeyAudit(ctx, "sk-never-existed")
+ require.NoError(t, err)
+ require.Nil(t, miss)
+}
diff --git a/backend/internal/repository/scheduler_cache.go b/backend/internal/repository/scheduler_cache.go
index ab01a863..921aa081 100644
--- a/backend/internal/repository/scheduler_cache.go
+++ b/backend/internal/repository/scheduler_cache.go
@@ -548,6 +548,18 @@ func filterSchedulerExtra(extra map[string]any) map[string]any {
"openai_ws_force_http",
"openai_responses_mode",
"openai_responses_supported",
+ "codex_5h_used_percent",
+ "codex_7d_used_percent",
+ "codex_5h_reset_at",
+ "codex_7d_reset_at",
+ "codex_5h_reset_after_seconds",
+ "codex_7d_reset_after_seconds",
+ "codex_usage_updated_at",
+ "auto_pause_5h_threshold",
+ "auto_pause_7d_threshold",
+ "auto_pause_5h_disabled",
+ "auto_pause_7d_disabled",
+ "model_rate_limits",
}
filtered := make(map[string]any)
for _, key := range keys {
diff --git a/backend/internal/repository/scheduler_cache_unit_test.go b/backend/internal/repository/scheduler_cache_unit_test.go
index 86de87c7..c14721cd 100644
--- a/backend/internal/repository/scheduler_cache_unit_test.go
+++ b/backend/internal/repository/scheduler_cache_unit_test.go
@@ -75,3 +75,62 @@ func TestBuildSchedulerMetadataAccount_KeepsSlimGroupMembership(t *testing.T) {
require.Equal(t, int64(11), got.AccountGroups[1].GroupID)
require.Nil(t, got.Groups)
}
+
+func TestBuildSchedulerMetadataAccount_KeepsQuotaAutoPauseFields(t *testing.T) {
+ account := service.Account{
+ ID: 88,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 12.34,
+ "codex_7d_used_percent": 56.78,
+ "codex_5h_reset_at": "2026-05-29T10:00:00Z",
+ "codex_7d_reset_at": "2026-06-01T10:00:00Z",
+ "codex_5h_reset_after_seconds": 300,
+ "codex_7d_reset_after_seconds": 600,
+ "codex_usage_updated_at": "2026-05-29T09:00:00Z",
+ "auto_pause_5h_threshold": 0.95,
+ "auto_pause_7d_threshold": 0.96,
+ "auto_pause_5h_disabled": true,
+ "auto_pause_7d_disabled": false,
+ },
+ }
+
+ got := buildSchedulerMetadataAccount(account)
+
+ require.Equal(t, 12.34, got.Extra["codex_5h_used_percent"])
+ require.Equal(t, 56.78, got.Extra["codex_7d_used_percent"])
+ require.Equal(t, "2026-05-29T10:00:00Z", got.Extra["codex_5h_reset_at"])
+ require.Equal(t, "2026-06-01T10:00:00Z", got.Extra["codex_7d_reset_at"])
+ require.Equal(t, 300, got.Extra["codex_5h_reset_after_seconds"])
+ require.Equal(t, 600, got.Extra["codex_7d_reset_after_seconds"])
+ require.Equal(t, "2026-05-29T09:00:00Z", got.Extra["codex_usage_updated_at"])
+ require.Equal(t, 0.95, got.Extra["auto_pause_5h_threshold"])
+ require.Equal(t, 0.96, got.Extra["auto_pause_7d_threshold"])
+ require.Equal(t, true, got.Extra["auto_pause_5h_disabled"])
+ require.Equal(t, false, got.Extra["auto_pause_7d_disabled"])
+}
+
+func TestBuildSchedulerMetadataAccount_KeepsModelRateLimits(t *testing.T) {
+ account := service.Account{
+ ID: 90,
+ Platform: service.PlatformAntigravity,
+ Extra: map[string]any{
+ "model_rate_limits": map[string]any{
+ "gemini-3-flash": map[string]any{
+ "rate_limit_reset_at": "2026-05-30T10:10:00Z",
+ },
+ "antigravity:gemini": map[string]any{
+ "rate_limit_reset_at": "2026-05-30T10:10:00Z",
+ },
+ },
+ "unused_large_field": "drop-me",
+ },
+ }
+
+ got := buildSchedulerMetadataAccount(account)
+
+ limits, ok := got.Extra["model_rate_limits"].(map[string]any)
+ require.True(t, ok)
+ require.Contains(t, limits, "gemini-3-flash")
+ require.Contains(t, limits, "antigravity:gemini")
+ require.Nil(t, got.Extra["unused_large_field"])
+}
diff --git a/backend/internal/repository/usage_log_repo.go b/backend/internal/repository/usage_log_repo.go
index f11910a0..b0992dae 100644
--- a/backend/internal/repository/usage_log_repo.go
+++ b/backend/internal/repository/usage_log_repo.go
@@ -17,6 +17,7 @@ import (
dbaccount "github.com/Wei-Shaw/sub2api/ent/account"
dbapikey "github.com/Wei-Shaw/sub2api/ent/apikey"
dbgroup "github.com/Wei-Shaw/sub2api/ent/group"
+ "github.com/Wei-Shaw/sub2api/ent/schema/mixins"
dbuser "github.com/Wei-Shaw/sub2api/ent/user"
dbusersub "github.com/Wei-Shaw/sub2api/ent/usersubscription"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
@@ -26,6 +27,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/lib/pq"
gocache "github.com/patrickmn/go-cache"
+ "golang.org/x/sync/errgroup"
)
const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at"
@@ -3537,24 +3539,6 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us
stats := &UsageStats{}
var totalAccountCost float64
- if err := scanSingleRow(
- ctx,
- r.sql,
- query,
- args,
- &stats.TotalRequests,
- &stats.TotalInputTokens,
- &stats.TotalOutputTokens,
- &stats.TotalCacheTokens,
- &stats.TotalCost,
- &stats.TotalActualCost,
- &totalAccountCost,
- &stats.AverageDurationMs,
- ); err != nil {
- return nil, err
- }
- stats.TotalAccountCost = &totalAccountCost
- stats.TotalTokens = stats.TotalInputTokens + stats.TotalOutputTokens + stats.TotalCacheTokens
start := time.Unix(0, 0).UTC()
if filters.StartTime != nil {
@@ -3565,21 +3549,76 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us
end = *filters.EndTime
}
- endpoints, endpointErr := r.GetEndpointStatsWithFilters(ctx, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
- if endpointErr != nil {
- logger.LegacyPrintf("repository.usage_log", "GetEndpointStatsWithFilters failed in GetStatsWithFilters: %v", endpointErr)
- endpoints = []EndpointStat{}
+ var endpoints, upstreamEndpoints, endpointPaths []EndpointStat
+
+ // 汇总查询:失败即致命。
+ runSummary := func(c context.Context) error {
+ return scanSingleRow(
+ c, r.sql, query, args,
+ &stats.TotalRequests,
+ &stats.TotalInputTokens,
+ &stats.TotalOutputTokens,
+ &stats.TotalCacheTokens,
+ &stats.TotalCost,
+ &stats.TotalActualCost,
+ &totalAccountCost,
+ &stats.AverageDurationMs,
+ )
}
- upstreamEndpoints, upstreamEndpointErr := r.GetUpstreamEndpointStatsWithFilters(ctx, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
- if upstreamEndpointErr != nil {
- logger.LegacyPrintf("repository.usage_log", "GetUpstreamEndpointStatsWithFilters failed in GetStatsWithFilters: %v", upstreamEndpointErr)
- upstreamEndpoints = []EndpointStat{}
+ // endpoint 明细:best-effort(失败 log + 返空),不致命。
+ runEndpoints := func(c context.Context) {
+ res, err := r.GetEndpointStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
+ if err != nil {
+ if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
+ logger.LegacyPrintf("repository.usage_log", "GetEndpointStatsWithFilters failed in GetStatsWithFilters: %v", err)
+ }
+ res = []EndpointStat{}
+ }
+ endpoints = res
}
- endpointPaths, endpointPathErr := r.getEndpointPathStatsWithFilters(ctx, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
- if endpointPathErr != nil {
- logger.LegacyPrintf("repository.usage_log", "getEndpointPathStatsWithFilters failed in GetStatsWithFilters: %v", endpointPathErr)
- endpointPaths = []EndpointStat{}
+ runUpstream := func(c context.Context) {
+ res, err := r.GetUpstreamEndpointStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
+ if err != nil {
+ if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
+ logger.LegacyPrintf("repository.usage_log", "GetUpstreamEndpointStatsWithFilters failed in GetStatsWithFilters: %v", err)
+ }
+ res = []EndpointStat{}
+ }
+ upstreamEndpoints = res
}
+ runPaths := func(c context.Context) {
+ res, err := r.getEndpointPathStatsWithFilters(c, start, end, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType)
+ if err != nil {
+ if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
+ logger.LegacyPrintf("repository.usage_log", "getEndpointPathStatsWithFilters failed in GetStatsWithFilters: %v", err)
+ }
+ res = []EndpointStat{}
+ }
+ endpointPaths = res
+ }
+
+ if r.db != nil {
+ // 生产路径:r.sql 是 *sql.DB 连接池,可并发。4 条查询并行,延迟取最大值。
+ g, gctx := errgroup.WithContext(ctx)
+ g.Go(func() error { return runSummary(gctx) })
+ g.Go(func() error { runEndpoints(gctx); return nil })
+ g.Go(func() error { runUpstream(gctx); return nil })
+ g.Go(func() error { runPaths(gctx); return nil })
+ if err := g.Wait(); err != nil {
+ return nil, err
+ }
+ } else {
+ // 事务路径(ent.Tx 不能并发查询):顺序执行,行为与重构前一致。
+ if err := runSummary(ctx); err != nil {
+ return nil, err
+ }
+ runEndpoints(ctx)
+ runUpstream(ctx)
+ runPaths(ctx)
+ }
+
+ stats.TotalAccountCost = &totalAccountCost
+ stats.TotalTokens = stats.TotalInputTokens + stats.TotalOutputTokens + stats.TotalCacheTokens
stats.Endpoints = endpoints
stats.UpstreamEndpoints = upstreamEndpoints
stats.EndpointPaths = endpointPaths
@@ -4121,7 +4160,8 @@ func (r *usageLogRepository) loadUsers(ctx context.Context, ids []int64) (map[in
if len(ids) == 0 {
return out, nil
}
- models, err := r.client.User.Query().Where(dbuser.IDIn(ids...)).All(ctx)
+ // 无条件穿透软删除:ids 来自调用方已按 user_id 筛选的日志行;普通用户路径强制 UserID=本人(本人必为活跃用户),不会借此解析他人已删身份;仅 admin 路径可借此显示已删用户。
+ models, err := r.client.User.Query().Where(dbuser.IDIn(ids...)).All(mixins.SkipSoftDelete(ctx))
if err != nil {
return nil, err
}
diff --git a/backend/internal/repository/usage_log_repo_deleted_user_integration_test.go b/backend/internal/repository/usage_log_repo_deleted_user_integration_test.go
new file mode 100644
index 00000000..70835b03
--- /dev/null
+++ b/backend/internal/repository/usage_log_repo_deleted_user_integration_test.go
@@ -0,0 +1,65 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUsageLog_ListWithFilters_ResolvesSoftDeletedUser(t *testing.T) {
+ ctx := context.Background()
+ tx := testEntTx(t)
+ client := tx.Client()
+ repo := newUsageLogRepositoryWithSQL(client, tx)
+
+ // 一个活跃用户、一个将被软删的用户,各一条日志。
+ active := mustCreateUser(t, client, &service.User{Email: "active-listfilter@test.com"})
+ deleted := mustCreateUser(t, client, &service.User{Email: "deleted-listfilter@test.com"})
+ apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: deleted.ID, Key: "sk-del-1", Name: "k"})
+ apiKey2 := mustCreateApiKey(t, client, &service.APIKey{UserID: active.ID, Key: "sk-act-1", Name: "k"})
+ account := mustCreateAccount(t, client, &service.Account{Name: "acc-listfilter"})
+
+ now := time.Now().UTC()
+ for _, u := range []struct {
+ uid int64
+ kid int64
+ }{{deleted.ID, apiKey.ID}, {active.ID, apiKey2.ID}} {
+ _, err := repo.Create(ctx, &service.UsageLog{
+ UserID: u.uid, APIKeyID: u.kid, AccountID: account.ID,
+ Model: "claude-3", InputTokens: 1, OutputTokens: 1,
+ TotalCost: 0.1, ActualCost: 0.1, CreatedAt: now,
+ })
+ require.NoError(t, err)
+ }
+
+ // 软删除该用户(触发 SoftDeleteMixin Hook → UPDATE deleted_at)。
+ require.NoError(t, client.User.DeleteOneID(deleted.ID).Exec(ctx))
+
+ logs, _, err := repo.ListWithFilters(ctx, pagination.PaginationParams{Page: 1, PageSize: 50},
+ usagestats.UsageLogFilters{ExactTotal: true})
+ require.NoError(t, err)
+
+ byUser := map[int64]service.UsageLog{}
+ for _, l := range logs {
+ byUser[l.UserID] = l
+ }
+
+ // 已删用户的日志行:富化后 User 非 nil、邮箱正确、DeletedAt 非 nil。
+ delLog, ok := byUser[deleted.ID]
+ require.True(t, ok, "deleted user's usage log must still be listed")
+ require.NotNil(t, delLog.User, "deleted user identity must resolve")
+ require.Equal(t, "deleted-listfilter@test.com", delLog.User.Email)
+ require.NotNil(t, delLog.User.DeletedAt, "DeletedAt must be set for soft-deleted user")
+
+ // 活跃用户:DeletedAt 为 nil。
+ actLog := byUser[active.ID]
+ require.NotNil(t, actLog.User)
+ require.Nil(t, actLog.User.DeletedAt)
+}
diff --git a/backend/internal/repository/usage_log_repo_stats_integration_test.go b/backend/internal/repository/usage_log_repo_stats_integration_test.go
new file mode 100644
index 00000000..09ac2aee
--- /dev/null
+++ b/backend/internal/repository/usage_log_repo_stats_integration_test.go
@@ -0,0 +1,51 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUsageLog_GetStatsWithFilters_AggregatesAndEndpoints(t *testing.T) {
+ ctx := context.Background()
+ tx := testEntTx(t)
+ client := tx.Client()
+ repo := newUsageLogRepositoryWithSQL(client, tx)
+
+ user := mustCreateUser(t, client, &service.User{Email: "stats@test.com"})
+ apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: user.ID, Key: "sk-stats-1", Name: "k"})
+ account := mustCreateAccount(t, client, &service.Account{Name: "acc-stats"})
+
+ now := time.Now().UTC()
+ inboundEndpoint := "/v1/messages"
+ upstreamEndpoint := "/v1/responses"
+ for i := 0; i < 3; i++ {
+ _, err := repo.Create(ctx, &service.UsageLog{
+ UserID: user.ID, APIKeyID: apiKey.ID, AccountID: account.ID,
+ Model: "claude-3", InputTokens: 2, OutputTokens: 3,
+ TotalCost: 0.5, ActualCost: 0.4, CreatedAt: now,
+ InboundEndpoint: &inboundEndpoint, UpstreamEndpoint: &upstreamEndpoint,
+ })
+ require.NoError(t, err)
+ }
+
+ start := now.Add(-1 * time.Hour)
+ end := now.Add(1 * time.Hour)
+ // 按本测试创建的 user 维度过滤:集成库为共享实例,其它用 testEntClient 的兄弟测试会留下
+ // 已提交的 usage_log 行(含零 token 的失败请求),不限定 user 会把它们计入 TotalRequests。
+ stats, err := repo.GetStatsWithFilters(ctx, usagestats.UsageLogFilters{UserID: user.ID, StartTime: &start, EndTime: &end})
+ require.NoError(t, err)
+ require.Equal(t, int64(3), stats.TotalRequests)
+ require.Equal(t, int64(6), stats.TotalInputTokens)
+ require.Equal(t, int64(9), stats.TotalOutputTokens)
+ require.InDelta(t, 1.2, stats.TotalActualCost, 1e-9)
+ require.NotEmpty(t, stats.Endpoints)
+ require.NotEmpty(t, stats.UpstreamEndpoints)
+ require.NotEmpty(t, stats.EndpointPaths)
+}
diff --git a/backend/internal/repository/user_platform_quota_adapter_test.go b/backend/internal/repository/user_platform_quota_adapter_test.go
index a55d2e9c..f31defe5 100644
--- a/backend/internal/repository/user_platform_quota_adapter_test.go
+++ b/backend/internal/repository/user_platform_quota_adapter_test.go
@@ -38,6 +38,9 @@ func (f *fakeRepoForAdapter) UpsertForUser(_ context.Context, userID int64, reco
f.upsertCalledWith = records
return f.upsertErr
}
+func (f *fakeRepoForAdapter) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error {
+ return nil
+}
func TestGenericAdapter_UpsertForUser_ForwardsRecords(t *testing.T) {
fake := &fakeRepoForAdapter{}
diff --git a/backend/internal/repository/user_platform_quota_repo.go b/backend/internal/repository/user_platform_quota_repo.go
index 1e2e7f51..ccba2330 100644
--- a/backend/internal/repository/user_platform_quota_repo.go
+++ b/backend/internal/repository/user_platform_quota_repo.go
@@ -2,6 +2,7 @@ package repository
import (
"context"
+ "errors"
"fmt"
"strings"
"time"
@@ -9,6 +10,7 @@ import (
dbent "github.com/Wei-Shaw/sub2api/ent"
"github.com/Wei-Shaw/sub2api/ent/userplatformquota"
"github.com/Wei-Shaw/sub2api/internal/pkg/timezone"
+ "github.com/lib/pq"
)
// UserPlatformQuotaRecord 是 repository 层的传输结构体,
@@ -30,6 +32,22 @@ type UserPlatformQuotaRecord struct {
// ErrUserPlatformQuotaNotFound 用于 ResetExpiredWindow 等需要"必须命中已有记录"的方法。
var ErrUserPlatformQuotaNotFound = fmt.Errorf("user platform quota record not found")
+// ErrUserPlatformQuotaFKViolation 当批量 UPSERT 中存在 user_id 不在 users 表的记录时返回。
+var ErrUserPlatformQuotaFKViolation = errors.New("user platform quota snapshot FK violation")
+
+// UserPlatformQuotaSnapshot 是 BatchSnapshotUsage 的输入结构体,
+// 表示 Redis 当前窗口快照(用于绝对值覆盖写入 DB)。
+type UserPlatformQuotaSnapshot struct {
+ UserID int64
+ Platform string
+ DailyUsageUSD float64
+ WeeklyUsageUSD float64
+ MonthlyUsageUSD float64
+ DailyWindowStart time.Time
+ WeeklyWindowStart time.Time
+ MonthlyWindowStart time.Time
+}
+
// UserPlatformQuotaRepository 定义用户平台配额的数据访问接口。
type UserPlatformQuotaRepository interface {
// BulkInsertInitial 幂等批量插入初始配额记录(ON CONFLICT DO NOTHING)。
@@ -44,6 +62,10 @@ type UserPlatformQuotaRepository interface {
ResetExpiredWindow(ctx context.Context, userID int64, platform string, window string, newStart time.Time) error
// UpsertForUser 全量替换该用户所有平台限额配置(详见 service.UserPlatformQuotaRepository.UpsertForUser)。
UpsertForUser(ctx context.Context, userID int64, records []UserPlatformQuotaRecord) error
+ // BatchSnapshotUsage 用一条多行 UPSERT 把整批 usage 以绝对值覆盖写入(非累加)。
+ // usage/window_start 直接取 EXCLUDED(Redis 当前窗口快照),无 CASE。整批共用 now 作 created/updated_at。
+ // 要求 snapshots 内 (user,platform) 不重复。FK 违反返回 ErrUserPlatformQuotaFKViolation。
+ BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error
}
type userPlatformQuotaRepository struct {
@@ -414,3 +436,73 @@ func insertLimitsRow(ctx context.Context, client *dbent.Client, userID int64, re
}
return nil
}
+
+// batchRows 是 BatchSnapshotUsage 每批最大行数(9 参/行 × 6000 ≈ 54000 参,低于 Postgres 65535 上限)。
+const batchRows = 6000
+
+// BatchSnapshotUsage 用一条多行 UPSERT 把整批 usage 以绝对值覆盖写入(非累加)。
+// 每批最多 batchRows 行;$1=now 共用;每行 8 个 per-row 参(user_id, platform, 3×usage, 3×window_start)。
+// FK 违反(user_id 不存在)返回 ErrUserPlatformQuotaFKViolation。
+//
+// 注意:snapshots 超过 batchRows 会分多条 SQL 执行且【非单事务】——若某子批 FK 失败,
+// 先前子批已写入无法回滚。调用方(flusher)应保证单次 batchSize ≤ batchRows
+// (默认 flush_batch_size=1000 < 6000,安全)。
+// 另注:启用 flusher 后,本绝对值覆盖与 admin 直写 DB(ResetExpiredWindow/UpsertForUser)存在覆盖竞态,
+// 详见 service/user_platform_quota_flusher.go 中 flushOneBatch 的"已知竞态"注释。
+func (r *userPlatformQuotaRepository) BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error {
+ if len(snapshots) == 0 {
+ return nil
+ }
+
+ client := clientFromContext(ctx, r.client)
+
+ for start := 0; start < len(snapshots); start += batchRows {
+ end := start + batchRows
+ if end > len(snapshots) {
+ end = len(snapshots)
+ }
+ batch := snapshots[start:end]
+
+ var sb strings.Builder
+ _, _ = sb.WriteString(
+ "INSERT INTO user_platform_quotas" +
+ " (user_id, platform, daily_usage_usd, weekly_usage_usd, monthly_usage_usd," +
+ " daily_window_start, weekly_window_start, monthly_window_start, created_at, updated_at)" +
+ " VALUES ")
+
+ // $1 = now(共用);每行 8 个 per-row 参,从 $2 起连续编号。
+ args := []any{now}
+ for i, s := range batch {
+ if i > 0 {
+ _, _ = sb.WriteString(",")
+ }
+ b := len(args) // 当前 per-row 第一个参数的 0-based 索引,实际占位符 = b+1
+ fmt.Fprintf(&sb, "($%d,$%d,$%d,$%d,$%d,$%d,$%d,$%d,$1,$1)",
+ b+1, b+2, b+3, b+4, b+5, b+6, b+7, b+8)
+ args = append(args,
+ s.UserID, s.Platform,
+ s.DailyUsageUSD, s.WeeklyUsageUSD, s.MonthlyUsageUSD,
+ s.DailyWindowStart, s.WeeklyWindowStart, s.MonthlyWindowStart,
+ )
+ }
+
+ _, _ = sb.WriteString(
+ " ON CONFLICT (user_id, platform) WHERE deleted_at IS NULL DO UPDATE SET" +
+ " daily_usage_usd = EXCLUDED.daily_usage_usd," +
+ " weekly_usage_usd = EXCLUDED.weekly_usage_usd," +
+ " monthly_usage_usd = EXCLUDED.monthly_usage_usd," +
+ " daily_window_start = EXCLUDED.daily_window_start," +
+ " weekly_window_start = EXCLUDED.weekly_window_start," +
+ " monthly_window_start = EXCLUDED.monthly_window_start," +
+ " updated_at = EXCLUDED.updated_at")
+
+ if _, err := client.ExecContext(ctx, sb.String(), args...); err != nil {
+ var pqErr *pq.Error
+ if errors.As(err, &pqErr) && pqErr.Code == "23503" {
+ return ErrUserPlatformQuotaFKViolation
+ }
+ return err
+ }
+ }
+ return nil
+}
diff --git a/backend/internal/repository/user_platform_quota_repo_integration_test.go b/backend/internal/repository/user_platform_quota_repo_integration_test.go
index f02eeaa9..39e2f6e0 100644
--- a/backend/internal/repository/user_platform_quota_repo_integration_test.go
+++ b/backend/internal/repository/user_platform_quota_repo_integration_test.go
@@ -267,3 +267,101 @@ func TestUserPlatformQuotaRepository_ResetExpiredWindow_NotFoundReturnsSentinel(
require.True(t, errors.Is(err, ErrUserPlatformQuotaNotFound),
"expected ErrUserPlatformQuotaNotFound, got %v", err)
}
+
+// TestBatchSnapshotUsage_InsertOverwriteMultiKey 验证 BatchSnapshotUsage 的绝对值覆盖语义:
+// 1. 首批插入 2 条(不同 user),验证 daily 等于首批值;
+// 2. 对同一 key 传不同值,验证 daily 等于新值(绝对覆盖,非累加)。
+func TestBatchSnapshotUsage_InsertOverwriteMultiKey(t *testing.T) {
+ ctx := context.Background()
+ // BatchSnapshotUsage 不开事务(直接写),使用独立 client 保证跨调用可见性。
+ client := testEntClient(t)
+
+ userID1 := mustCreateUserForQuota(t, client)
+ userID2 := mustCreateUserForQuota(t, client)
+
+ repo := NewUserPlatformQuotaRepository(client)
+
+ now := time.Date(2026, 5, 29, 12, 0, 0, 0, time.UTC)
+ dailyStart := time.Date(2026, 5, 29, 0, 0, 0, 0, time.UTC)
+ weeklyStart := time.Date(2026, 5, 25, 0, 0, 0, 0, time.UTC) // 当周一
+ monthlyStart := time.Date(2026, 5, 1, 0, 0, 0, 0, time.UTC)
+
+ // ── 第一批:插入 2 行 ──────────────────────────────────────────────────────
+ firstBatch := []UserPlatformQuotaSnapshot{
+ {
+ UserID: userID1,
+ Platform: "anthropic",
+ DailyUsageUSD: 1.0,
+ WeeklyUsageUSD: 3.0,
+ MonthlyUsageUSD: 5.0,
+ DailyWindowStart: dailyStart,
+ WeeklyWindowStart: weeklyStart,
+ MonthlyWindowStart: monthlyStart,
+ },
+ {
+ UserID: userID2,
+ Platform: "openai",
+ DailyUsageUSD: 2.0,
+ WeeklyUsageUSD: 4.0,
+ MonthlyUsageUSD: 6.0,
+ DailyWindowStart: dailyStart,
+ WeeklyWindowStart: weeklyStart,
+ MonthlyWindowStart: monthlyStart,
+ },
+ }
+ require.NoError(t, repo.BatchSnapshotUsage(ctx, firstBatch, now), "first batch upsert")
+
+ // 验证首批 daily 值
+ rec1, err := repo.GetByUserPlatform(ctx, userID1, "anthropic")
+ require.NoError(t, err)
+ require.NotNil(t, rec1, "user1/anthropic should exist after first batch")
+ require.InDelta(t, 1.0, rec1.DailyUsageUSD, 1e-9, "user1 daily after first batch")
+ require.InDelta(t, 3.0, rec1.WeeklyUsageUSD, 1e-9, "user1 weekly after first batch")
+ require.InDelta(t, 5.0, rec1.MonthlyUsageUSD, 1e-9, "user1 monthly after first batch")
+
+ rec2, err := repo.GetByUserPlatform(ctx, userID2, "openai")
+ require.NoError(t, err)
+ require.NotNil(t, rec2, "user2/openai should exist after first batch")
+ require.InDelta(t, 2.0, rec2.DailyUsageUSD, 1e-9, "user2 daily after first batch")
+
+ // ── 第二批:对同一 key 传不同值,验证绝对覆盖(非累加)──────────────────
+ now2 := now.Add(5 * time.Minute)
+ secondBatch := []UserPlatformQuotaSnapshot{
+ {
+ UserID: userID1,
+ Platform: "anthropic",
+ DailyUsageUSD: 9.9, // 新值,不是 1.0+9.9=10.9
+ WeeklyUsageUSD: 19.9, // 新值,不是 3.0+19.9=22.9
+ MonthlyUsageUSD: 29.9, // 新值
+ DailyWindowStart: dailyStart,
+ WeeklyWindowStart: weeklyStart,
+ MonthlyWindowStart: monthlyStart,
+ },
+ {
+ UserID: userID2,
+ Platform: "openai",
+ DailyUsageUSD: 8.8,
+ WeeklyUsageUSD: 18.8,
+ MonthlyUsageUSD: 28.8,
+ DailyWindowStart: dailyStart,
+ WeeklyWindowStart: weeklyStart,
+ MonthlyWindowStart: monthlyStart,
+ },
+ }
+ require.NoError(t, repo.BatchSnapshotUsage(ctx, secondBatch, now2), "second batch upsert")
+
+ // 验证第二批覆盖:daily 应为新值,不是累加
+ rec1After, err := repo.GetByUserPlatform(ctx, userID1, "anthropic")
+ require.NoError(t, err)
+ require.NotNil(t, rec1After)
+ require.InDelta(t, 9.9, rec1After.DailyUsageUSD, 1e-9, "user1 daily must be overwritten to 9.9 (not accumulated)")
+ require.InDelta(t, 19.9, rec1After.WeeklyUsageUSD, 1e-9, "user1 weekly must be overwritten to 19.9")
+ require.InDelta(t, 29.9, rec1After.MonthlyUsageUSD, 1e-9, "user1 monthly must be overwritten to 29.9")
+
+ rec2After, err := repo.GetByUserPlatform(ctx, userID2, "openai")
+ require.NoError(t, err)
+ require.NotNil(t, rec2After)
+ require.InDelta(t, 8.8, rec2After.DailyUsageUSD, 1e-9, "user2 daily must be overwritten to 8.8 (not accumulated)")
+ require.InDelta(t, 18.8, rec2After.WeeklyUsageUSD, 1e-9, "user2 weekly must be overwritten to 18.8")
+ require.InDelta(t, 28.8, rec2After.MonthlyUsageUSD, 1e-9, "user2 monthly must be overwritten to 28.8")
+}
diff --git a/backend/internal/repository/user_platform_quota_service_adapter.go b/backend/internal/repository/user_platform_quota_service_adapter.go
index 7495cd26..5240bb54 100644
--- a/backend/internal/repository/user_platform_quota_service_adapter.go
+++ b/backend/internal/repository/user_platform_quota_service_adapter.go
@@ -94,6 +94,29 @@ func (a *userPlatformQuotaServiceAdapter) ResetExpiredWindow(ctx context.Context
return err
}
+// BatchSnapshotUsage 转换 []service.UserPlatformQuotaSnapshot → []UserPlatformQuotaSnapshot,
+// 调底层 repo,并将 repository FK sentinel 包装为 service sentinel。
+func (a *userPlatformQuotaServiceAdapter) BatchSnapshotUsage(ctx context.Context, snapshots []service.UserPlatformQuotaSnapshot, now time.Time) error {
+ repoSnaps := make([]UserPlatformQuotaSnapshot, len(snapshots))
+ for i, s := range snapshots {
+ repoSnaps[i] = UserPlatformQuotaSnapshot{
+ UserID: s.UserID,
+ Platform: s.Platform,
+ DailyUsageUSD: s.DailyUsageUSD,
+ WeeklyUsageUSD: s.WeeklyUsageUSD,
+ MonthlyUsageUSD: s.MonthlyUsageUSD,
+ DailyWindowStart: s.DailyWindowStart,
+ WeeklyWindowStart: s.WeeklyWindowStart,
+ MonthlyWindowStart: s.MonthlyWindowStart,
+ }
+ }
+ err := a.inner.BatchSnapshotUsage(ctx, repoSnaps, now)
+ if errors.Is(err, ErrUserPlatformQuotaFKViolation) {
+ return fmt.Errorf("%w: %v", service.ErrUserPlatformQuotaFKViolation, err)
+ }
+ return err
+}
+
// genericUserPlatformQuotaAdapter 通过通用接口适配(用于测试 fake 或非标准实现)。
type genericUserPlatformQuotaAdapter struct {
inner UserPlatformQuotaRepository
@@ -167,6 +190,29 @@ func (a *genericUserPlatformQuotaAdapter) ResetExpiredWindow(ctx context.Context
return err
}
+// BatchSnapshotUsage 转换 []service.UserPlatformQuotaSnapshot → []UserPlatformQuotaSnapshot(通用 adapter),
+// 并将 repository FK sentinel 包装为 service sentinel。
+func (a *genericUserPlatformQuotaAdapter) BatchSnapshotUsage(ctx context.Context, snapshots []service.UserPlatformQuotaSnapshot, now time.Time) error {
+ repoSnaps := make([]UserPlatformQuotaSnapshot, len(snapshots))
+ for i, s := range snapshots {
+ repoSnaps[i] = UserPlatformQuotaSnapshot{
+ UserID: s.UserID,
+ Platform: s.Platform,
+ DailyUsageUSD: s.DailyUsageUSD,
+ WeeklyUsageUSD: s.WeeklyUsageUSD,
+ MonthlyUsageUSD: s.MonthlyUsageUSD,
+ DailyWindowStart: s.DailyWindowStart,
+ WeeklyWindowStart: s.WeeklyWindowStart,
+ MonthlyWindowStart: s.MonthlyWindowStart,
+ }
+ }
+ err := a.inner.BatchSnapshotUsage(ctx, repoSnaps, now)
+ if errors.Is(err, ErrUserPlatformQuotaFKViolation) {
+ return fmt.Errorf("%w: %v", service.ErrUserPlatformQuotaFKViolation, err)
+ }
+ return err
+}
+
// toServiceRecord 将 repository.UserPlatformQuotaRecord 转换为 service.UserPlatformQuotaRecord。
func toServiceRecord(rec *UserPlatformQuotaRecord) *service.UserPlatformQuotaRecord {
return &service.UserPlatformQuotaRecord{
diff --git a/backend/internal/repository/user_repo.go b/backend/internal/repository/user_repo.go
index 610d9a7b..fb05452d 100644
--- a/backend/internal/repository/user_repo.go
+++ b/backend/internal/repository/user_repo.go
@@ -16,6 +16,7 @@ import (
dbgroup "github.com/Wei-Shaw/sub2api/ent/group"
"github.com/Wei-Shaw/sub2api/ent/identityadoptiondecision"
"github.com/Wei-Shaw/sub2api/ent/predicate"
+ "github.com/Wei-Shaw/sub2api/ent/schema/mixins"
dbuser "github.com/Wei-Shaw/sub2api/ent/user"
"github.com/Wei-Shaw/sub2api/ent/userallowedgroup"
"github.com/Wei-Shaw/sub2api/ent/usersubscription"
@@ -133,6 +134,23 @@ func (r *userRepository) GetByID(ctx context.Context, id int64) (*service.User,
return out, nil
}
+func (r *userRepository) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ ctx = mixins.SkipSoftDelete(ctx)
+ m, err := r.client.User.Query().Where(dbuser.IDEQ(id)).Only(ctx)
+ if err != nil {
+ return nil, translatePersistenceError(err, service.ErrUserNotFound, nil)
+ }
+ out := userEntityToService(m)
+ groups, err := r.loadAllowedGroups(ctx, []int64{id})
+ if err != nil {
+ return nil, err
+ }
+ if v, ok := groups[id]; ok {
+ out.AllowedGroups = v
+ }
+ return out, nil
+}
+
func (r *userRepository) GetByEmail(ctx context.Context, email string) (*service.User, error) {
matches, err := r.client.User.Query().
Where(userEmailLookupPredicate(email)).
@@ -405,6 +423,12 @@ func (r *userRepository) List(ctx context.Context, params pagination.PaginationP
}
func (r *userRepository) ListWithFilters(ctx context.Context, params pagination.PaginationParams, filters service.UserListFilters) ([]service.User, *pagination.PaginationResult, error) {
+ // SkipSoftDelete 仅作用于 User 身份解析(下方 Count/All);订阅、分组等关联实体沿用原始 ctx,避免穿透到这些同样带软删除的实体而带出已删除行。
+ userCtx := ctx
+ if filters.IncludeDeleted {
+ userCtx = mixins.SkipSoftDelete(ctx)
+ }
+
q := r.client.User.Query()
if filters.Status != "" {
@@ -445,7 +469,7 @@ func (r *userRepository) ListWithFilters(ctx context.Context, params pagination.
q = q.Where(dbuser.IDIn(allowedUserIDs...))
}
- total, err := q.Clone().Count(ctx)
+ total, err := q.Clone().Count(userCtx)
if err != nil {
return nil, nil, err
}
@@ -457,7 +481,7 @@ func (r *userRepository) ListWithFilters(ctx context.Context, params pagination.
usersQuery = usersQuery.Order(order)
}
- users, err := usersQuery.All(ctx)
+ users, err := usersQuery.All(userCtx)
if err != nil {
return nil, nil, err
}
diff --git a/backend/internal/repository/user_repo_include_deleted_integration_test.go b/backend/internal/repository/user_repo_include_deleted_integration_test.go
new file mode 100644
index 00000000..014b24f9
--- /dev/null
+++ b/backend/internal/repository/user_repo_include_deleted_integration_test.go
@@ -0,0 +1,69 @@
+//go:build integration
+
+package repository
+
+import (
+ "context"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
+ "github.com/Wei-Shaw/sub2api/internal/service"
+ "github.com/stretchr/testify/require"
+)
+
+func TestUserRepo_ListWithFilters_IncludeDeleted(t *testing.T) {
+ ctx := context.Background()
+ tx := testEntTx(t)
+ client := tx.Client()
+ repo := NewUserRepository(client, integrationDB)
+
+ active := mustCreateUser(t, client, &service.User{Email: "shared-keyword-active@test.com"})
+ deleted := mustCreateUser(t, client, &service.User{Email: "shared-keyword-deleted@test.com"})
+ require.NoError(t, client.User.DeleteOneID(deleted.ID).Exec(ctx))
+
+ params := pagination.PaginationParams{Page: 1, PageSize: 50, SortBy: "email", SortOrder: "asc"}
+
+ // 默认(不含已删):只返回活跃用户。
+ usersDefault, resDefault, err := repo.ListWithFilters(ctx, params,
+ service.UserListFilters{Search: "shared-keyword-"})
+ require.NoError(t, err)
+ require.Len(t, usersDefault, 1)
+ require.Equal(t, active.ID, usersDefault[0].ID)
+ require.EqualValues(t, 1, resDefault.Total)
+
+ // IncludeDeleted=true:两个都返回,且 Total 与结果集一致。
+ usersAll, resAll, err := repo.ListWithFilters(ctx, params,
+ service.UserListFilters{Search: "shared-keyword-", IncludeDeleted: true})
+ require.NoError(t, err)
+ require.Len(t, usersAll, 2)
+ require.EqualValues(t, 2, resAll.Total, "Count 必须与结果集行数一致")
+
+ var delUser *service.User
+ for i := range usersAll {
+ if usersAll[i].ID == deleted.ID {
+ delUser = &usersAll[i]
+ }
+ }
+ require.NotNil(t, delUser)
+ require.NotNil(t, delUser.DeletedAt)
+}
+
+func TestUserRepo_GetByIDIncludeDeleted(t *testing.T) {
+ ctx := context.Background()
+ tx := testEntTx(t)
+ client := tx.Client()
+ repo := NewUserRepository(client, integrationDB)
+
+ u := mustCreateUser(t, client, &service.User{Email: "getbyid-deleted@test.com"})
+ require.NoError(t, client.User.DeleteOneID(u.ID).Exec(ctx))
+
+ // 默认 GetByID:找不到(被软删过滤)。
+ _, err := repo.GetByID(ctx, u.ID)
+ require.ErrorIs(t, err, service.ErrUserNotFound)
+
+ // GetByIDIncludeDeleted:找得到,且 DeletedAt 非空。
+ got, err := repo.GetByIDIncludeDeleted(ctx, u.ID)
+ require.NoError(t, err)
+ require.Equal(t, "getbyid-deleted@test.com", got.Email)
+ require.NotNil(t, got.DeletedAt)
+}
diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go
index 8bc9e280..766225ff 100644
--- a/backend/internal/server/api_contract_test.go
+++ b/backend/internal/server/api_contract_test.go
@@ -843,6 +843,7 @@ func TestAPIContracts(t *testing.T) {
"payment_visible_method_wxpay_enabled": false,
"openai_advanced_scheduler_enabled": true,
"openai_codex_user_agent": "",
+ "openai_allow_claude_code_codex_plugin": false,
"openai_fast_policy_settings": {
"rules": []
},
@@ -895,7 +896,8 @@ func TestAPIContracts(t *testing.T) {
"wechat_connect_mobile_app_secret_configured": false,
"wechat_connect_redirect_url": "",
"wechat_connect_frontend_redirect_url": "/auth/wechat/callback",
- "wechat_connect_scopes": "snsapi_login"
+ "wechat_connect_scopes": "snsapi_login",
+ "allow_user_view_error_requests": false
}
}`,
},
@@ -1079,6 +1081,7 @@ func TestAPIContracts(t *testing.T) {
"payment_visible_method_wxpay_enabled": false,
"openai_advanced_scheduler_enabled": false,
"openai_codex_user_agent": "",
+ "openai_allow_claude_code_codex_plugin": false,
"openai_fast_policy_settings": {
"rules": []
},
@@ -1165,7 +1168,8 @@ func TestAPIContracts(t *testing.T) {
"auth_source_default_dingtalk_subscriptions": [],
"auth_source_default_dingtalk_grant_on_signup": false,
"auth_source_default_dingtalk_grant_on_first_bind": false,
- "force_email_on_third_party_signup": false
+ "force_email_on_third_party_signup": false,
+ "allow_user_view_error_requests": false
}
}`,
},
@@ -1277,7 +1281,7 @@ func newContractDeps(t *testing.T) *contractDeps {
adminService := service.NewAdminService(userRepo, groupRepo, &accountRepo, proxyRepo, apiKeyRepo, redeemRepo, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
authHandler := handler.NewAuthHandler(cfg, nil, userService, settingService, nil, redeemService, nil, nil)
apiKeyHandler := handler.NewAPIKeyHandler(apiKeyService)
- usageHandler := handler.NewUsageHandler(usageService, apiKeyService)
+ usageHandler := handler.NewUsageHandler(usageService, apiKeyService, nil, nil)
adminSettingHandler := adminhandler.NewSettingHandler(settingService, nil, nil, nil, nil, nil, nil)
adminAccountHandler := adminhandler.NewAccountHandler(adminService, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil)
@@ -1490,6 +1494,10 @@ func (r *stubUserRepo) DisableTotp(ctx context.Context, userID int64) error {
return errors.New("not implemented")
}
+func (r *stubUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ panic("unexpected GetByIDIncludeDeleted call")
+}
+
type stubApiKeyCache struct{}
func (stubApiKeyCache) GetCreateAttemptCount(ctx context.Context, userID int64) (int, error) {
@@ -1731,7 +1739,7 @@ func (s *stubAccountRepo) SetRateLimited(ctx context.Context, id int64, resetAt
return errors.New("not implemented")
}
-func (s *stubAccountRepo) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time) error {
+func (s *stubAccountRepo) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error {
return errors.New("not implemented")
}
@@ -2096,6 +2104,10 @@ func (r *stubApiKeyRepo) Delete(ctx context.Context, id int64) error {
return nil
}
+func (r *stubApiKeyRepo) DeleteWithAudit(ctx context.Context, id int64) error {
+ return r.Delete(ctx, id)
+}
+
func (r *stubApiKeyRepo) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, _ service.APIKeyListFilters) ([]service.APIKey, *pagination.PaginationResult, error) {
ids := make([]int64, 0, len(r.byID))
for id := range r.byID {
diff --git a/backend/internal/server/middleware/admin_auth_test.go b/backend/internal/server/middleware/admin_auth_test.go
index 303d0db8..3110c6c1 100644
--- a/backend/internal/server/middleware/admin_auth_test.go
+++ b/backend/internal/server/middleware/admin_auth_test.go
@@ -236,3 +236,7 @@ func (s *stubUserRepo) EnableTotp(ctx context.Context, userID int64) error {
func (s *stubUserRepo) DisableTotp(ctx context.Context, userID int64) error {
panic("unexpected DisableTotp call")
}
+
+func (s *stubUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ panic("unexpected GetByIDIncludeDeleted call")
+}
diff --git a/backend/internal/server/middleware/api_key_auth.go b/backend/internal/server/middleware/api_key_auth.go
index d33ccbf5..ba43d126 100644
--- a/backend/internal/server/middleware/api_key_auth.go
+++ b/backend/internal/server/middleware/api_key_auth.go
@@ -76,6 +76,10 @@ func apiKeyAuthWithSubscription(apiKeyService *service.APIKeyService, subscripti
return
}
+ // apiKey 已加载(含 User/Group)。即便后续因分组停用/Key 停用/用户停用/
+ // IP 限制等早退中断,也让 Ops 错误日志能回退取到 user/group/platform。
+ SetOpsFallbackAPIKey(c, apiKey)
+
// ── 3. 基础鉴权(始终执行) ─────────────────────────────────
// disabled / 未知状态 → 无条件拦截(expired 和 quota_exhausted 留给计费阶段)
@@ -237,6 +241,26 @@ func GetAPIKeyFromContext(c *gin.Context) (*service.APIKey, bool) {
return apiKey, ok
}
+// SetOpsFallbackAPIKey 记录已加载的 API Key,供 Ops 错误日志在鉴权早退时回退使用。
+// 与 ContextKeyAPIKey 区分:写入它不代表请求已通过鉴权,因此不影响 handler、
+// 审计日志等对“已鉴权”的判断。
+func SetOpsFallbackAPIKey(c *gin.Context, apiKey *service.APIKey) {
+ if c == nil || apiKey == nil {
+ return
+ }
+ c.Set(string(ContextKeyOpsFallbackAPIKey), apiKey)
+}
+
+// GetOpsFallbackAPIKey 读取 Ops 错误日志专用的回退 API Key。
+func GetOpsFallbackAPIKey(c *gin.Context) (*service.APIKey, bool) {
+ value, exists := c.Get(string(ContextKeyOpsFallbackAPIKey))
+ if !exists {
+ return nil, false
+ }
+ apiKey, ok := value.(*service.APIKey)
+ return apiKey, ok
+}
+
// GetSubscriptionFromContext 从上下文中获取订阅信息
func GetSubscriptionFromContext(c *gin.Context) (*service.UserSubscription, bool) {
value, exists := c.Get(string(ContextKeySubscription))
diff --git a/backend/internal/server/middleware/api_key_auth_google.go b/backend/internal/server/middleware/api_key_auth_google.go
index 596bed52..97f3936c 100644
--- a/backend/internal/server/middleware/api_key_auth_google.go
+++ b/backend/internal/server/middleware/api_key_auth_google.go
@@ -42,6 +42,10 @@ func APIKeyAuthWithSubscriptionGoogle(apiKeyService *service.APIKeyService, subs
return
}
+ // 同 api_key_auth.go:早退中断前也写入 Ops 回退 key,便于错误日志展示
+ // user/group/platform。
+ SetOpsFallbackAPIKey(c, apiKey)
+
if !apiKey.IsActive() {
abortWithGoogleError(c, 401, "API key is disabled")
return
diff --git a/backend/internal/server/middleware/api_key_auth_google_test.go b/backend/internal/server/middleware/api_key_auth_google_test.go
index feadd27d..32e7e70f 100644
--- a/backend/internal/server/middleware/api_key_auth_google_test.go
+++ b/backend/internal/server/middleware/api_key_auth_google_test.go
@@ -56,6 +56,9 @@ func (f fakeAPIKeyRepo) Update(ctx context.Context, key *service.APIKey) error {
func (f fakeAPIKeyRepo) Delete(ctx context.Context, id int64) error {
return errors.New("not implemented")
}
+func (f fakeAPIKeyRepo) DeleteWithAudit(ctx context.Context, id int64) error {
+ return errors.New("not implemented")
+}
func (f fakeAPIKeyRepo) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, _ service.APIKeyListFilters) ([]service.APIKey, *pagination.PaginationResult, error) {
return nil, nil, errors.New("not implemented")
}
diff --git a/backend/internal/server/middleware/api_key_auth_test.go b/backend/internal/server/middleware/api_key_auth_test.go
index 76a24192..5d48bed2 100644
--- a/backend/internal/server/middleware/api_key_auth_test.go
+++ b/backend/internal/server/middleware/api_key_auth_test.go
@@ -419,6 +419,138 @@ func TestAPIKeyAuthRejectsUnavailableGroup(t *testing.T) {
}
}
+func TestAPIKeyAuthSetsOpsFallbackKeyOnEarlyAbort(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(101)
+ user := &service.User{
+ ID: 7,
+ Role: service.RoleUser,
+ Status: service.StatusActive,
+ Balance: 10,
+ Concurrency: 3,
+ }
+ apiKey := &service.APIKey{
+ ID: 100,
+ UserID: user.ID,
+ GroupID: &groupID,
+ Key: "test-key",
+ Status: service.StatusActive,
+ User: user,
+ Group: &service.Group{
+ ID: groupID,
+ Name: "disabled",
+ Status: service.StatusDisabled,
+ Platform: service.PlatformAnthropic,
+ Hydrated: true,
+ },
+ }
+ apiKeyRepo := &stubApiKeyRepo{
+ getByKey: func(ctx context.Context, key string) (*service.APIKey, error) {
+ if key != apiKey.Key {
+ return nil, service.ErrAPIKeyNotFound
+ }
+ clone := *apiKey
+ return &clone, nil
+ },
+ }
+ cfg := &config.Config{RunMode: config.RunModeStandard}
+ apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
+
+ router := gin.New()
+ var fallback *service.APIKey
+ var fallbackOK bool
+ router.Use(func(c *gin.Context) {
+ c.Next()
+ fallback, fallbackOK = GetOpsFallbackAPIKey(c)
+ })
+ router.Use(gin.HandlerFunc(NewAPIKeyAuthMiddleware(apiKeyService, nil, cfg)))
+ router.GET("/t", func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"ok": true})
+ })
+
+ w := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, "/t", nil)
+ req.Header.Set("x-api-key", apiKey.Key)
+ router.ServeHTTP(w, req)
+
+ // 分组停用 → 早退中断,但 ops fallback key 仍应写入,含 user/group/platform。
+ require.Equal(t, http.StatusForbidden, w.Code)
+ require.Contains(t, w.Body.String(), "GROUP_DISABLED")
+ require.True(t, fallbackOK, "鉴权早退时也应写入 ops fallback api key")
+ require.NotNil(t, fallback)
+ require.Equal(t, apiKey.ID, fallback.ID)
+ require.NotNil(t, fallback.User)
+ require.Equal(t, user.ID, fallback.User.ID)
+ require.NotNil(t, fallback.GroupID)
+ require.Equal(t, groupID, *fallback.GroupID)
+ require.NotNil(t, fallback.Group)
+ require.Equal(t, service.PlatformAnthropic, fallback.Group.Platform)
+}
+
+func TestAPIKeyAuthGoogleSetsOpsFallbackKeyOnEarlyAbort(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ groupID := int64(202)
+ user := &service.User{
+ ID: 9,
+ Role: service.RoleUser,
+ Status: service.StatusActive,
+ Balance: 10,
+ Concurrency: 3,
+ }
+ apiKey := &service.APIKey{
+ ID: 200,
+ UserID: user.ID,
+ GroupID: &groupID,
+ Key: "g-key",
+ Status: service.StatusActive,
+ User: user,
+ Group: &service.Group{
+ ID: groupID,
+ Name: "disabled",
+ Status: service.StatusDisabled,
+ Platform: service.PlatformGemini,
+ Hydrated: true,
+ },
+ }
+ apiKeyRepo := &stubApiKeyRepo{
+ getByKey: func(ctx context.Context, key string) (*service.APIKey, error) {
+ if key != apiKey.Key {
+ return nil, service.ErrAPIKeyNotFound
+ }
+ clone := *apiKey
+ return &clone, nil
+ },
+ }
+ cfg := &config.Config{RunMode: config.RunModeStandard}
+ apiKeyService := service.NewAPIKeyService(apiKeyRepo, nil, nil, nil, nil, nil, cfg)
+
+ router := gin.New()
+ var fallback *service.APIKey
+ var fallbackOK bool
+ router.Use(func(c *gin.Context) {
+ c.Next()
+ fallback, fallbackOK = GetOpsFallbackAPIKey(c)
+ })
+ router.Use(gin.HandlerFunc(APIKeyAuthWithSubscriptionGoogle(apiKeyService, nil, cfg)))
+ router.GET("/t", func(c *gin.Context) {
+ c.JSON(http.StatusOK, gin.H{"ok": true})
+ })
+
+ w := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, "/t", nil)
+ req.Header.Set("x-goog-api-key", apiKey.Key)
+ router.ServeHTTP(w, req)
+
+ require.Equal(t, http.StatusForbidden, w.Code)
+ require.True(t, fallbackOK, "Google 鉴权早退时也应写入 ops fallback api key")
+ require.NotNil(t, fallback)
+ require.Equal(t, apiKey.ID, fallback.ID)
+ require.NotNil(t, fallback.User)
+ require.Equal(t, user.ID, fallback.User.ID)
+}
+
func TestRequireGroupAssignmentMarksUngroupedKeyBusinessLimited(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -761,6 +893,10 @@ func (r *stubApiKeyRepo) Delete(ctx context.Context, id int64) error {
return errors.New("not implemented")
}
+func (r *stubApiKeyRepo) DeleteWithAudit(ctx context.Context, id int64) error {
+ return errors.New("not implemented")
+}
+
func (r *stubApiKeyRepo) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, _ service.APIKeyListFilters) ([]service.APIKey, *pagination.PaginationResult, error) {
return nil, nil, errors.New("not implemented")
}
diff --git a/backend/internal/server/middleware/client_request_id.go b/backend/internal/server/middleware/client_request_id.go
index 6838d6af..5f886646 100644
--- a/backend/internal/server/middleware/client_request_id.go
+++ b/backend/internal/server/middleware/client_request_id.go
@@ -11,6 +11,8 @@ import (
"go.uber.org/zap"
)
+const clientRequestIDHeader = "X-Client-Request-ID"
+
// ClientRequestID ensures every request has a unique client_request_id in request.Context().
//
// This is used by the Ops monitoring module for end-to-end request correlation.
@@ -21,12 +23,14 @@ func ClientRequestID() gin.HandlerFunc {
return
}
- if v := c.Request.Context().Value(ctxkey.ClientRequestID); v != nil {
+ if v, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string); strings.TrimSpace(v) != "" {
+ c.Header(clientRequestIDHeader, strings.TrimSpace(v))
c.Next()
return
}
id := uuid.New().String()
+ c.Header(clientRequestIDHeader, id)
ctx := context.WithValue(c.Request.Context(), ctxkey.ClientRequestID, id)
requestLogger := logger.FromContext(ctx).With(zap.String("client_request_id", strings.TrimSpace(id)))
ctx = logger.IntoContext(ctx, requestLogger)
diff --git a/backend/internal/server/middleware/client_request_id_test.go b/backend/internal/server/middleware/client_request_id_test.go
new file mode 100644
index 00000000..394c1612
--- /dev/null
+++ b/backend/internal/server/middleware/client_request_id_test.go
@@ -0,0 +1,50 @@
+package middleware
+
+import (
+ "context"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+)
+
+func TestClientRequestIDGeneratesAndExposesID(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ router := gin.New()
+ router.Use(ClientRequestID())
+ router.GET("/", func(c *gin.Context) {
+ value, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
+ c.String(http.StatusOK, value)
+ })
+
+ req := httptest.NewRequest(http.MethodGet, "/", nil)
+ w := httptest.NewRecorder()
+
+ router.ServeHTTP(w, req)
+
+ require.Equal(t, http.StatusOK, w.Code)
+ require.NotEmpty(t, w.Body.String())
+ require.Equal(t, w.Body.String(), w.Header().Get(clientRequestIDHeader))
+}
+
+func TestClientRequestIDPreservesExistingContextID(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ router := gin.New()
+ router.Use(ClientRequestID())
+ router.GET("/", func(c *gin.Context) {
+ value, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string)
+ c.String(http.StatusOK, value)
+ })
+
+ w := httptest.NewRecorder()
+ req := httptest.NewRequest(http.MethodGet, "/", nil)
+ req = req.WithContext(context.WithValue(req.Context(), ctxkey.ClientRequestID, "existing-client-request-id"))
+ router.ServeHTTP(w, req)
+
+ require.Equal(t, http.StatusOK, w.Code)
+ require.Equal(t, "existing-client-request-id", w.Body.String())
+ require.Equal(t, "existing-client-request-id", w.Header().Get(clientRequestIDHeader))
+}
diff --git a/backend/internal/server/middleware/middleware.go b/backend/internal/server/middleware/middleware.go
index d42eacec..9efe78a3 100644
--- a/backend/internal/server/middleware/middleware.go
+++ b/backend/internal/server/middleware/middleware.go
@@ -24,6 +24,11 @@ const (
ContextKeySubscription ContextKey = "subscription"
// ContextKeyForcePlatform 强制平台(用于 /antigravity 路由)
ContextKeyForcePlatform ContextKey = "force_platform"
+ // ContextKeyOpsFallbackAPIKey 运维错误日志专用回退键。
+ // 鉴权早退(分组停用/删除、Key 停用/过期/额度、用户停用、IP 限制等)时,
+ // apiKey 已加载但尚未写入 ContextKeyAPIKey;该键让 Ops 错误日志仍能取到
+ // user/group/platform。仅供 Ops 错误日志读取,不代表请求已通过鉴权。
+ ContextKeyOpsFallbackAPIKey ContextKey = "ops_fallback_api_key"
)
// ForcePlatform 返回设置强制平台的中间件
diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go
index 349c520c..9a3253b5 100644
--- a/backend/internal/server/routes/admin.go
+++ b/backend/internal/server/routes/admin.go
@@ -259,6 +259,7 @@ func registerGroupRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
groups.GET("/usage-summary", h.Admin.Group.GetUsageSummary)
groups.GET("/capacity-summary", h.Admin.Group.GetCapacitySummary)
groups.PUT("/sort-order", h.Admin.Group.UpdateSortOrder)
+ groups.GET("/:id/models-list-candidates", h.Admin.Group.GetModelsListCandidates)
groups.GET("/:id", h.Admin.Group.GetByID)
groups.POST("", h.Admin.Group.Create)
groups.PUT("/:id", h.Admin.Group.Update)
@@ -301,6 +302,7 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
accounts.GET("/:id/temp-unschedulable", h.Admin.Account.GetTempUnschedulable)
accounts.DELETE("/:id/temp-unschedulable", h.Admin.Account.ClearTempUnschedulable)
accounts.POST("/:id/schedulable", h.Admin.Account.SetSchedulable)
+ accounts.POST("/models/sync-upstream-preview", h.Admin.Account.SyncUpstreamModelsPreview)
accounts.GET("/:id/models", h.Admin.Account.GetAvailableModels)
accounts.POST("/:id/models/sync-upstream", h.Admin.Account.SyncUpstreamModels)
accounts.POST("/batch", h.Admin.Account.BatchCreate)
diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go
index efc0687f..b039a6ec 100644
--- a/backend/internal/server/routes/gateway.go
+++ b/backend/internal/server/routes/gateway.go
@@ -89,6 +89,19 @@ func RegisterGatewayRoutes(
}
h.Gateway.ChatCompletions(c)
})
+ gateway.POST("/embeddings", func(c *gin.Context) {
+ if getGroupPlatform(c) != service.PlatformOpenAI {
+ service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
+ c.JSON(http.StatusNotFound, gin.H{
+ "error": gin.H{
+ "type": "not_found_error",
+ "message": "Embeddings API is not supported for this platform",
+ },
+ })
+ return
+ }
+ h.OpenAIGateway.Embeddings(c)
+ })
gateway.POST("/images/generations", func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
@@ -158,6 +171,19 @@ func RegisterGatewayRoutes(
}
h.Gateway.ChatCompletions(c)
})
+ r.POST("/embeddings", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
+ if getGroupPlatform(c) != service.PlatformOpenAI {
+ service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
+ c.JSON(http.StatusNotFound, gin.H{
+ "error": gin.H{
+ "type": "not_found_error",
+ "message": "Embeddings API is not supported for this platform",
+ },
+ })
+ return
+ }
+ h.OpenAIGateway.Embeddings(c)
+ })
r.POST("/images/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, func(c *gin.Context) {
if getGroupPlatform(c) != service.PlatformOpenAI {
service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate)
diff --git a/backend/internal/server/routes/user.go b/backend/internal/server/routes/user.go
index 07ae33de..0f3758f7 100644
--- a/backend/internal/server/routes/user.go
+++ b/backend/internal/server/routes/user.go
@@ -82,6 +82,8 @@ func RegisterUserRoutes(
usage := authenticated.Group("/usage")
{
usage.GET("", h.Usage.List)
+ usage.GET("/errors", h.Usage.ListErrors)
+ usage.GET("/errors/:id", h.Usage.GetErrorDetail)
usage.GET("/:id", h.Usage.GetByID)
usage.GET("/stats", h.Usage.Stats)
// User dashboard endpoints
diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go
index cd06ffa3..fb95201f 100644
--- a/backend/internal/service/account.go
+++ b/backend/internal/service/account.go
@@ -66,6 +66,15 @@ type Account struct {
modelMappingCacheRawSig uint64
}
+type OpenAIEndpointCapability string
+
+const (
+ OpenAIEndpointCapabilityChatCompletions OpenAIEndpointCapability = "chat_completions"
+ OpenAIEndpointCapabilityEmbeddings OpenAIEndpointCapability = "embeddings"
+)
+
+const openAIEndpointCapabilitiesCredentialKey = "openai_capabilities"
+
type TempUnschedulableRule struct {
ErrorCode int `json:"error_code"`
Keywords []string `json:"keywords"`
@@ -890,14 +899,90 @@ func parsePoolModeRetryCount(value any) int {
return defaultPoolModeRetryCount
}
-// isPoolModeRetryableStatus 池模式下应触发同账号重试的状态码
+// defaultPoolModeRetryableStatusCodes 池模式下默认触发同账号重试的状态码。
+// 未在 Account.Credentials 中显式配置 pool_mode_retry_status_codes 时使用。
+var defaultPoolModeRetryableStatusCodes = []int{401, 403, 429}
+
+// isPoolModeRetryableStatus 池模式下应触发同账号重试的状态码(默认列表)。
func isPoolModeRetryableStatus(statusCode int) bool {
- switch statusCode {
- case 401, 403, 429:
- return true
- default:
- return false
+ for _, c := range defaultPoolModeRetryableStatusCodes {
+ if c == statusCode {
+ return true
+ }
}
+ return false
+}
+
+// GetPoolModeRetryStatusCodes 返回账号自定义的池模式同账号重试状态码列表。
+//
+// 返回值语义:
+// - nil:未配置 → 调用方应回退到默认值 [401, 403, 429]
+// - 长度为 0 的切片:管理员显式置空 → 关闭按状态码触发的同账号重试
+// - 非空切片:去重、过滤为合法 HTTP 状态码(100-599)后的覆盖列表
+func (a *Account) GetPoolModeRetryStatusCodes() []int {
+ if a == nil || a.Credentials == nil {
+ return nil
+ }
+ raw, ok := a.Credentials["pool_mode_retry_status_codes"]
+ if !ok || raw == nil {
+ return nil
+ }
+ arr, ok := raw.([]any)
+ if !ok {
+ return nil
+ }
+ seen := make(map[int]struct{}, len(arr))
+ codes := make([]int, 0, len(arr))
+ for _, v := range arr {
+ var code int
+ switch n := v.(type) {
+ case float64:
+ code = int(n)
+ case int:
+ code = n
+ case int64:
+ code = int(n)
+ case json.Number:
+ i, err := n.Int64()
+ if err != nil {
+ continue
+ }
+ code = int(i)
+ case string:
+ i, err := strconv.Atoi(strings.TrimSpace(n))
+ if err != nil {
+ continue
+ }
+ code = i
+ default:
+ continue
+ }
+ if code < 100 || code > 599 {
+ continue
+ }
+ if _, exists := seen[code]; exists {
+ continue
+ }
+ seen[code] = struct{}{}
+ codes = append(codes, code)
+ }
+ sort.Ints(codes)
+ return codes
+}
+
+// IsPoolModeRetryableStatus 在账号上下文中判断给定状态码是否应触发同账号重试。
+// 若账号未配置 pool_mode_retry_status_codes,则回退到默认列表。
+func (a *Account) IsPoolModeRetryableStatus(statusCode int) bool {
+ codes := a.GetPoolModeRetryStatusCodes()
+ if codes == nil {
+ return isPoolModeRetryableStatus(statusCode)
+ }
+ for _, c := range codes {
+ if c == statusCode {
+ return true
+ }
+ }
+ return false
}
func (a *Account) GetCustomErrorCodes() []int {
@@ -1046,6 +1131,80 @@ func (a *Account) GetOpenAISessionID() string {
return strings.TrimSpace(a.GetExtraString("openai_session_id"))
}
+func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapability) bool {
+ if a == nil {
+ return false
+ }
+ if capability == "" {
+ return true
+ }
+ if !a.IsOpenAI() {
+ return false
+ }
+ switch capability {
+ case OpenAIEndpointCapabilityChatCompletions:
+ case OpenAIEndpointCapabilityEmbeddings:
+ if a.Type != AccountTypeAPIKey {
+ return false
+ }
+ default:
+ return false
+ }
+
+ configured, found := a.openAIEndpointCapabilitySet()
+ if !found {
+ return true
+ }
+ return configured[string(capability)]
+}
+
+func (a *Account) openAIEndpointCapabilitySet() (map[string]bool, bool) {
+ if a == nil || a.Credentials == nil {
+ return nil, false
+ }
+ raw, found := a.Credentials[openAIEndpointCapabilitiesCredentialKey]
+ if !found || raw == nil {
+ return nil, false
+ }
+
+ result := make(map[string]bool)
+ add := func(value string) {
+ value = strings.ToLower(strings.TrimSpace(value))
+ if value == "" {
+ return
+ }
+ result[value] = true
+ }
+
+ switch capabilities := raw.(type) {
+ case []any:
+ for _, item := range capabilities {
+ if value, ok := item.(string); ok {
+ add(value)
+ }
+ }
+ case []string:
+ for _, value := range capabilities {
+ add(value)
+ }
+ case map[string]any:
+ for key, value := range capabilities {
+ enabled, ok := value.(bool)
+ if ok && enabled {
+ add(key)
+ }
+ }
+ case map[string]bool:
+ for key, enabled := range capabilities {
+ if enabled {
+ add(key)
+ }
+ }
+ }
+
+ return result, true
+}
+
func (a *Account) SupportsOpenAIImageCapability(capability OpenAIImagesCapability) bool {
if !a.IsOpenAI() {
return false
@@ -1366,6 +1525,38 @@ func (a *Account) IsCodexCLIOnlyEnabled() bool {
return ok && enabled
}
+// GetCodexCLIOnlyAllowedClients 返回 codex_cli_only 之上额外放行的命名客户端预设 ID 列表。
+// 仅 OpenAI OAuth 账号生效;缺失或类型不符时返回空。预设 ID 的具体匹配规则由
+// openai 包的 registry 固化,配置只能引用预设键、不能自定义规则。
+func (a *Account) GetCodexCLIOnlyAllowedClients() []string {
+ if a == nil || !a.IsOpenAIOAuth() || a.Extra == nil {
+ return nil
+ }
+ raw, ok := a.Extra["codex_cli_only_allowed_clients"]
+ if !ok || raw == nil {
+ return nil
+ }
+ switch v := raw.(type) {
+ case []string:
+ result := make([]string, 0, len(v))
+ for _, s := range v {
+ if strings.TrimSpace(s) != "" {
+ result = append(result, s)
+ }
+ }
+ return result
+ case []any:
+ result := make([]string, 0, len(v))
+ for _, item := range v {
+ if s, ok := item.(string); ok && strings.TrimSpace(s) != "" {
+ result = append(result, s)
+ }
+ }
+ return result
+ }
+ return nil
+}
+
// WindowCostSchedulability 窗口费用调度状态
type WindowCostSchedulability int
diff --git a/backend/internal/service/account_codex_cli_only_allowed_clients_test.go b/backend/internal/service/account_codex_cli_only_allowed_clients_test.go
new file mode 100644
index 00000000..c835ea27
--- /dev/null
+++ b/backend/internal/service/account_codex_cli_only_allowed_clients_test.go
@@ -0,0 +1,68 @@
+package service
+
+import (
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestAccount_GetCodexCLIOnlyAllowedClients(t *testing.T) {
+ t.Run("OAuth 账号读取 []any 字符串列表", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{"codex_cli_only_allowed_clients": []any{"claude_code"}},
+ }
+ require.Equal(t, []string{"claude_code"}, account.GetCodexCLIOnlyAllowedClients())
+ })
+
+ t.Run("OAuth 账号读取 []string 列表", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{"codex_cli_only_allowed_clients": []string{"claude_code"}},
+ }
+ require.Equal(t, []string{"claude_code"}, account.GetCodexCLIOnlyAllowedClients())
+ })
+
+ t.Run("[]string 跳过空白元素", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{"codex_cli_only_allowed_clients": []string{"claude_code", "", " "}},
+ }
+ require.Equal(t, []string{"claude_code"}, account.GetCodexCLIOnlyAllowedClients())
+ })
+
+ t.Run("跳过非字符串与空白元素", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{"codex_cli_only_allowed_clients": []any{"claude_code", 123, "", " "}},
+ }
+ require.Equal(t, []string{"claude_code"}, account.GetCodexCLIOnlyAllowedClients())
+ })
+
+ t.Run("非 OAuth 账号返回空", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Extra: map[string]any{"codex_cli_only_allowed_clients": []any{"claude_code"}},
+ }
+ require.Empty(t, account.GetCodexCLIOnlyAllowedClients())
+ })
+
+ t.Run("Extra 为空返回空", func(t *testing.T) {
+ account := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}
+ require.Empty(t, account.GetCodexCLIOnlyAllowedClients())
+ })
+
+ t.Run("字段缺失返回空", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{},
+ }
+ require.Empty(t, account.GetCodexCLIOnlyAllowedClients())
+ })
+}
diff --git a/backend/internal/service/account_pool_retry_status_codes_test.go b/backend/internal/service/account_pool_retry_status_codes_test.go
new file mode 100644
index 00000000..c0b9d7ab
--- /dev/null
+++ b/backend/internal/service/account_pool_retry_status_codes_test.go
@@ -0,0 +1,193 @@
+//go:build unit
+
+package service
+
+import (
+ "encoding/json"
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestGetPoolModeRetryStatusCodes(t *testing.T) {
+ tests := []struct {
+ name string
+ account *Account
+ expected []int
+ }{
+ {
+ name: "nil_account_returns_nil",
+ account: nil,
+ expected: nil,
+ },
+ {
+ name: "nil_credentials_returns_nil",
+ account: &Account{
+ Type: AccountTypeAPIKey,
+ Platform: PlatformOpenAI,
+ },
+ expected: nil,
+ },
+ {
+ name: "missing_key_returns_nil",
+ account: &Account{
+ Type: AccountTypeAPIKey,
+ Platform: PlatformOpenAI,
+ Credentials: map[string]any{"pool_mode": true},
+ },
+ expected: nil,
+ },
+ {
+ name: "empty_slice_is_preserved",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{},
+ },
+ },
+ expected: []int{},
+ },
+ {
+ name: "float64_values_from_json_are_normalized",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{float64(429), float64(401), float64(403)},
+ },
+ },
+ expected: []int{401, 403, 429},
+ },
+ {
+ name: "json_number_values_supported",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{json.Number("502"), json.Number("503")},
+ },
+ },
+ expected: []int{502, 503},
+ },
+ {
+ name: "string_values_supported",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{"520", "529"},
+ },
+ },
+ expected: []int{520, 529},
+ },
+ {
+ name: "duplicates_are_deduped",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{float64(429), float64(429), float64(401)},
+ },
+ },
+ expected: []int{401, 429},
+ },
+ {
+ name: "out_of_range_values_dropped",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{float64(99), float64(600), float64(429)},
+ },
+ },
+ expected: []int{429},
+ },
+ {
+ name: "invalid_string_dropped",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{"oops", float64(429)},
+ },
+ },
+ expected: []int{429},
+ },
+ {
+ name: "non_array_value_returns_nil",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": "not-an-array",
+ },
+ },
+ expected: nil,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ require.Equal(t, tt.expected, tt.account.GetPoolModeRetryStatusCodes())
+ })
+ }
+}
+
+func TestIsPoolModeRetryableStatus_Account(t *testing.T) {
+ tests := []struct {
+ name string
+ account *Account
+ statusCode int
+ expected bool
+ }{
+ {
+ name: "nil_account_falls_back_to_default_401",
+ account: nil,
+ statusCode: 401,
+ expected: true,
+ },
+ {
+ name: "nil_account_falls_back_to_default_500",
+ account: nil,
+ statusCode: 500,
+ expected: false,
+ },
+ {
+ name: "unconfigured_uses_default_403",
+ account: &Account{
+ Credentials: map[string]any{"pool_mode": true},
+ },
+ statusCode: 403,
+ expected: true,
+ },
+ {
+ name: "unconfigured_uses_default_502_false",
+ account: &Account{
+ Credentials: map[string]any{"pool_mode": true},
+ },
+ statusCode: 502,
+ expected: false,
+ },
+ {
+ name: "configured_list_overrides_default_401_dropped",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{float64(502), float64(503)},
+ },
+ },
+ statusCode: 401,
+ expected: false,
+ },
+ {
+ name: "configured_list_overrides_default_502_added",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{float64(502), float64(503)},
+ },
+ },
+ statusCode: 502,
+ expected: true,
+ },
+ {
+ name: "empty_list_disables_all_default_codes",
+ account: &Account{
+ Credentials: map[string]any{
+ "pool_mode_retry_status_codes": []any{},
+ },
+ },
+ statusCode: 429,
+ expected: false,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ require.Equal(t, tt.expected, tt.account.IsPoolModeRetryableStatus(tt.statusCode))
+ })
+ }
+}
diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go
index 3189a729..748840b7 100644
--- a/backend/internal/service/account_service.go
+++ b/backend/internal/service/account_service.go
@@ -60,7 +60,7 @@ type AccountRepository interface {
ListSchedulableUngroupedByPlatforms(ctx context.Context, platforms []string) ([]Account, error)
SetRateLimited(ctx context.Context, id int64, resetAt time.Time) error
- SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time) error
+ SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error
SetOverloaded(ctx context.Context, id int64, until time.Time) error
SetTempUnschedulable(ctx context.Context, id int64, until time.Time, reason string) error
ClearTempUnschedulable(ctx context.Context, id int64) error
diff --git a/backend/internal/service/account_service_delete_test.go b/backend/internal/service/account_service_delete_test.go
index 81169a02..d72554ce 100644
--- a/backend/internal/service/account_service_delete_test.go
+++ b/backend/internal/service/account_service_delete_test.go
@@ -159,7 +159,7 @@ func (s *accountRepoStub) SetRateLimited(ctx context.Context, id int64, resetAt
panic("unexpected SetRateLimited call")
}
-func (s *accountRepoStub) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time) error {
+func (s *accountRepoStub) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error {
panic("unexpected SetModelRateLimit call")
}
diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go
index fc8f3fbb..ae9dd8f6 100644
--- a/backend/internal/service/admin_service.go
+++ b/backend/internal/service/admin_service.go
@@ -17,9 +17,13 @@ import (
dbent "github.com/Wei-Shaw/sub2api/ent"
"github.com/Wei-Shaw/sub2api/ent/authidentity"
"github.com/Wei-Shaw/sub2api/ent/authidentitychannel"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/claude"
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/geminicli"
"github.com/Wei-Shaw/sub2api/internal/pkg/httpclient"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/openai"
"github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
"github.com/Wei-Shaw/sub2api/internal/util/httputil"
)
@@ -29,6 +33,7 @@ type AdminService interface {
// User management
ListUsers(ctx context.Context, page, pageSize int, filters UserListFilters, sortBy, sortOrder string) ([]User, int64, error)
GetUser(ctx context.Context, id int64) (*User, error)
+ GetUserIncludeDeleted(ctx context.Context, id int64) (*User, error)
CreateUser(ctx context.Context, input *CreateUserInput) (*User, error)
UpdateUser(ctx context.Context, id int64, input *UpdateUserInput) (*User, error)
DeleteUser(ctx context.Context, id int64) error
@@ -48,6 +53,7 @@ type AdminService interface {
GetAllGroups(ctx context.Context) ([]Group, error)
GetAllGroupsByPlatform(ctx context.Context, platform string) ([]Group, error)
GetGroup(ctx context.Context, id int64) (*Group, error)
+ GetGroupModelsListCandidates(ctx context.Context, id int64, platform string) ([]string, error)
CreateGroup(ctx context.Context, input *CreateGroupInput) (*Group, error)
UpdateGroup(ctx context.Context, id int64, input *UpdateGroupInput) (*Group, error)
DeleteGroup(ctx context.Context, id int64) error
@@ -123,7 +129,7 @@ type CreateUserInput struct {
Password string
Username string
Notes string
- Balance float64
+ Balance *float64
Concurrency int
RPMLimit int
AllowedGroups []int64
@@ -215,6 +221,7 @@ type CreateGroupInput struct {
RequireOAuthOnly bool
RequirePrivacySet bool
MessagesDispatchModelConfig OpenAIMessagesDispatchModelConfig
+ ModelsListConfig GroupModelsListConfig
// RPMLimit 分组 RPM 上限(0 = 不限制)
RPMLimit int
// 从指定分组复制账号(创建分组后在同一事务内绑定)
@@ -223,7 +230,7 @@ type CreateGroupInput struct {
type UpdateGroupInput struct {
Name string
- Description string
+ Description *string
Platform string
RateMultiplier *float64 // 使用指针以支持设置为0
IsExclusive *bool
@@ -255,6 +262,7 @@ type UpdateGroupInput struct {
RequireOAuthOnly *bool
RequirePrivacySet *bool
MessagesDispatchModelConfig *OpenAIMessagesDispatchModelConfig
+ ModelsListConfig *GroupModelsListConfig
// RPMLimit 分组 RPM 上限(0 = 不限制),nil 表示未提供不改动。
RPMLimit *int
// 从指定分组复制账号(同步操作:先清空当前分组的账号绑定,再绑定源分组的账号)
@@ -667,13 +675,24 @@ func (s *adminServiceImpl) GetUser(ctx context.Context, id int64) (*User, error)
return user, nil
}
+func (s *adminServiceImpl) GetUserIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return s.userRepo.GetByIDIncludeDeleted(ctx, id)
+}
+
func (s *adminServiceImpl) CreateUser(ctx context.Context, input *CreateUserInput) (*User, error) {
+ balance := 0.0
+ if input.Balance != nil {
+ balance = *input.Balance
+ } else if s.settingService != nil {
+ balance = s.settingService.GetDefaultBalance(ctx)
+ }
+
user := &User{
Email: input.Email,
Username: input.Username,
Notes: input.Notes,
Role: RoleUser, // Always create as regular user, never admin
- Balance: input.Balance,
+ Balance: balance,
Concurrency: input.Concurrency,
RPMLimit: input.RPMLimit,
Status: StatusActive,
@@ -1582,6 +1601,80 @@ func (s *adminServiceImpl) GetGroup(ctx context.Context, id int64) (*Group, erro
return s.groupRepo.GetByID(ctx, id)
}
+func (s *adminServiceImpl) GetGroupModelsListCandidates(ctx context.Context, id int64, platform string) ([]string, error) {
+ platform = strings.TrimSpace(platform)
+ if id > 0 {
+ group, err := s.groupRepo.GetByIDLite(ctx, id)
+ if err != nil {
+ return nil, err
+ }
+ if platform == "" {
+ platform = group.Platform
+ }
+ }
+ if platform == "" {
+ platform = PlatformAnthropic
+ }
+
+ candidates := defaultModelsListCandidateIDs(platform)
+ if id <= 0 || s.accountRepo == nil {
+ return candidates, nil
+ }
+
+ accounts, err := s.accountRepo.ListSchedulableByGroupID(ctx, id)
+ if err != nil {
+ return nil, err
+ }
+
+ seen := make(map[string]struct{}, len(candidates))
+ for _, model := range candidates {
+ seen[model] = struct{}{}
+ }
+ for _, acc := range accounts {
+ if acc.Platform != platform {
+ continue
+ }
+ for model := range acc.GetModelMapping() {
+ model = strings.TrimSpace(model)
+ if model == "" {
+ continue
+ }
+ if _, ok := seen[model]; ok {
+ continue
+ }
+ seen[model] = struct{}{}
+ candidates = append(candidates, model)
+ }
+ }
+ return candidates, nil
+}
+
+func defaultModelsListCandidateIDs(platform string) []string {
+ switch platform {
+ case PlatformOpenAI:
+ return openai.DefaultModelIDs()
+ case PlatformGemini:
+ ids := make([]string, 0, len(geminicli.DefaultModels))
+ for _, model := range geminicli.DefaultModels {
+ ids = append(ids, model.ID)
+ }
+ return ids
+ case PlatformAntigravity:
+ models := antigravity.DefaultModels()
+ ids := make([]string, 0, len(models))
+ for _, model := range models {
+ ids = append(ids, model.ID)
+ }
+ return ids
+ default:
+ ids := make([]string, 0, len(claude.DefaultModels))
+ for _, model := range claude.DefaultModels {
+ ids = append(ids, model.ID)
+ }
+ return ids
+ }
+}
+
func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupInput) (*Group, error) {
if input.RateMultiplier <= 0 {
return nil, errors.New("rate_multiplier must be > 0")
@@ -1697,6 +1790,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn
RequirePrivacySet: input.RequirePrivacySet,
DefaultMappedModel: input.DefaultMappedModel,
MessagesDispatchModelConfig: normalizeOpenAIMessagesDispatchModelConfig(input.MessagesDispatchModelConfig),
+ ModelsListConfig: normalizeGroupModelsListConfig(input.ModelsListConfig),
RPMLimit: input.RPMLimit,
}
sanitizeGroupMessagesDispatchFields(group)
@@ -1830,8 +1924,8 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd
if input.Name != "" {
group.Name = input.Name
}
- if input.Description != "" {
- group.Description = input.Description
+ if input.Description != nil {
+ group.Description = *input.Description
}
if input.Platform != "" {
group.Platform = input.Platform
@@ -1944,6 +2038,9 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd
if input.MessagesDispatchModelConfig != nil {
group.MessagesDispatchModelConfig = normalizeOpenAIMessagesDispatchModelConfig(*input.MessagesDispatchModelConfig)
}
+ if input.ModelsListConfig != nil {
+ group.ModelsListConfig = normalizeGroupModelsListConfig(*input.ModelsListConfig)
+ }
if input.RPMLimit != nil {
group.RPMLimit = *input.RPMLimit
}
diff --git a/backend/internal/service/admin_service_apikey_test.go b/backend/internal/service/admin_service_apikey_test.go
index 3b3dbc21..ccc8d221 100644
--- a/backend/internal/service/admin_service_apikey_test.go
+++ b/backend/internal/service/admin_service_apikey_test.go
@@ -69,8 +69,12 @@ func (s *userRepoStubForGroupUpdate) UpdateConcurrency(context.Context, int64, i
panic("unexpected")
}
-func (s *userRepoStubForGroupUpdate) BatchSetConcurrency(context.Context, []int64, int) (int, error) { return 0, nil }
-func (s *userRepoStubForGroupUpdate) BatchAddConcurrency(context.Context, []int64, int) (int, error) { return 0, nil }
+func (s *userRepoStubForGroupUpdate) BatchSetConcurrency(context.Context, []int64, int) (int, error) {
+ return 0, nil
+}
+func (s *userRepoStubForGroupUpdate) BatchAddConcurrency(context.Context, []int64, int) (int, error) {
+ return 0, nil
+}
func (s *userRepoStubForGroupUpdate) ExistsByEmail(context.Context, string) (bool, error) {
panic("unexpected")
}
@@ -82,6 +86,9 @@ func (s *userRepoStubForGroupUpdate) UpdateTotpSecret(context.Context, int64, *s
}
func (s *userRepoStubForGroupUpdate) EnableTotp(context.Context, int64) error { panic("unexpected") }
func (s *userRepoStubForGroupUpdate) DisableTotp(context.Context, int64) error { panic("unexpected") }
+func (s *userRepoStubForGroupUpdate) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ panic("unexpected GetByIDIncludeDeleted call")
+}
func (s *userRepoStubForGroupUpdate) ListUserAuthIdentities(context.Context, int64) ([]UserAuthIdentityRecord, error) {
panic("unexpected")
}
@@ -139,6 +146,9 @@ func (s *apiKeyRepoStubForGroupUpdate) GetByKeyForAuth(context.Context, string)
panic("unexpected")
}
func (s *apiKeyRepoStubForGroupUpdate) Delete(context.Context, int64) error { panic("unexpected") }
+func (s *apiKeyRepoStubForGroupUpdate) DeleteWithAudit(context.Context, int64) error {
+ panic("unexpected")
+}
func (s *apiKeyRepoStubForGroupUpdate) ListByUserID(context.Context, int64, pagination.PaginationParams, APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error) {
panic("unexpected")
}
diff --git a/backend/internal/service/admin_service_create_user_test.go b/backend/internal/service/admin_service_create_user_test.go
index c5b1e38d..5e9578ab 100644
--- a/backend/internal/service/admin_service_create_user_test.go
+++ b/backend/internal/service/admin_service_create_user_test.go
@@ -14,13 +14,14 @@ import (
func TestAdminService_CreateUser_Success(t *testing.T) {
repo := &userRepoStub{nextID: 10}
svc := &adminServiceImpl{userRepo: repo}
+ balance := 12.5
input := &CreateUserInput{
Email: "user@test.com",
Password: "strong-pass",
Username: "tester",
Notes: "note",
- Balance: 12.5,
+ Balance: &balance,
Concurrency: 7,
AllowedGroups: []int64{3, 5},
}
@@ -32,7 +33,7 @@ func TestAdminService_CreateUser_Success(t *testing.T) {
require.Equal(t, input.Email, user.Email)
require.Equal(t, input.Username, user.Username)
require.Equal(t, input.Notes, user.Notes)
- require.Equal(t, input.Balance, user.Balance)
+ require.Equal(t, balance, user.Balance)
require.Equal(t, input.Concurrency, user.Concurrency)
require.Equal(t, input.AllowedGroups, user.AllowedGroups)
require.Equal(t, RoleUser, user.Role)
@@ -42,6 +43,56 @@ func TestAdminService_CreateUser_Success(t *testing.T) {
require.Equal(t, user, repo.created[0])
}
+func TestAdminService_CreateUser_UsesDefaultBalanceWhenBalanceOmitted(t *testing.T) {
+ repo := &userRepoStub{nextID: 11}
+ cfg := &config.Config{
+ Default: config.DefaultConfig{
+ UserBalance: 0,
+ },
+ }
+ settingService := NewSettingService(&settingRepoStub{values: map[string]string{
+ SettingKeyDefaultBalance: "0.02",
+ }}, cfg)
+ svc := &adminServiceImpl{userRepo: repo, settingService: settingService}
+
+ user, err := svc.CreateUser(context.Background(), &CreateUserInput{
+ Email: "default-balance@test.com",
+ Password: "strong-pass",
+ })
+
+ require.NoError(t, err)
+ require.NotNil(t, user)
+ require.Equal(t, 0.02, user.Balance)
+ require.Len(t, repo.created, 1)
+ require.Equal(t, 0.02, repo.created[0].Balance)
+}
+
+func TestAdminService_CreateUser_ExplicitZeroBalanceOverridesDefault(t *testing.T) {
+ repo := &userRepoStub{nextID: 12}
+ cfg := &config.Config{
+ Default: config.DefaultConfig{
+ UserBalance: 0,
+ },
+ }
+ settingService := NewSettingService(&settingRepoStub{values: map[string]string{
+ SettingKeyDefaultBalance: "0.02",
+ }}, cfg)
+ svc := &adminServiceImpl{userRepo: repo, settingService: settingService}
+ balance := 0.0
+
+ user, err := svc.CreateUser(context.Background(), &CreateUserInput{
+ Email: "zero-balance@test.com",
+ Password: "strong-pass",
+ Balance: &balance,
+ })
+
+ require.NoError(t, err)
+ require.NotNil(t, user)
+ require.Equal(t, 0.0, user.Balance)
+ require.Len(t, repo.created, 1)
+ require.Equal(t, 0.0, repo.created[0].Balance)
+}
+
func TestAdminService_CreateUser_EmailExists(t *testing.T) {
repo := &userRepoStub{createErr: ErrEmailExists}
svc := &adminServiceImpl{userRepo: repo}
diff --git a/backend/internal/service/admin_service_delete_test.go b/backend/internal/service/admin_service_delete_test.go
index d01b11e6..150c4f53 100644
--- a/backend/internal/service/admin_service_delete_test.go
+++ b/backend/internal/service/admin_service_delete_test.go
@@ -173,6 +173,10 @@ func (s *userRepoStub) DisableTotp(ctx context.Context, userID int64) error {
panic("unexpected DisableTotp call")
}
+func (s *userRepoStub) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return s.GetByID(ctx, id)
+}
+
type groupRepoStub struct {
affectedUserIDs []int64
deleteErr error
@@ -471,10 +475,22 @@ func (s *billingCacheStub) DeleteUserPlatformQuotaCache(ctx context.Context, use
panic("unexpected DeleteUserPlatformQuotaCache call")
}
-func (s *billingCacheStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error {
+func (s *billingCacheStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error {
panic("unexpected IncrUserPlatformQuotaUsageCache call")
}
+func (s *billingCacheStub) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) {
+ panic("unexpected PopDirtyUserPlatformQuotaKeys call")
+}
+
+func (s *billingCacheStub) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error {
+ panic("unexpected ReaddDirtyUserPlatformQuotaKeys call")
+}
+
+func (s *billingCacheStub) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) {
+ panic("unexpected BatchGetUserPlatformQuotaCache call")
+}
+
func waitForInvalidations(t *testing.T, ch <-chan subscriptionInvalidateCall, expected int) []subscriptionInvalidateCall {
t.Helper()
calls := make([]subscriptionInvalidateCall, 0, expected)
diff --git a/backend/internal/service/admin_service_email_identity_sync_test.go b/backend/internal/service/admin_service_email_identity_sync_test.go
index c791b747..c3737f5a 100644
--- a/backend/internal/service/admin_service_email_identity_sync_test.go
+++ b/backend/internal/service/admin_service_email_identity_sync_test.go
@@ -113,8 +113,12 @@ func (s *emailSyncRepoStub) RemoveGroupFromAllowedGroups(context.Context, int64)
return 0, nil
}
-func (s *emailSyncRepoStub) BatchSetConcurrency(context.Context, []int64, int) (int, error) { return 0, nil }
-func (s *emailSyncRepoStub) BatchAddConcurrency(context.Context, []int64, int) (int, error) { return 0, nil }
+func (s *emailSyncRepoStub) BatchSetConcurrency(context.Context, []int64, int) (int, error) {
+ return 0, nil
+}
+func (s *emailSyncRepoStub) BatchAddConcurrency(context.Context, []int64, int) (int, error) {
+ return 0, nil
+}
func (s *emailSyncRepoStub) AddGroupToAllowedGroups(context.Context, int64, int64) error { return nil }
@@ -133,6 +137,9 @@ func (s *emailSyncRepoStub) UpdateTotpSecret(context.Context, int64, *string) er
func (s *emailSyncRepoStub) EnableTotp(context.Context, int64) error { return nil }
func (s *emailSyncRepoStub) DisableTotp(context.Context, int64) error { return nil }
+func (s *emailSyncRepoStub) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return s.GetByID(ctx, id)
+}
func (s *emailSyncRepoStub) EnsureEmailAuthIdentity(_ context.Context, userID int64, email string) error {
s.ensureCalls = append(s.ensureCalls, ensureEmailCall{userID: userID, email: email})
diff --git a/backend/internal/service/admin_service_get_deleted_test.go b/backend/internal/service/admin_service_get_deleted_test.go
new file mode 100644
index 00000000..6ad17f60
--- /dev/null
+++ b/backend/internal/service/admin_service_get_deleted_test.go
@@ -0,0 +1,22 @@
+//go:build unit
+
+package service
+
+import (
+ "context"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/require"
+)
+
+func TestAdminService_GetUserIncludeDeleted(t *testing.T) {
+ ts := time.Date(2026, 5, 28, 0, 0, 0, 0, time.UTC)
+ repo := &userRepoStub{user: &User{ID: 7, Email: "del@test.com", DeletedAt: &ts}}
+ svc := &adminServiceImpl{userRepo: repo}
+
+ got, err := svc.GetUserIncludeDeleted(context.Background(), 7)
+ require.NoError(t, err)
+ require.Equal(t, int64(7), got.ID)
+ require.NotNil(t, got.DeletedAt)
+}
diff --git a/backend/internal/service/admin_service_group_test.go b/backend/internal/service/admin_service_group_test.go
index 0a2020ea..eb3eff7f 100644
--- a/backend/internal/service/admin_service_group_test.go
+++ b/backend/internal/service/admin_service_group_test.go
@@ -280,8 +280,9 @@ func TestAdminService_UpdateGroup_PreservesImageGenerationControlsWhenOmitted(t
repo := &groupRepoStubForAdmin{getByID: existingGroup}
svc := &adminServiceImpl{groupRepo: repo}
+ updatedDesc := "updated"
group, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{
- Description: "updated",
+ Description: &updatedDesc,
})
require.NoError(t, err)
require.NotNil(t, group)
@@ -291,6 +292,45 @@ func TestAdminService_UpdateGroup_PreservesImageGenerationControlsWhenOmitted(t
require.InDelta(t, 0.5, repo.updated.ImageRateMultiplier, 1e-12)
}
+func TestAdminService_UpdateGroup_ClearsDescriptionWhenEmptyString(t *testing.T) {
+ existingGroup := &Group{
+ ID: 1,
+ Name: "existing-group",
+ Description: "Auto-created default group",
+ Platform: PlatformOpenAI,
+ Status: StatusActive,
+ }
+ repo := &groupRepoStubForAdmin{getByID: existingGroup}
+ svc := &adminServiceImpl{groupRepo: repo}
+
+ empty := ""
+ _, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{
+ Description: &empty,
+ })
+ require.NoError(t, err)
+ require.NotNil(t, repo.updated)
+ require.Equal(t, "", repo.updated.Description, "empty string should clear description")
+}
+
+func TestAdminService_UpdateGroup_PreservesDescriptionWhenNil(t *testing.T) {
+ existingGroup := &Group{
+ ID: 1,
+ Name: "existing-group",
+ Description: "keep me",
+ Platform: PlatformOpenAI,
+ Status: StatusActive,
+ }
+ repo := &groupRepoStubForAdmin{getByID: existingGroup}
+ svc := &adminServiceImpl{groupRepo: repo}
+
+ _, err := svc.UpdateGroup(context.Background(), 1, &UpdateGroupInput{
+ Description: nil,
+ })
+ require.NoError(t, err)
+ require.NotNil(t, repo.updated)
+ require.Equal(t, "keep me", repo.updated.Description, "nil should preserve existing description")
+}
+
func TestAdminService_UpdateGroup_RejectsNegativeImageRateMultiplier(t *testing.T) {
existingGroup := &Group{
ID: 1,
diff --git a/backend/internal/service/anthropic_session.go b/backend/internal/service/anthropic_session.go
index 26544c68..bca8cc7f 100644
--- a/backend/internal/service/anthropic_session.go
+++ b/backend/internal/service/anthropic_session.go
@@ -4,6 +4,8 @@ import (
"encoding/json"
"strings"
"time"
+
+ "github.com/tidwall/gjson"
)
// Anthropic 会话 Fallback 相关常量
@@ -30,30 +32,39 @@ func BuildAnthropicDigestChain(parsed *ParsedRequest) string {
var parts []string
- // 1. system prompt
- if parsed.System != nil {
- systemData, _ := json.Marshal(parsed.System)
- if len(systemData) > 0 && string(systemData) != "null" {
- parts = append(parts, "s:"+shortHash(systemData))
- }
+ if systemRaw := parsed.SystemRaw(); len(systemRaw) > 0 && string(systemRaw) != "null" {
+ parts = append(parts, "s:"+shortHash(canonicalAnthropicDigestJSON(systemRaw)))
}
- // 2. messages
- for _, msg := range parsed.Messages {
- msgMap, ok := msg.(map[string]any)
- if !ok {
- continue
- }
- role, _ := msgMap["role"].(string)
- prefix := rolePrefix(role)
- content := msgMap["content"]
- contentData, _ := json.Marshal(content)
- parts = append(parts, prefix+":"+shortHash(contentData))
+ messages := parsed.MessagesRaw()
+ if len(messages) > 0 {
+ gjson.ParseBytes(messages).ForEach(func(_, msg gjson.Result) bool {
+ prefix := rolePrefix(msg.Get("role").String())
+ content := msg.Get("content")
+ parts = append(parts, prefix+":"+shortHash(canonicalAnthropicDigestJSON([]byte(content.Raw))))
+ return true
+ })
}
return strings.Join(parts, "-")
}
+// canonicalAnthropicDigestJSON 保持 digest 对 JSON key 顺序和空白不敏感。
+func canonicalAnthropicDigestJSON(raw []byte) []byte {
+ if len(raw) == 0 {
+ return raw
+ }
+ var value any
+ if err := json.Unmarshal(raw, &value); err != nil {
+ return raw
+ }
+ canonical, err := json.Marshal(value)
+ if err != nil {
+ return raw
+ }
+ return canonical
+}
+
// rolePrefix 将 Anthropic 的 role 映射为单字符前缀
func rolePrefix(role string) string {
switch role {
diff --git a/backend/internal/service/anthropic_session_test.go b/backend/internal/service/anthropic_session_test.go
index 10406643..4d88fc8b 100644
--- a/backend/internal/service/anthropic_session_test.go
+++ b/backend/internal/service/anthropic_session_test.go
@@ -1,3 +1,5 @@
+//go:build unit
+
package service
import (
@@ -5,6 +7,15 @@ import (
"testing"
)
+func mustParseAnthropicDigestRequest(t *testing.T, body string) *ParsedRequest {
+ t.Helper()
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(body)), "")
+ if err != nil {
+ t.Fatalf("ParseGatewayRequest failed: %v", err)
+ }
+ return parsed
+}
+
func TestBuildAnthropicDigestChain_NilRequest(t *testing.T) {
result := BuildAnthropicDigestChain(nil)
if result != "" {
@@ -13,9 +24,7 @@ func TestBuildAnthropicDigestChain_NilRequest(t *testing.T) {
}
func TestBuildAnthropicDigestChain_EmptyMessages(t *testing.T) {
- parsed := &ParsedRequest{
- Messages: []any{},
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"messages":[]}`)
result := BuildAnthropicDigestChain(parsed)
if result != "" {
t.Errorf("expected empty string for empty messages, got: %s", result)
@@ -23,11 +32,7 @@ func TestBuildAnthropicDigestChain_EmptyMessages(t *testing.T) {
}
func TestBuildAnthropicDigestChain_SingleUserMessage(t *testing.T) {
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":"hello"}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 1 {
@@ -39,12 +44,7 @@ func TestBuildAnthropicDigestChain_SingleUserMessage(t *testing.T) {
}
func TestBuildAnthropicDigestChain_UserAndAssistant(t *testing.T) {
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi there"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"hi there"}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 2 {
@@ -59,12 +59,7 @@ func TestBuildAnthropicDigestChain_UserAndAssistant(t *testing.T) {
}
func TestBuildAnthropicDigestChain_WithSystemString(t *testing.T) {
- parsed := &ParsedRequest{
- System: "You are a helpful assistant",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"system":"You are a helpful assistant","messages":[{"role":"user","content":"hello"}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 2 {
@@ -79,14 +74,7 @@ func TestBuildAnthropicDigestChain_WithSystemString(t *testing.T) {
}
func TestBuildAnthropicDigestChain_WithSystemContentBlocks(t *testing.T) {
- parsed := &ParsedRequest{
- System: []any{
- map[string]any{"type": "text", "text": "You are a helpful assistant"},
- },
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"system":[{"type":"text","text":"You are a helpful assistant"}],"messages":[{"role":"user","content":"hello"}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 2 {
@@ -100,74 +88,33 @@ func TestBuildAnthropicDigestChain_WithSystemContentBlocks(t *testing.T) {
func TestBuildAnthropicDigestChain_ConversationPrefixRelationship(t *testing.T) {
// 核心测试:验证对话增长时链的前缀关系
// 上一轮的完整链一定是下一轮链的前缀
- system := "You are a helpful assistant"
-
- // 第 1 轮: system + user
- round1 := &ParsedRequest{
- System: system,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ round1 := mustParseAnthropicDigestRequest(t, `{"system":"You are a helpful assistant","messages":[{"role":"user","content":"hello"}]}`)
chain1 := BuildAnthropicDigestChain(round1)
- // 第 2 轮: system + user + assistant + user
- round2 := &ParsedRequest{
- System: system,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi there"},
- map[string]any{"role": "user", "content": "how are you?"},
- },
- }
+ round2 := mustParseAnthropicDigestRequest(t, `{"system":"You are a helpful assistant","messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"hi there"},{"role":"user","content":"how are you?"}]}`)
chain2 := BuildAnthropicDigestChain(round2)
- // 第 3 轮: system + user + assistant + user + assistant + user
- round3 := &ParsedRequest{
- System: system,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi there"},
- map[string]any{"role": "user", "content": "how are you?"},
- map[string]any{"role": "assistant", "content": "I'm doing well"},
- map[string]any{"role": "user", "content": "great"},
- },
- }
+ round3 := mustParseAnthropicDigestRequest(t, `{"system":"You are a helpful assistant","messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"hi there"},{"role":"user","content":"how are you?"},{"role":"assistant","content":"I'm doing well"},{"role":"user","content":"great"}]}`)
chain3 := BuildAnthropicDigestChain(round3)
t.Logf("Chain1: %s", chain1)
t.Logf("Chain2: %s", chain2)
t.Logf("Chain3: %s", chain3)
- // chain1 是 chain2 的前缀
if !strings.HasPrefix(chain2, chain1) {
t.Errorf("chain1 should be prefix of chain2:\n chain1: %s\n chain2: %s", chain1, chain2)
}
-
- // chain2 是 chain3 的前缀
if !strings.HasPrefix(chain3, chain2) {
t.Errorf("chain2 should be prefix of chain3:\n chain2: %s\n chain3: %s", chain2, chain3)
}
-
- // chain1 也是 chain3 的前缀(传递性)
if !strings.HasPrefix(chain3, chain1) {
t.Errorf("chain1 should be prefix of chain3:\n chain1: %s\n chain3: %s", chain1, chain3)
}
}
func TestBuildAnthropicDigestChain_DifferentSystemProducesDifferentChain(t *testing.T) {
- parsed1 := &ParsedRequest{
- System: "System A",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
- parsed2 := &ParsedRequest{
- System: "System B",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed1 := mustParseAnthropicDigestRequest(t, `{"system":"System A","messages":[{"role":"user","content":"hello"}]}`)
+ parsed2 := mustParseAnthropicDigestRequest(t, `{"system":"System B","messages":[{"role":"user","content":"hello"}]}`)
chain1 := BuildAnthropicDigestChain(parsed1)
chain2 := BuildAnthropicDigestChain(parsed2)
@@ -176,7 +123,6 @@ func TestBuildAnthropicDigestChain_DifferentSystemProducesDifferentChain(t *test
t.Error("Different system prompts should produce different chains")
}
- // 但 user 部分的 hash 应该相同
parts1 := splitChain(chain1)
parts2 := splitChain(chain2)
if parts1[1] != parts2[1] {
@@ -185,20 +131,8 @@ func TestBuildAnthropicDigestChain_DifferentSystemProducesDifferentChain(t *test
}
func TestBuildAnthropicDigestChain_DifferentContentProducesDifferentChain(t *testing.T) {
- parsed1 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "ORIGINAL reply"},
- map[string]any{"role": "user", "content": "next"},
- },
- }
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "TAMPERED reply"},
- map[string]any{"role": "user", "content": "next"},
- },
- }
+ parsed1 := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"ORIGINAL reply"},{"role":"user","content":"next"}]}`)
+ parsed2 := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"TAMPERED reply"},{"role":"user","content":"next"}]}`)
chain1 := BuildAnthropicDigestChain(parsed1)
chain2 := BuildAnthropicDigestChain(parsed2)
@@ -209,24 +143,16 @@ func TestBuildAnthropicDigestChain_DifferentContentProducesDifferentChain(t *tes
parts1 := splitChain(chain1)
parts2 := splitChain(chain2)
- // 第一个 user message hash 应该相同
if parts1[0] != parts2[0] {
t.Error("First user message hash should be the same")
}
- // assistant reply hash 应该不同
if parts1[1] == parts2[1] {
t.Error("Assistant reply hash should differ")
}
}
func TestBuildAnthropicDigestChain_Deterministic(t *testing.T) {
- parsed := &ParsedRequest{
- System: "test system",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi"},
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"system":"test system","messages":[{"role":"user","content":"hello"},{"role":"assistant","content":"hi"}]}`)
chain1 := BuildAnthropicDigestChain(parsed)
chain2 := BuildAnthropicDigestChain(parsed)
@@ -236,6 +162,18 @@ func TestBuildAnthropicDigestChain_Deterministic(t *testing.T) {
}
}
+func TestBuildAnthropicDigestChain_CanonicalJSON(t *testing.T) {
+ parsed1 := mustParseAnthropicDigestRequest(t, `{"system":[{"type":"text","text":"system"}],"messages":[{"role":"user","content":{"type":"text","text":"hello"}}]}`)
+ parsed2 := mustParseAnthropicDigestRequest(t, `{"system":[{"text":"system","type":"text"}],"messages":[{"role":"user","content":{"text":"hello","type":"text"}}]}`)
+
+ chain1 := BuildAnthropicDigestChain(parsed1)
+ chain2 := BuildAnthropicDigestChain(parsed2)
+
+ if chain1 != chain2 {
+ t.Errorf("semantically equivalent JSON should produce same chain: %s vs %s", chain1, chain2)
+ }
+}
+
func TestGenerateAnthropicDigestSessionKey(t *testing.T) {
tests := []struct {
name string
@@ -278,7 +216,6 @@ func TestGenerateAnthropicDigestSessionKey(t *testing.T) {
})
}
- // 验证不同 uuid 产生不同 sessionKey
t.Run("different uuid different key", func(t *testing.T) {
hash := "sameprefix123456"
result1 := GenerateAnthropicDigestSessionKey(hash, "uuid0001-session-a")
@@ -297,18 +234,7 @@ func TestAnthropicSessionTTL(t *testing.T) {
}
func TestBuildAnthropicDigestChain_ContentBlocks(t *testing.T) {
- // 测试 content 为 content blocks 数组的情况
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "content": []any{
- map[string]any{"type": "text", "text": "describe this image"},
- map[string]any{"type": "image", "source": map[string]any{"type": "base64"}},
- },
- },
- },
- }
+ parsed := mustParseAnthropicDigestRequest(t, `{"messages":[{"role":"user","content":[{"type":"text","text":"describe this image"},{"type":"image","source":{"type":"base64"}}]}]}`)
result := BuildAnthropicDigestChain(parsed)
parts := splitChain(result)
if len(parts) != 1 {
diff --git a/backend/internal/service/antigravity_default_test_stubs_test.go b/backend/internal/service/antigravity_default_test_stubs_test.go
new file mode 100644
index 00000000..d3c2c57a
--- /dev/null
+++ b/backend/internal/service/antigravity_default_test_stubs_test.go
@@ -0,0 +1,61 @@
+//go:build !unit
+
+package service
+
+import (
+ "context"
+ "time"
+)
+
+type defaultRateLimitCall struct {
+ accountID int64
+ resetAt time.Time
+}
+
+type defaultModelRateLimitCall struct {
+ accountID int64
+ modelKey string
+ resetAt time.Time
+}
+
+type defaultExtraUpdateCall struct {
+ accountID int64
+ updates map[string]any
+}
+
+type stubAntigravityAccountRepo struct {
+ AccountRepository
+ rateCalls []defaultRateLimitCall
+ modelRateLimitCalls []defaultModelRateLimitCall
+ extraUpdateCalls []defaultExtraUpdateCall
+}
+
+func (s *stubAntigravityAccountRepo) SetRateLimited(_ context.Context, id int64, resetAt time.Time) error {
+ s.rateCalls = append(s.rateCalls, defaultRateLimitCall{accountID: id, resetAt: resetAt})
+ return nil
+}
+
+func (s *stubAntigravityAccountRepo) SetModelRateLimit(_ context.Context, id int64, modelKey string, resetAt time.Time, _ ...string) error {
+ s.modelRateLimitCalls = append(s.modelRateLimitCalls, defaultModelRateLimitCall{accountID: id, modelKey: modelKey, resetAt: resetAt})
+ return nil
+}
+
+func (s *stubAntigravityAccountRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error {
+ s.extraUpdateCalls = append(s.extraUpdateCalls, defaultExtraUpdateCall{accountID: id, updates: updates})
+ return nil
+}
+
+type defaultDeleteSessionCall struct {
+ groupID int64
+ sessionHash string
+}
+
+type stubSmartRetryCache struct {
+ GatewayCache
+ deleteCalls []defaultDeleteSessionCall
+}
+
+func (c *stubSmartRetryCache) DeleteSessionAccountID(_ context.Context, groupID int64, sessionHash string) error {
+ c.deleteCalls = append(c.deleteCalls, defaultDeleteSessionCall{groupID: groupID, sessionHash: sessionHash})
+ return nil
+}
diff --git a/backend/internal/service/antigravity_gateway_service.go b/backend/internal/service/antigravity_gateway_service.go
index 5a90a195..f79d20a2 100644
--- a/backend/internal/service/antigravity_gateway_service.go
+++ b/backend/internal/service/antigravity_gateway_service.go
@@ -228,12 +228,11 @@ func (s *AntigravityGatewayService) handleSmartRetry(p antigravityRetryLoopParam
p.prefix, resp.StatusCode, modelName, p.account.ID, rateLimitDuration, truncateForLog(respBody, 200))
resetAt := time.Now().Add(rateLimitDuration)
- if !setModelRateLimitByModelName(p.ctx, p.accountRepo, p.account.ID, modelName, p.prefix, resp.StatusCode, resetAt, false) {
+ if !s.setAntigravityModelRateLimits(p.ctx, p.accountRepo, p.account, modelName, p.prefix, resp.StatusCode, resetAt, false) {
p.handleError(p.ctx, p.prefix, p.account, resp.StatusCode, resp.Header, respBody, p.requestedModel, p.groupID, p.sessionHash, p.isStickySession)
logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d rate_limited account=%d (no model mapping)", p.prefix, resp.StatusCode, p.account.ID)
- } else {
- s.updateAccountModelRateLimitInCache(p.ctx, p.account, modelName, resetAt)
}
+ s.clearStickySession(p.ctx, p.groupID, p.sessionHash)
// 返回账号切换信号,让上层切换账号重试
return &smartRetryResult{
@@ -392,20 +391,10 @@ func (s *AntigravityGatewayService) handleSmartRetry(p antigravityRetryLoopParam
p.prefix, resp.StatusCode, maxAttempts, modelName, p.account.ID, rateLimitDuration, truncateForLog(retryBody, 200))
resetAt := time.Now().Add(rateLimitDuration)
- if p.accountRepo != nil && modelName != "" {
- if err := p.accountRepo.SetModelRateLimit(p.ctx, p.account.ID, modelName, resetAt); err != nil {
- logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d model_rate_limit_failed model=%s error=%v", p.prefix, resp.StatusCode, modelName, err)
- } else {
- logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d model_rate_limited_after_smart_retry model=%s account=%d reset_in=%v",
- p.prefix, resp.StatusCode, modelName, p.account.ID, rateLimitDuration)
- s.updateAccountModelRateLimitInCache(p.ctx, p.account, modelName, resetAt)
- }
- }
+ s.setAntigravityModelRateLimits(p.ctx, p.accountRepo, p.account, modelName, p.prefix, resp.StatusCode, resetAt, true)
// 清除粘性会话绑定,避免下次请求仍命中限流账号
- if s.cache != nil && p.sessionHash != "" {
- _ = s.cache.DeleteSessionAccountID(p.ctx, p.groupID, p.sessionHash)
- }
+ s.clearStickySession(p.ctx, p.groupID, p.sessionHash)
// 返回账号切换信号,让上层切换账号重试
return &smartRetryResult{
@@ -662,7 +651,7 @@ urlFallbackLoop:
// 统一处理错误响应
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
if overagesInjected && shouldMarkCreditsExhausted(resp, respBody, nil) {
@@ -875,6 +864,22 @@ type AntigravityGatewayService struct {
internal500Cache Internal500CounterCache // INTERNAL 500 渐进惩罚计数器
}
+func (s *AntigravityGatewayService) upstreamErrorBodyReadLimit() int64 {
+ limit := gatewayUpstreamErrorBodyReadLimit
+ if s != nil && s.settingService != nil && s.settingService.cfg != nil && s.settingService.cfg.Gateway.LogUpstreamErrorBody && s.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes > int(limit) {
+ limit = int64(s.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes)
+ }
+ return limit
+}
+
+func (s *AntigravityGatewayService) readUpstreamErrorBody(resp *http.Response) []byte {
+ if resp == nil || resp.Body == nil {
+ return nil
+ }
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, s.upstreamErrorBodyReadLimit()))
+ return body
+}
+
func NewAntigravityGatewayService(
accountRepo AccountRepository,
cache GatewayCache,
@@ -938,8 +943,14 @@ func (s *AntigravityGatewayService) checkErrorPolicy(ctx context.Context, accoun
func (s *AntigravityGatewayService) applyErrorPolicy(p antigravityRetryLoopParams, statusCode int, headers http.Header, respBody []byte) (handled bool, outStatus int, retErr error) {
switch s.checkErrorPolicy(p.ctx, p.account, statusCode, respBody) {
case ErrorPolicySkipped:
+ if s.handleAntigravityModelRateLimitBeforePolicy(p, statusCode, headers, respBody) {
+ return true, statusCode, nil
+ }
return true, http.StatusInternalServerError, nil
case ErrorPolicyMatched:
+ if s.handleAntigravityModelRateLimitBeforePolicy(p, statusCode, headers, respBody) {
+ return true, statusCode, nil
+ }
_ = p.handleError(p.ctx, p.prefix, p.account, statusCode, headers, respBody,
p.requestedModel, p.groupID, p.sessionHash, p.isStickySession)
return true, statusCode, nil
@@ -951,6 +962,31 @@ func (s *AntigravityGatewayService) applyErrorPolicy(p antigravityRetryLoopParam
return false, statusCode, nil
}
+func (s *AntigravityGatewayService) handleAntigravityModelRateLimitBeforePolicy(p antigravityRetryLoopParams, statusCode int, headers http.Header, respBody []byte) bool {
+ if statusCode != http.StatusTooManyRequests && statusCode != http.StatusServiceUnavailable {
+ return false
+ }
+ if p.account == nil || p.account.Platform != PlatformAntigravity {
+ return false
+ }
+ _, shouldRateLimitModel, waitDuration, modelName, isModelCapacityExhausted := shouldTriggerAntigravitySmartRetry(p.account, respBody)
+ if isModelCapacityExhausted || !shouldRateLimitModel || strings.TrimSpace(modelName) == "" {
+ return false
+ }
+ rateLimitDuration := waitDuration
+ if rateLimitDuration <= 0 {
+ rateLimitDuration = antigravityDefaultRateLimitDuration
+ }
+ resetAt := time.Now().Add(rateLimitDuration)
+ if !s.setAntigravityModelRateLimits(p.ctx, p.accountRepo, p.account, modelName, p.prefix, statusCode, resetAt, false) {
+ return false
+ }
+ s.clearStickySession(p.ctx, p.groupID, p.sessionHash)
+ logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d model_rate_limited_before_error_policy model=%s account=%d reset_in=%v",
+ p.prefix, statusCode, modelName, p.account.ID, rateLimitDuration)
+ return true
+}
+
// mapAntigravityModel 获取映射后的模型名
// 完全依赖映射配置:账户映射(通配符)→ 默认映射兜底(DefaultAntigravityModelMapping)
// 注意:返回空字符串表示模型不被支持,调度时会过滤掉该账号
@@ -958,6 +994,7 @@ func mapAntigravityModel(account *Account, requestedModel string) string {
if account == nil {
return ""
}
+ requestedModel = strings.TrimPrefix(requestedModel, "models/")
// 获取映射表(未配置时自动使用 DefaultAntigravityModelMapping)
mapping := account.GetModelMapping()
@@ -1090,7 +1127,7 @@ func (s *AntigravityGatewayService) TestConnection(ctx context.Context, account
}
defer func() { _ = result.resp.Body.Close() }()
- respBody, err := io.ReadAll(io.LimitReader(result.resp.Body, 2<<20))
+ respBody, err := io.ReadAll(io.LimitReader(result.resp.Body, s.upstreamErrorBodyReadLimit()))
if err != nil {
return nil, fmt.Errorf("读取响应失败: %w", err)
}
@@ -1312,22 +1349,6 @@ func (s *AntigravityGatewayService) unwrapV1InternalResponse(body []byte) ([]byt
return body, nil
}
-// isModelNotFoundError 检测是否为模型不存在的 404 错误
-func isModelNotFoundError(statusCode int, body []byte) bool {
- if statusCode != 404 {
- return false
- }
-
- bodyStr := strings.ToLower(string(body))
- keywords := []string{"model not found", "unknown model", "not found"}
- for _, keyword := range keywords {
- if strings.Contains(bodyStr, keyword) {
- return true
- }
- }
- return true // 404 without specific message also treated as model not found
-}
-
// Forward 转发 Claude 协议请求(Claude → Gemini 转换)
//
// 限流处理流程:
@@ -1443,7 +1464,7 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context,
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
// 优先检测 thinking block 的 signature 相关错误(400)并重试一次:
// Antigravity /v1internal 链路在部分场景会对 thought/thinking signature 做严格校验,
@@ -1638,7 +1659,7 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context,
resp = retryResp
respBody = nil
} else {
- retryBody, _ := io.ReadAll(io.LimitReader(retryResp.Body, 2<<20))
+ retryBody := s.readUpstreamErrorBody(retryResp)
_ = retryResp.Body.Close()
respBody = retryBody
resp = &http.Response{
@@ -2073,8 +2094,28 @@ func stripSignatureSensitiveBlocksFromClaudeRequest(req *antigravity.ClaudeReque
// └─ retryDelay < 7s → 等待后重试 1 次
// ├─ 成功 → 正常返回
// └─ 失败 → 设置模型限流 + 清除粘性绑定 → 切换账号
-func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Context, account *Account, originalModel string, action string, stream bool, body []byte, isStickySession bool) (*ForwardResult, error) {
+type ForwardGeminiOption func(*forwardGeminiOptions)
+
+type forwardGeminiOptions struct {
+ groupID int64
+ sessionHash string
+}
+
+func WithForwardGeminiSession(groupID int64, sessionHash string) ForwardGeminiOption {
+ return func(opts *forwardGeminiOptions) {
+ opts.groupID = groupID
+ opts.sessionHash = sessionHash
+ }
+}
+
+func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Context, account *Account, originalModel string, action string, stream bool, body []byte, isStickySession bool, options ...ForwardGeminiOption) (*ForwardResult, error) {
startTime := time.Now()
+ forwardOpts := forwardGeminiOptions{}
+ for _, apply := range options {
+ if apply != nil {
+ apply(&forwardOpts)
+ }
+ }
sessionID := getSessionID(c)
prefix := logPrefix(sessionID, account.Name)
@@ -2179,8 +2220,8 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
handleError: s.handleUpstreamError,
requestedModel: originalModel,
isStickySession: isStickySession, // ForwardGemini 由上层判断粘性会话
- groupID: 0, // ForwardGemini 方法没有 groupID,由上层处理粘性会话清除
- sessionHash: "", // ForwardGemini 方法没有 sessionHash,由上层处理粘性会话清除
+ groupID: forwardOpts.groupID,
+ sessionHash: forwardOpts.sessionHash,
})
if err != nil {
// 检查是否是账号切换信号,转换为 UpstreamFailoverError 让 Handler 切换账号
@@ -2205,7 +2246,7 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
// 处理错误响应
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
contentType := resp.Header.Get("Content-Type")
// 尽早关闭原始响应体,释放连接;后续逻辑仍可能需要读取 body,因此用内存副本重新包装。
_ = resp.Body.Close()
@@ -2278,15 +2319,15 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
handleError: s.handleUpstreamError,
requestedModel: originalModel,
isStickySession: isStickySession,
- groupID: 0,
- sessionHash: "",
+ groupID: forwardOpts.groupID,
+ sessionHash: forwardOpts.sessionHash,
})
if retryErr == nil {
retryResp := retryResult.resp
if retryResp.StatusCode < 400 {
resp = retryResp
} else {
- retryRespBody, _ := io.ReadAll(io.LimitReader(retryResp.Body, 2<<20))
+ retryRespBody := s.readUpstreamErrorBody(retryResp)
_ = retryResp.Body.Close()
retryOpsBody := retryRespBody
if retryUnwrapped, unwrapErr := s.unwrapV1InternalResponse(retryRespBody); unwrapErr == nil && len(retryUnwrapped) > 0 {
@@ -2355,7 +2396,7 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
if unwrapErr != nil || len(unwrappedForOps) == 0 {
unwrappedForOps = respBody
}
- s.handleUpstreamError(ctx, prefix, account, resp.StatusCode, resp.Header, respBody, originalModel, 0, "", isStickySession)
+ s.handleUpstreamError(ctx, prefix, account, resp.StatusCode, resp.Header, respBody, originalModel, forwardOpts.groupID, forwardOpts.sessionHash, isStickySession)
upstreamMsg := strings.TrimSpace(extractAntigravityErrorMessage(unwrappedForOps))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
upstreamDetail := s.getUpstreamErrorDetail(unwrappedForOps)
@@ -2566,6 +2607,34 @@ func setModelRateLimitByModelName(ctx context.Context, repo AccountRepository, a
return true
}
+func (s *AntigravityGatewayService) setAntigravityModelRateLimits(ctx context.Context, repo AccountRepository, account *Account, modelName, prefix string, statusCode int, resetAt time.Time, afterSmartRetry bool) bool {
+ if account == nil || repo == nil {
+ return false
+ }
+ keys := antigravityModelRateLimitKeys(modelName)
+ if len(keys) == 0 {
+ return false
+ }
+
+ success := false
+ for _, key := range keys {
+ if setModelRateLimitByModelName(ctx, repo, account.ID, key, prefix, statusCode, resetAt, afterSmartRetry) {
+ s.updateAccountModelRateLimitInCache(ctx, account, key, resetAt)
+ success = true
+ }
+ }
+ return success
+}
+
+func (s *AntigravityGatewayService) clearStickySession(ctx context.Context, groupID int64, sessionHash string) {
+ if s == nil || s.cache == nil || strings.TrimSpace(sessionHash) == "" {
+ return
+ }
+ if err := s.cache.DeleteSessionAccountID(ctx, groupID, sessionHash); err != nil {
+ logger.LegacyPrintf("service.antigravity_gateway", "[antigravity-Forward] sticky_session_clear_failed group_id=%d session=%s err=%v", groupID, shortSessionHash(sessionHash), err)
+ }
+}
+
func antigravityFallbackCooldownSeconds() (time.Duration, bool) {
raw := strings.TrimSpace(os.Getenv(antigravityFallbackSecondsEnv))
if raw == "" {
@@ -2644,7 +2713,7 @@ func parseAntigravitySmartRetryInfo(body []byte) *antigravitySmartRetryInfo {
if atType == googleRPCTypeErrorInfo {
if meta, ok := dm["metadata"].(map[string]any); ok {
if model, ok := meta["model"].(string); ok {
- modelName = model
+ modelName = normalizeAntigravityModelName(model)
}
}
// 检查 reason
@@ -2818,13 +2887,7 @@ func (s *AntigravityGatewayService) setModelRateLimitAndClearSession(p *handleMo
logger.LegacyPrintf("service.antigravity_gateway", "%s status=%d model_rate_limited model=%s account=%d reset_in=%v",
p.prefix, p.statusCode, info.ModelName, p.account.ID, info.RetryDelay)
- // 设置模型限流状态(数据库)
- if err := s.accountRepo.SetModelRateLimit(p.ctx, p.account.ID, info.ModelName, resetAt); err != nil {
- logger.LegacyPrintf("service.antigravity_gateway", "%s model_rate_limit_failed model=%s error=%v", p.prefix, info.ModelName, err)
- }
-
- // 立即更新 Redis 快照中账号的限流状态,避免并发请求重复选中
- s.updateAccountModelRateLimitInCache(p.ctx, p.account, info.ModelName, resetAt)
+ s.setAntigravityModelRateLimits(p.ctx, s.accountRepo, p.account, info.ModelName, p.prefix, p.statusCode, resetAt, false)
// 清除粘性会话绑定
if p.cache != nil && p.sessionHash != "" {
@@ -2914,12 +2977,11 @@ func (s *AntigravityGatewayService) handleUpstreamError(
}
if modelKey != "" {
ra := s.resolveResetTime(resetAt, defaultDur)
- if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, modelKey, ra); err != nil {
- logger.LegacyPrintf("service.antigravity_gateway", "%s status=429 model_rate_limit_set_failed model=%s error=%v", prefix, modelKey, err)
+ if !s.setAntigravityModelRateLimits(ctx, s.accountRepo, account, modelKey, prefix, statusCode, ra, false) {
+ logger.LegacyPrintf("service.antigravity_gateway", "%s status=429 model_rate_limit_set_failed model=%s", prefix, modelKey)
} else {
logger.LegacyPrintf("service.antigravity_gateway", "%s status=429 model_rate_limited model=%s account=%d reset_at=%v reset_in=%v",
prefix, modelKey, account.ID, ra.Format("15:04:05"), time.Until(ra).Truncate(time.Second))
- s.updateAccountModelRateLimitInCache(ctx, account, modelKey, ra)
}
return nil
}
@@ -4225,6 +4287,14 @@ func (s *AntigravityGatewayService) ForwardUpstream(ctx context.Context, c *gin.
// 构建上游请求 URL
upstreamURL := baseURL + "/v1/messages"
+ // 能力维度 sanitize:Anthropic-compatible 上游透传路径也需要保证 body↔beta header
+ // 对称。客户端 anthropic-beta header 不含 context-management-2025-06-27 但 body 带
+ // context_management 时 strip,与 Anthropic 直连 / Bedrock / Vertex 路径保持一致。
+ clientBeta := c.GetHeader("anthropic-beta")
+ if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed {
+ body = sanitized
+ }
+
// 创建请求
req, err := http.NewRequestWithContext(ctx, http.MethodPost, upstreamURL, bytes.NewReader(body))
if err != nil {
@@ -4240,7 +4310,7 @@ func (s *AntigravityGatewayService) ForwardUpstream(ctx context.Context, c *gin.
if v := c.GetHeader("anthropic-version"); v != "" {
req.Header.Set("anthropic-version", v)
}
- if v := c.GetHeader("anthropic-beta"); v != "" {
+ if v := clientBeta; v != "" {
req.Header.Set("anthropic-beta", v)
}
@@ -4260,7 +4330,7 @@ func (s *AntigravityGatewayService) ForwardUpstream(ctx context.Context, c *gin.
// 处理错误响应
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
// 429 错误时标记账号限流
if resp.StatusCode == http.StatusTooManyRequests {
@@ -4463,6 +4533,14 @@ func (s *AntigravityGatewayService) streamUpstreamResponse(c *gin.Context, resp
}
// extractSSEUsage 从 SSE data 行中提取 Claude usage(用于流式透传场景)
+//
+// Anthropic streaming 的 usage 字段分布在两类事件中:
+// - message_start:嵌套在 event.message.usage(input_tokens、cache_creation_input_tokens、
+// cache_read_input_tokens 等输入侧字段)
+// - message_delta:位于顶层 event.usage(流结束时的最终 output_tokens)
+//
+// 仅读取顶层 event.usage 会漏掉 message_start 的输入侧字段,导致流式透传请求落库的
+// usage_logs 记录 input_tokens=0。
func (s *AntigravityGatewayService) extractSSEUsage(line string, usage *ClaudeUsage) {
if !strings.HasPrefix(line, "data: ") {
return
@@ -4472,8 +4550,15 @@ func (s *AntigravityGatewayService) extractSSEUsage(line string, usage *ClaudeUs
if json.Unmarshal([]byte(dataStr), &event) != nil {
return
}
- u, ok := event["usage"].(map[string]any)
- if !ok {
+ var u map[string]any
+ if eventType, _ := event["type"].(string); eventType == "message_start" {
+ if msg, ok := event["message"].(map[string]any); ok {
+ u, _ = msg["usage"].(map[string]any)
+ }
+ } else {
+ u, _ = event["usage"].(map[string]any)
+ }
+ if u == nil {
return
}
if v, ok := u["input_tokens"].(float64); ok && int(v) > 0 {
diff --git a/backend/internal/service/antigravity_gateway_service_test.go b/backend/internal/service/antigravity_gateway_service_test.go
index 1eb1451e..0fac7a1e 100644
--- a/backend/internal/service/antigravity_gateway_service_test.go
+++ b/backend/internal/service/antigravity_gateway_service_test.go
@@ -42,6 +42,15 @@ func newAntigravityTestService(cfg *config.Config) *AntigravityGatewayService {
}
}
+func TestAntigravityUpstreamErrorBodyReadLimit_RespectsDiagnosticLimit(t *testing.T) {
+ svc := newAntigravityTestService(&config.Config{Gateway: config.GatewayConfig{
+ LogUpstreamErrorBody: true,
+ LogUpstreamErrorBodyMaxBytes: int(gatewayUpstreamErrorBodyReadLimit) + 1024,
+ }})
+
+ require.Equal(t, int64(svc.settingService.cfg.Gateway.LogUpstreamErrorBodyMaxBytes), svc.upstreamErrorBodyReadLimit())
+}
+
func TestStripSignatureSensitiveBlocksFromClaudeRequest(t *testing.T) {
req := &antigravity.ClaudeRequest{
Model: "claude-sonnet-4-5",
@@ -491,6 +500,86 @@ func TestAntigravityGatewayService_ForwardGemini_StickySessionForceCacheBilling(
require.True(t, failoverErr.ForceCacheBilling, "ForceCacheBilling should be true for sticky session switch")
}
+func TestAntigravityGatewayService_ForwardGemini_ClearsStickySessionOnGeminiRateLimit(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ writer := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(writer)
+
+ body, err := json.Marshal(map[string]any{
+ "contents": []map[string]any{
+ {"role": "user", "parts": []map[string]any{{"text": "hi"}}},
+ },
+ })
+ require.NoError(t, err)
+
+ req := httptest.NewRequest(http.MethodPost, "/v1beta/models/gemini-3-flash-preview:generateContent", bytes.NewReader(body))
+ c.Request = req
+
+ respBody := []byte(`{
+ "error": {
+ "status": "RESOURCE_EXHAUSTED",
+ "details": [
+ {"@type": "type.googleapis.com/google.rpc.ErrorInfo", "metadata": {"model": "gemini-3-flash"}, "reason": "RATE_LIMIT_EXCEEDED"},
+ {"@type": "type.googleapis.com/google.rpc.RetryInfo", "retryDelay": "15s"}
+ ]
+ }
+ }`)
+ upstream := &httpUpstreamStub{resp: &http.Response{
+ StatusCode: http.StatusTooManyRequests,
+ Header: http.Header{},
+ Body: io.NopCloser(bytes.NewReader(respBody)),
+ }}
+ repo := &stubAntigravityAccountRepo{}
+ cache := &stubSmartRetryCache{}
+ svc := &AntigravityGatewayService{
+ tokenProvider: &AntigravityTokenProvider{},
+ httpUpstream: upstream,
+ accountRepo: repo,
+ cache: cache,
+ }
+
+ account := &Account{
+ ID: 44,
+ Name: "acc-gemini-runtime-rate-limited",
+ Platform: PlatformAntigravity,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "token",
+ "expires_at": time.Now().Add(time.Hour).Format(time.RFC3339),
+ "project_id": "proj",
+ },
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ },
+ }
+
+ result, err := svc.ForwardGemini(
+ context.Background(),
+ c,
+ account,
+ "gemini-3-flash-preview",
+ "generateContent",
+ false,
+ body,
+ true,
+ WithForwardGeminiSession(77, "gemini:sticky-runtime"),
+ )
+
+ require.Nil(t, result)
+ var failoverErr *UpstreamFailoverError
+ require.ErrorAs(t, err, &failoverErr)
+ require.Equal(t, http.StatusServiceUnavailable, failoverErr.StatusCode)
+ require.Len(t, repo.modelRateLimitCalls, 2)
+ require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
+ require.Len(t, cache.deleteCalls, 1)
+ require.Equal(t, int64(77), cache.deleteCalls[0].groupID)
+ require.Equal(t, "gemini:sticky-runtime", cache.deleteCalls[0].sessionHash)
+}
+
// TestAntigravityGatewayService_Forward_BillsWithMappedModel
// 验证:Antigravity Claude 转发返回的计费模型使用映射后的模型
func TestAntigravityGatewayService_Forward_BillsWithMappedModel(t *testing.T) {
@@ -1301,6 +1390,19 @@ func TestExtractSSEUsage(t *testing.T) {
line: `data: {"usage":{"input_tokens":10,"output_tokens":20,"cache_read_input_tokens":5,"cache_creation_input_tokens":3}}`,
expected: ClaudeUsage{InputTokens: 10, OutputTokens: 20, CacheReadInputTokens: 5, CacheCreationInputTokens: 3},
},
+ {
+ // Anthropic message_start 把 usage 嵌套在 message.usage 下,
+ // 必须从这里提取输入侧字段(含 cache_read/cache_creation_input_tokens)。
+ name: "message_start nested usage with input/cache tokens",
+ line: `data: {"type":"message_start","message":{"id":"msg_01","usage":{"input_tokens":35576,"cache_creation_input_tokens":0,"cache_read_input_tokens":12000,"output_tokens":1}}}`,
+ expected: ClaudeUsage{InputTokens: 35576, OutputTokens: 1, CacheReadInputTokens: 12000},
+ },
+ {
+ // message_start.message.usage.cache_creation 内的 5m/1h 明细也要解析。
+ name: "message_start nested usage with cache_creation breakdown",
+ line: `data: {"type":"message_start","message":{"usage":{"input_tokens":100,"cache_creation":{"ephemeral_5m_input_tokens":30,"ephemeral_1h_input_tokens":70}}}}`,
+ expected: ClaudeUsage{InputTokens: 100, CacheCreation5mTokens: 30, CacheCreation1hTokens: 70},
+ },
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
@@ -1311,6 +1413,29 @@ func TestExtractSSEUsage(t *testing.T) {
}
}
+// TestExtractSSEUsage_StreamingSequence 复现 issue #2332:完整的 Anthropic streaming
+// 序列(message_start → message_delta)必须把两类事件中的 usage 字段都汇入同一份累计值,
+// 否则透传账号产出的 usage_logs 会出现 input_tokens=0、仅有 output_tokens 的"残缺"记录。
+func TestExtractSSEUsage_StreamingSequence(t *testing.T) {
+ svc := &AntigravityGatewayService{}
+ usage := &ClaudeUsage{}
+
+ // 1) message_start:携带完整输入侧 usage(input_tokens + cache_read)
+ svc.extractSSEUsage(
+ `data: {"type":"message_start","message":{"id":"msg_01","type":"message","role":"assistant","content":[],"model":"claude-opus-4-6","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":35576,"cache_creation_input_tokens":0,"cache_read_input_tokens":12000,"output_tokens":1}}}`,
+ usage,
+ )
+ // 2) message_delta:流结束时只带 output_tokens(无 input_tokens 字段)
+ svc.extractSSEUsage(
+ `data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"output_tokens":816}}`,
+ usage,
+ )
+
+ require.Equal(t, 35576, usage.InputTokens, "message_start 的 input_tokens 必须被记录,否则记账会缺失输入侧 token (#2332)")
+ require.Equal(t, 12000, usage.CacheReadInputTokens, "message_start 的 cache_read_input_tokens 必须被记录")
+ require.Equal(t, 816, usage.OutputTokens, "message_delta 的最终 output_tokens 必须被记录")
+}
+
// TestAntigravityClientWriter 验证 antigravityClientWriter 的断开检测
func TestAntigravityClientWriter(t *testing.T) {
t.Run("normal write succeeds", func(t *testing.T) {
diff --git a/backend/internal/service/antigravity_model_mapping_test.go b/backend/internal/service/antigravity_model_mapping_test.go
index a29000e7..652b6d66 100644
--- a/backend/internal/service/antigravity_model_mapping_test.go
+++ b/backend/internal/service/antigravity_model_mapping_test.go
@@ -88,6 +88,18 @@ func TestAntigravityGatewayService_GetMappedModel(t *testing.T) {
accountMapping: nil,
expected: "claude-sonnet-4-5",
},
+ {
+ name: "默认映射透传 - claude-opus-4-8",
+ requestedModel: "claude-opus-4-8",
+ accountMapping: nil,
+ expected: "claude-opus-4-8",
+ },
+ {
+ name: "默认映射透传 - claude-opus-4-7",
+ requestedModel: "claude-opus-4-7",
+ accountMapping: nil,
+ expected: "claude-opus-4-7",
+ },
{
name: "默认映射透传 - claude-opus-4-6-thinking",
requestedModel: "claude-opus-4-6-thinking",
@@ -210,6 +222,7 @@ func TestAntigravityGatewayService_IsModelSupported(t *testing.T) {
{"直接支持 - gemini-3-flash", "gemini-3-flash", true},
// 可映射(有明确前缀映射)
+ {"可映射 - claude-opus-4-8", "claude-opus-4-8", true},
{"可映射 - claude-opus-4-6", "claude-opus-4-6", true},
// 前缀透传(claude 和 gemini 前缀)
diff --git a/backend/internal/service/antigravity_quota_scope.go b/backend/internal/service/antigravity_quota_scope.go
index b536d16c..75862633 100644
--- a/backend/internal/service/antigravity_quota_scope.go
+++ b/backend/internal/service/antigravity_quota_scope.go
@@ -8,7 +8,17 @@ import (
func normalizeAntigravityModelName(model string) string {
normalized := strings.ToLower(strings.TrimSpace(model))
- normalized = strings.TrimPrefix(normalized, "models/")
+ if idx := strings.LastIndex(normalized, "/publishers/google/models/"); idx != -1 {
+ normalized = normalized[idx+len("/publishers/google/models/"):]
+ } else if idx := strings.LastIndex(normalized, "/publishers/anthropic/models/"); idx != -1 {
+ normalized = normalized[idx+len("/publishers/anthropic/models/"):]
+ } else if idx := strings.LastIndex(normalized, "/models/"); idx != -1 {
+ normalized = normalized[idx+len("/models/"):]
+ } else {
+ normalized = strings.TrimPrefix(normalized, "publishers/google/models/")
+ normalized = strings.TrimPrefix(normalized, "publishers/anthropic/models/")
+ normalized = strings.TrimPrefix(normalized, "models/")
+ }
return normalized
}
diff --git a/backend/internal/service/antigravity_rate_limit_test.go b/backend/internal/service/antigravity_rate_limit_test.go
index 35e130dc..374b29f6 100644
--- a/backend/internal/service/antigravity_rate_limit_test.go
+++ b/backend/internal/service/antigravity_rate_limit_test.go
@@ -94,7 +94,7 @@ func (s *stubAntigravityAccountRepo) SetRateLimited(ctx context.Context, id int6
return nil
}
-func (s *stubAntigravityAccountRepo) SetModelRateLimit(ctx context.Context, id int64, modelKey string, resetAt time.Time) error {
+func (s *stubAntigravityAccountRepo) SetModelRateLimit(ctx context.Context, id int64, modelKey string, resetAt time.Time, reason ...string) error {
s.modelRateLimitCalls = append(s.modelRateLimitCalls, modelRateLimitCall{accountID: id, modelKey: modelKey, resetAt: resetAt})
return nil
}
@@ -821,6 +821,51 @@ func TestSetModelRateLimitByModelName_NotConvertToScope(t *testing.T) {
require.NotEqual(t, "claude_sonnet", call.modelKey, "should NOT be scope")
}
+func TestSetAntigravityModelRateLimits_GeminiWritesFamilyScope(t *testing.T) {
+ repo := &stubAntigravityAccountRepo{}
+ svc := &AntigravityGatewayService{}
+ account := &Account{ID: 789, Platform: PlatformAntigravity}
+ resetAt := time.Now().Add(30 * time.Second)
+
+ success := svc.setAntigravityModelRateLimits(
+ context.Background(),
+ repo,
+ account,
+ "gemini-3-pro",
+ "[test]",
+ 429,
+ resetAt,
+ false,
+ )
+
+ require.True(t, success)
+ require.Len(t, repo.modelRateLimitCalls, 2)
+ require.Equal(t, "gemini-3-pro", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
+}
+
+func TestSetAntigravityModelRateLimits_ClaudeDoesNotWriteGeminiScope(t *testing.T) {
+ repo := &stubAntigravityAccountRepo{}
+ svc := &AntigravityGatewayService{}
+ account := &Account{ID: 790, Platform: PlatformAntigravity}
+ resetAt := time.Now().Add(30 * time.Second)
+
+ success := svc.setAntigravityModelRateLimits(
+ context.Background(),
+ repo,
+ account,
+ "claude-sonnet-4-5",
+ "[test]",
+ 429,
+ resetAt,
+ false,
+ )
+
+ require.True(t, success)
+ require.Len(t, repo.modelRateLimitCalls, 1)
+ require.Equal(t, "claude-sonnet-4-5", repo.modelRateLimitCalls[0].modelKey)
+}
+
func TestAntigravityRetryLoop_PreCheck_SwitchesWhenRateLimited(t *testing.T) {
upstream := &recordingOKUpstream{}
account := &Account{
@@ -1124,3 +1169,53 @@ func TestSchedulerSnapshotService_UpdateAccountInCache(t *testing.T) {
require.ErrorIs(t, err, expectedErr)
})
}
+func TestNormalizeAntigravityModelName(t *testing.T) {
+ tests := []struct {
+ name string
+ model string
+ expected string
+ }{
+ {
+ name: "plain model name",
+ model: "gemini-1.5-pro",
+ expected: "gemini-1.5-pro",
+ },
+ {
+ name: "models/ prefix",
+ model: "models/gemini-1.5-pro",
+ expected: "gemini-1.5-pro",
+ },
+ {
+ name: "publishers/google/models/ prefix",
+ model: "publishers/google/models/gemini-1.5-pro",
+ expected: "gemini-1.5-pro",
+ },
+ {
+ name: "projects/.../publishers/google/models/ path",
+ model: "projects/my-proj/locations/us-central1/publishers/google/models/gemini-2.5-flash",
+ expected: "gemini-2.5-flash",
+ },
+ {
+ name: "publishers/anthropic/models/ prefix",
+ model: "publishers/anthropic/models/claude-sonnet-4-5",
+ expected: "claude-sonnet-4-5",
+ },
+ {
+ name: "projects/.../publishers/anthropic/models/ path",
+ model: "projects/my-proj/locations/global/publishers/anthropic/models/claude-sonnet-4-5",
+ expected: "claude-sonnet-4-5",
+ },
+ {
+ name: "mixed case and spaces",
+ model: " Models/Gemini-1.5-Pro ",
+ expected: "gemini-1.5-pro",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ actual := normalizeAntigravityModelName(tt.model)
+ require.Equal(t, tt.expected, actual)
+ })
+ }
+}
diff --git a/backend/internal/service/antigravity_single_account_retry_test.go b/backend/internal/service/antigravity_single_account_retry_test.go
index 675e9c0c..6e58ab75 100644
--- a/backend/internal/service/antigravity_single_account_retry_test.go
+++ b/backend/internal/service/antigravity_single_account_retry_test.go
@@ -196,8 +196,10 @@ func TestHandleSmartRetry_503_LongDelay_NoSingleAccountRetry_StillSwitches(t *te
require.Nil(t, result.resp, "should not return resp when switchError is set")
// 对照:多账号模式应设模型限流
- require.Len(t, repo.modelRateLimitCalls, 1,
+ require.Len(t, repo.modelRateLimitCalls, 2,
"multi-account mode SHOULD set model rate limit")
+ require.Equal(t, "gemini-3-pro-high", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
}
// TestHandleSmartRetry_429_LongDelay_SingleAccountRetry_StillSwitches
@@ -412,8 +414,10 @@ func TestHandleSmartRetry_503_ShortDelay_NoSingleAccountRetry_SetsRateLimit(t *t
// 对照:多账号模式应返回 switchError
require.NotNil(t, result.switchError, "multi-account mode should return switchError for 503")
// 对照:多账号模式应设模型限流
- require.Len(t, repo.modelRateLimitCalls, 1,
+ require.Len(t, repo.modelRateLimitCalls, 2,
"multi-account mode should set model rate limit")
+ require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
}
// ---------------------------------------------------------------------------
diff --git a/backend/internal/service/antigravity_smart_retry_test.go b/backend/internal/service/antigravity_smart_retry_test.go
index e3b60a27..9f06b13b 100644
--- a/backend/internal/service/antigravity_smart_retry_test.go
+++ b/backend/internal/service/antigravity_smart_retry_test.go
@@ -328,9 +328,10 @@ func TestHandleSmartRetry_ShortDelay_SmartRetryFailed_ReturnsSwitchError(t *test
require.Equal(t, "gemini-3-flash", result.switchError.RateLimitedModel)
require.False(t, result.switchError.IsStickySession)
- // 验证模型限流已设置
- require.Len(t, repo.modelRateLimitCalls, 1)
+ // 验证模型限流已设置:Gemini 同时写入精确模型和家族级 scope
+ require.Len(t, repo.modelRateLimitCalls, 2)
require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
require.Len(t, upstream.calls, 1, "should have made one retry call (max attempts)")
}
@@ -1104,10 +1105,9 @@ func TestHandleSmartRetry_ShortDelay_StickySession_SuccessRetry_NoDeleteSession(
require.Len(t, cache.deleteCalls, 0, "should NOT call DeleteSessionAccountID on successful retry")
}
-// TestHandleSmartRetry_LongDelay_StickySession_NoDeleteInHandleSmartRetry
-// 长延迟路径(情况1)在 handleSmartRetry 中不直接调用 DeleteSessionAccountID
-// (清除由 handler 层的 shouldClearStickySession 在下次请求时处理)
-func TestHandleSmartRetry_LongDelay_StickySession_NoDeleteInHandleSmartRetry(t *testing.T) {
+// TestHandleSmartRetry_LongDelay_StickySession_ClearsSession
+// 长延迟路径(情况1)应立即清除 sticky 绑定,避免下一次请求继续命中已限流账号。
+func TestHandleSmartRetry_LongDelay_StickySession_ClearsSession(t *testing.T) {
repo := &stubAntigravityAccountRepo{}
cache := &stubSmartRetryCache{}
account := &Account{
@@ -1159,10 +1159,9 @@ func TestHandleSmartRetry_LongDelay_StickySession_NoDeleteInHandleSmartRetry(t *
require.NotNil(t, result.switchError)
require.True(t, result.switchError.IsStickySession)
- // 长延迟路径不在 handleSmartRetry 中调用 DeleteSessionAccountID
- // (由上游 handler 的 shouldClearStickySession 处理)
- require.Len(t, cache.deleteCalls, 0,
- "long delay path should NOT call DeleteSessionAccountID in handleSmartRetry (handled by handler layer)")
+ require.Len(t, cache.deleteCalls, 1, "long delay path should clear sticky session in handleSmartRetry")
+ require.Equal(t, int64(42), cache.deleteCalls[0].groupID)
+ require.Equal(t, "sticky-hash-long-delay", cache.deleteCalls[0].sessionHash)
}
// TestHandleSmartRetry_ShortDelay_NetworkError_StickySession_ClearsSession
@@ -1227,6 +1226,10 @@ func TestHandleSmartRetry_ShortDelay_NetworkError_StickySession_ClearsSession(t
require.Len(t, cache.deleteCalls, 1, "should call DeleteSessionAccountID after network error exhausts retry")
require.Equal(t, int64(99), cache.deleteCalls[0].groupID)
require.Equal(t, "sticky-net-error", cache.deleteCalls[0].sessionHash)
+
+ require.Len(t, repo.modelRateLimitCalls, 2)
+ require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
}
// TestHandleSmartRetry_ShortDelay_503_StickySession_FailedRetry_ClearsSession
@@ -1308,9 +1311,10 @@ func TestHandleSmartRetry_ShortDelay_503_StickySession_FailedRetry_ClearsSession
require.Equal(t, int64(77), cache.deleteCalls[0].groupID)
require.Equal(t, "sticky-503-short", cache.deleteCalls[0].sessionHash)
- // 验证模型限流已设置
- require.Len(t, repo.modelRateLimitCalls, 1)
+ // 验证模型限流已设置:Gemini 同时写入精确模型和家族级 scope
+ require.Len(t, repo.modelRateLimitCalls, 2)
require.Equal(t, "gemini-3-pro", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
}
// TestAntigravityRetryLoop_SmartRetryFailed_StickySession_SwitchErrorPropagates
diff --git a/backend/internal/service/api_key_auth_cache.go b/backend/internal/service/api_key_auth_cache.go
index 3553a18a..74163179 100644
--- a/backend/internal/service/api_key_auth_cache.go
+++ b/backend/internal/service/api_key_auth_cache.go
@@ -87,6 +87,7 @@ type APIKeyAuthGroupSnapshot struct {
AllowMessagesDispatch bool `json:"allow_messages_dispatch"`
DefaultMappedModel string `json:"default_mapped_model,omitempty"`
MessagesDispatchModelConfig OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config,omitempty"`
+ ModelsListConfig GroupModelsListConfig `json:"models_list_config,omitempty"`
// RPMLimit 分组级每分钟请求数上限(0 = 不限制);用于 billing_cache_service.checkRPM 级联判断。
RPMLimit int `json:"rpm_limit"`
diff --git a/backend/internal/service/api_key_auth_cache_impl.go b/backend/internal/service/api_key_auth_cache_impl.go
index c752ce28..69c6086f 100644
--- a/backend/internal/service/api_key_auth_cache_impl.go
+++ b/backend/internal/service/api_key_auth_cache_impl.go
@@ -14,7 +14,7 @@ import (
"github.com/dgraph-io/ristretto"
)
-const apiKeyAuthSnapshotVersion = 10 // v10: reload snapshots for group availability checks
+const apiKeyAuthSnapshotVersion = 11 // v11: reload snapshots for custom models_list_config
type apiKeyAuthCacheConfig struct {
l1Size int
@@ -272,6 +272,7 @@ func (s *APIKeyService) snapshotFromAPIKey(ctx context.Context, apiKey *APIKey)
AllowMessagesDispatch: apiKey.Group.AllowMessagesDispatch,
DefaultMappedModel: apiKey.Group.DefaultMappedModel,
MessagesDispatchModelConfig: apiKey.Group.MessagesDispatchModelConfig,
+ ModelsListConfig: apiKey.Group.ModelsListConfig,
RPMLimit: apiKey.Group.RPMLimit,
}
}
@@ -342,6 +343,7 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho
AllowMessagesDispatch: snapshot.Group.AllowMessagesDispatch,
DefaultMappedModel: snapshot.Group.DefaultMappedModel,
MessagesDispatchModelConfig: snapshot.Group.MessagesDispatchModelConfig,
+ ModelsListConfig: snapshot.Group.ModelsListConfig,
RPMLimit: snapshot.Group.RPMLimit,
}
}
diff --git a/backend/internal/service/api_key_auth_cache_version_test.go b/backend/internal/service/api_key_auth_cache_version_test.go
new file mode 100644
index 00000000..5982e526
--- /dev/null
+++ b/backend/internal/service/api_key_auth_cache_version_test.go
@@ -0,0 +1,43 @@
+package service
+
+import "testing"
+
+func TestAPIKeyService_RejectsV10AuthSnapshotWithoutModelsListConfig(t *testing.T) {
+ groupID := int64(9)
+ svc := &APIKeyService{}
+
+ apiKey, ok, err := svc.applyAuthCacheEntry("k-legacy-models-list", &APIKeyAuthCacheEntry{
+ Snapshot: &APIKeyAuthSnapshot{
+ Version: 10,
+ APIKeyID: 1,
+ UserID: 2,
+ GroupID: &groupID,
+ Status: StatusActive,
+ User: APIKeyAuthUserSnapshot{
+ ID: 2,
+ Status: StatusActive,
+ Role: RoleUser,
+ Balance: 10,
+ Concurrency: 3,
+ },
+ Group: &APIKeyAuthGroupSnapshot{
+ ID: groupID,
+ Name: "openai",
+ Platform: PlatformOpenAI,
+ Status: StatusActive,
+ SubscriptionType: SubscriptionTypeStandard,
+ RateMultiplier: 1,
+ },
+ },
+ })
+
+ if err != nil {
+ t.Fatalf("expected stale snapshot to be ignored without error, got %v", err)
+ }
+ if ok {
+ t.Fatalf("expected v10 auth snapshot to be rejected after models_list_config was added")
+ }
+ if apiKey != nil {
+ t.Fatalf("expected no API key from stale snapshot, got %#v", apiKey)
+ }
+}
diff --git a/backend/internal/service/api_key_service.go b/backend/internal/service/api_key_service.go
index 48e0ab2f..dc008b8a 100644
--- a/backend/internal/service/api_key_service.go
+++ b/backend/internal/service/api_key_service.go
@@ -55,6 +55,8 @@ type APIKeyRepository interface {
GetByKeyForAuth(ctx context.Context, key string) (*APIKey, error)
Update(ctx context.Context, key *APIKey) error
Delete(ctx context.Context, id int64) error
+ // DeleteWithAudit 在同一事务内先写 deleted_api_key_audits 审计、再软删除该 key。
+ DeleteWithAudit(ctx context.Context, id int64) error
ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, filters APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error)
VerifyOwnership(ctx context.Context, userID int64, apiKeyIDs []int64) ([]int64, error)
@@ -648,15 +650,16 @@ func (s *APIKeyService) Delete(ctx context.Context, id int64, userID int64) erro
return ErrInsufficientPerms
}
- // 清除Redis缓存(使用 userID 而非 apiKey.UserID)
+ // 事务内:写审计 + 软删除(tombstone)。
+ if err := s.apiKeyRepo.DeleteWithAudit(ctx, id); err != nil {
+ return fmt.Errorf("delete api key: %w", err)
+ }
+
+ // 删除成功后再清理缓存,避免"缓存已清但删除失败"的竞态。
if s.cache != nil {
_ = s.cache.DeleteCreateAttemptCount(ctx, userID)
}
s.InvalidateAuthCacheByKey(ctx, key)
-
- if err := s.apiKeyRepo.Delete(ctx, id); err != nil {
- return fmt.Errorf("delete api key: %w", err)
- }
s.lastUsedTouchL1.Delete(id)
return nil
diff --git a/backend/internal/service/api_key_service_cache_test.go b/backend/internal/service/api_key_service_cache_test.go
index eaac9a1c..a1dfbcb0 100644
--- a/backend/internal/service/api_key_service_cache_test.go
+++ b/backend/internal/service/api_key_service_cache_test.go
@@ -53,6 +53,10 @@ func (s *authRepoStub) Delete(ctx context.Context, id int64) error {
panic("unexpected Delete call")
}
+func (s *authRepoStub) DeleteWithAudit(ctx context.Context, id int64) error {
+ panic("unexpected DeleteWithAudit call")
+}
+
func (s *authRepoStub) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, filters APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error) {
panic("unexpected ListByUserID call")
}
diff --git a/backend/internal/service/api_key_service_delete_test.go b/backend/internal/service/api_key_service_delete_test.go
index 392d52b9..b8511c35 100644
--- a/backend/internal/service/api_key_service_delete_test.go
+++ b/backend/internal/service/api_key_service_delete_test.go
@@ -79,6 +79,12 @@ func (s *apiKeyRepoStub) Delete(ctx context.Context, id int64) error {
return s.deleteErr
}
+// DeleteWithAudit 与 Delete 一样记录被删除的 ID,供 service 测试断言。
+func (s *apiKeyRepoStub) DeleteWithAudit(ctx context.Context, id int64) error {
+ s.deletedIDs = append(s.deletedIDs, id)
+ return s.deleteErr
+}
+
// 以下是接口要求实现但本测试不关心的方法
func (s *apiKeyRepoStub) ListByUserID(ctx context.Context, userID int64, params pagination.PaginationParams, filters APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error) {
@@ -274,8 +280,8 @@ func TestApiKeyService_Delete_NotFound(t *testing.T) {
// 预期行为:
// - GetKeyAndOwnerID 返回正确的所有者 ID
// - 所有权验证通过
-// - 缓存被清除(在删除之前)
-// - Delete 被调用但返回错误
+// - DeleteWithAudit 被调用但返回错误
+// - 删除失败时缓存不被清除(缓存清理在删除成功后执行,消除竞态)
// - 返回包含 "delete api key" 的错误信息
func TestApiKeyService_Delete_DeleteFails(t *testing.T) {
repo := &apiKeyRepoStub{
@@ -288,7 +294,7 @@ func TestApiKeyService_Delete_DeleteFails(t *testing.T) {
err := svc.Delete(context.Background(), 3, 3) // API Key ID=3, 调用者 userID=3
require.Error(t, err)
require.ErrorContains(t, err, "delete api key")
- require.Equal(t, []int64{3}, repo.deletedIDs) // 验证删除操作被调用
- require.Equal(t, []int64{3}, cache.invalidated) // 验证缓存已被清除(即使删除失败)
- require.Equal(t, []string{svc.authCacheKey("k")}, cache.deleteAuthKeys)
+ require.Equal(t, []int64{3}, repo.deletedIDs) // 验证 DeleteWithAudit 被调用
+ require.Empty(t, cache.invalidated) // 验证删除失败时缓存未被清除(新顺序:先删后清)
+ require.Empty(t, cache.deleteAuthKeys) // 验证删除失败时 auth 缓存未被清除
}
diff --git a/backend/internal/service/api_key_service_quota_test.go b/backend/internal/service/api_key_service_quota_test.go
index cf05e16c..4d1d6f00 100644
--- a/backend/internal/service/api_key_service_quota_test.go
+++ b/backend/internal/service/api_key_service_quota_test.go
@@ -101,6 +101,9 @@ func (s *quotaBaseAPIKeyRepoStub) Update(context.Context, *APIKey) error {
func (s *quotaBaseAPIKeyRepoStub) Delete(context.Context, int64) error {
panic("unexpected Delete call")
}
+func (s *quotaBaseAPIKeyRepoStub) DeleteWithAudit(context.Context, int64) error {
+ panic("unexpected DeleteWithAudit call")
+}
func (s *quotaBaseAPIKeyRepoStub) ListByUserID(context.Context, int64, pagination.PaginationParams, APIKeyListFilters) ([]APIKey, *pagination.PaginationResult, error) {
panic("unexpected ListByUserID call")
}
diff --git a/backend/internal/service/auth_service_email_bind_test.go b/backend/internal/service/auth_service_email_bind_test.go
index 87867395..28bb0a3b 100644
--- a/backend/internal/service/auth_service_email_bind_test.go
+++ b/backend/internal/service/auth_service_email_bind_test.go
@@ -850,6 +850,9 @@ func (s *emailBindUserRepoStub) UnbindUserAuthProvider(context.Context, int64, s
func (s *emailBindUserRepoStub) UpdateTotpSecret(context.Context, int64, *string) error { return nil }
func (s *emailBindUserRepoStub) EnableTotp(context.Context, int64) error { return nil }
func (s *emailBindUserRepoStub) DisableTotp(context.Context, int64) error { return nil }
+func (s *emailBindUserRepoStub) GetByIDIncludeDeleted(ctx context.Context, id int64) (*service.User, error) {
+ return s.GetByID(ctx, id)
+}
func cloneEmailBindUser(user *service.User) *service.User {
if user == nil {
diff --git a/backend/internal/service/auth_service_platform_quota_test.go b/backend/internal/service/auth_service_platform_quota_test.go
index f58dc48c..46069814 100644
--- a/backend/internal/service/auth_service_platform_quota_test.go
+++ b/backend/internal/service/auth_service_platform_quota_test.go
@@ -43,6 +43,10 @@ func (f *fakeInsertRecorder) ResetExpiredWindow(_ context.Context, _ int64, _ st
return nil
}
+func (f *fakeInsertRecorder) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error {
+ return nil
+}
+
func TestSnapshotPlatformQuotaDefaults_PassesToRepoBulkInsert(t *testing.T) {
fakeRepo := &fakeInsertRecorder{}
s := &AuthService{userPlatformQuotaRepo: fakeRepo}
diff --git a/backend/internal/service/auth_service_register_test.go b/backend/internal/service/auth_service_register_test.go
index a7c0d260..2ee9f21a 100644
--- a/backend/internal/service/auth_service_register_test.go
+++ b/backend/internal/service/auth_service_register_test.go
@@ -105,6 +105,10 @@ func (s *userPlatformQuotaRepoStub) ResetExpiredWindow(context.Context, int64, s
panic("unexpected ResetExpiredWindow call")
}
+func (s *userPlatformQuotaRepoStub) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error {
+ return nil
+}
+
func (s *defaultSubscriptionAssignerStub) AssignOrExtendSubscription(_ context.Context, input *AssignSubscriptionInput) (*UserSubscription, bool, error) {
if input != nil {
s.calls = append(s.calls, *input)
diff --git a/backend/internal/service/bedrock_request_test.go b/backend/internal/service/bedrock_request_test.go
index 94f1a118..c5b71da7 100644
--- a/backend/internal/service/bedrock_request_test.go
+++ b/backend/internal/service/bedrock_request_test.go
@@ -174,6 +174,7 @@ func TestIsBedrockClaude45OrNewer(t *testing.T) {
expect bool
}{
{"us.anthropic.claude-opus-4-6-v1", true},
+ {"us.anthropic.claude-opus-4-8-v1", true},
{"us.anthropic.claude-sonnet-4-6", true},
{"us.anthropic.claude-sonnet-4-5-20250929-v1:0", true},
{"us.anthropic.claude-opus-4-5-20251101-v1:0", true},
@@ -511,6 +512,20 @@ func TestResolveBedrockModelID(t *testing.T) {
assert.Equal(t, "au.anthropic.claude-opus-4-6-v1", modelID)
})
+ t.Run("default opus 4.8 mapping uses regional Bedrock model id", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformAnthropic,
+ Type: AccountTypeBedrock,
+ Credentials: map[string]any{
+ "aws_region": "eu-west-1",
+ },
+ }
+
+ modelID, ok := ResolveBedrockModelID(account, "claude-opus-4-8")
+ require.True(t, ok)
+ assert.Equal(t, "eu.anthropic.claude-opus-4-8-v1", modelID)
+ })
+
t.Run("force global rewrites anthropic regional model id", func(t *testing.T) {
account := &Account{
Platform: PlatformAnthropic,
@@ -714,6 +729,7 @@ func TestIsBedrockOpus47OrNewer(t *testing.T) {
modelID string
expect bool
}{
+ {"us.anthropic.claude-opus-4-8-v1", true},
{"us.anthropic.claude-opus-4-7-v1", true},
{"us.anthropic.claude-opus-4-6-v1", false},
{"us.anthropic.claude-opus-4-5-20251101-v1:0", false},
@@ -886,10 +902,12 @@ func TestIsBedrockOpus47OrNewer_EdgeCases(t *testing.T) {
modelID string
expect bool
}{
+ {"anthropic.claude-opus-4-8-v1", true},
{"anthropic.claude-opus-4-7-v1", true},
{"us.anthropic.claude-opus-4-7-20270101-v1:0", true},
{"", false},
// Forward() passes parsed.Model (standard names), not Bedrock IDs
+ {"claude-opus-4-8", true},
{"claude-opus-4-7", true},
{"claude-opus-4-6", false},
{"claude-sonnet-4-7", false},
diff --git a/backend/internal/service/billing_cache_service.go b/backend/internal/service/billing_cache_service.go
index 2b7c06ba..b734fab1 100644
--- a/backend/internal/service/billing_cache_service.go
+++ b/backend/internal/service/billing_cache_service.go
@@ -689,7 +689,8 @@ func (s *BillingCacheService) IncrementUserPlatformQuotaUsage(userID int64, plat
ctx, cancel := context.WithTimeout(context.Background(), cacheWriteTimeout)
defer cancel()
ttl := time.Duration(s.cfg.Billing.UserPlatformQuotaCacheTTLSeconds) * time.Second
- if err := s.cache.IncrUserPlatformQuotaUsageCache(ctx, userID, platform, cost, ttl); err != nil {
+ markDirty := s.cfg.Database.UserPlatformQuotaFlusherEnabled
+ if err := s.cache.IncrUserPlatformQuotaUsageCache(ctx, userID, platform, cost, ttl, markDirty); err != nil {
logger.LegacyPrintf("service.billing_cache",
"ALERT: incr user platform quota cache failed user=%d platform=%s cost=%f: %v",
userID, platform, cost, err)
@@ -1096,7 +1097,12 @@ func (s *BillingCacheService) checkUserPlatformQuotaEligibility(
// 超时 50ms:覆盖正常路径与可接受抖动;Redis 异常时 hot path 不阻塞超过此值。
// 用 context.Background()+短超时,避免请求 ctx 取消导致刷新丢失。
// 显式 setCancel()(而非 defer):缩短 context 生命周期,避免 defer 延迟到函数返回。
- if windowExpired && s.cache != nil {
+ // isSentinel 判定「该 entry 无任何 limit」,涵盖两类,跨窗口命中时都跳过 refresh:
+ // 1) A3 回填的 sentinel(DB 无行,短 TTL):refresh 会把短 TTL 误升级为 86400s,有害;
+ // 2) DB 有行但三 limit 全未配置的用户(TTL 86400s):refresh 纯属无意义(TTL 升级本身无害)。
+ // 两类的 enforcement(下方 limit!=nil 比较)都因 limit 全 nil 永远放行,跳过 refresh 均正确。
+ isSentinel := entry.DailyLimitUSD == nil && entry.WeeklyLimitUSD == nil && entry.MonthlyLimitUSD == nil
+ if windowExpired && s.cache != nil && !isSentinel {
refreshed := &UserPlatformQuotaCacheEntry{
DailyUsageUSD: dailyUsage,
WeeklyUsageUSD: weeklyUsage,
@@ -1159,6 +1165,33 @@ func (s *BillingCacheService) checkUserPlatformQuotaEligibility(
}
rec, _ := v.(*UserPlatformQuotaRecord)
if rec == nil {
+ // 仅在 cache 可用且本次 GET 未出错时回填 sentinel:Redis GET 故障(cacheErr!=nil)
+ // 时不回填,与下方 line ~1201 "Redis 故障时 fail-open:不回填" 保持一致,
+ // 避免在 Redis 异常期做一次注定失败的 SET。
+ if s.cache != nil && cacheErr == nil {
+ now := time.Now()
+ startOfDay := timezone.StartOfDay(now)
+ startOfWeek := timezone.StartOfWeek(now)
+ sentinel := &UserPlatformQuotaCacheEntry{
+ SchemaVersion: UserPlatformQuotaCacheSchemaV1,
+ DailyWindowStart: &startOfDay,
+ WeeklyWindowStart: &startOfWeek,
+ MonthlyWindowStart: &now,
+ // limits 全 nil, usage 全 0(零值)
+ }
+ sentinelTTL := time.Duration(s.cfg.Billing.UserPlatformQuotaSentinelTTLSeconds) * time.Second
+ if sentinelTTL <= 0 {
+ // 防御:TTL<=0 时 Redis EXPIRE 会立即删除整个 key(见 billing_cache.go 的 pipe.Expire),
+ // sentinel 不持久化 → 每请求击穿 DB。配置缺失/误配为 0 时 fallback 到 1h。
+ sentinelTTL = time.Hour
+ }
+ setCtx, setCancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
+ if setErr := s.cache.SetUserPlatformQuotaCache(setCtx, userID, platform, sentinel, sentinelTTL); setErr != nil {
+ userPlatformQuotaSentinelSetCacheErrorTotal.Add(1)
+ logger.LegacyPrintf("service.billing_cache", "Warning: set sentinel quota cache failed user=%d platform=%s: %v", userID, platform, setErr)
+ }
+ setCancel()
+ }
return nil
}
@@ -1278,3 +1311,20 @@ func monthlyQuotaWindowExpired(start *time.Time, now time.Time) bool {
}
return now.Sub(*start) >= 30*24*time.Hour
}
+
+// HasUserPlatformQuotaLimit 判断该 user×platform 是否设了任一非 nil limit。
+// 写入点守卫:无 limit 直接跳过 Redis 写 + 脏集标记,消除无谓写入。
+// fail-safe:任何不确定(simple 模式除外)都返回 true 维持写入。
+func (s *BillingCacheService) HasUserPlatformQuotaLimit(ctx context.Context, userID int64, platform string) bool {
+ if s.cfg.RunMode == config.RunModeSimple {
+ return false
+ }
+ if s.cache == nil {
+ return true
+ }
+ entry, ok, err := s.cache.GetUserPlatformQuotaCache(ctx, userID, platform)
+ if err != nil || !ok || entry == nil {
+ return true
+ }
+ return entry.DailyLimitUSD != nil || entry.WeeklyLimitUSD != nil || entry.MonthlyLimitUSD != nil
+}
diff --git a/backend/internal/service/billing_cache_service_singleflight_test.go b/backend/internal/service/billing_cache_service_singleflight_test.go
index b443d97e..235b13a6 100644
--- a/backend/internal/service/billing_cache_service_singleflight_test.go
+++ b/backend/internal/service/billing_cache_service_singleflight_test.go
@@ -79,10 +79,22 @@ func (s *billingCacheMissStub) DeleteUserPlatformQuotaCache(ctx context.Context,
return nil
}
-func (s *billingCacheMissStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error {
+func (s *billingCacheMissStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error {
return nil
}
+func (s *billingCacheMissStub) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) {
+ return nil, nil
+}
+
+func (s *billingCacheMissStub) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error {
+ return nil
+}
+
+func (s *billingCacheMissStub) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) {
+ return nil, nil
+}
+
type balanceLoadUserRepoStub struct {
mockUserRepo
calls atomic.Int64
diff --git a/backend/internal/service/billing_cache_service_test.go b/backend/internal/service/billing_cache_service_test.go
index bcd086fa..c344b417 100644
--- a/backend/internal/service/billing_cache_service_test.go
+++ b/backend/internal/service/billing_cache_service_test.go
@@ -80,10 +80,22 @@ func (b *billingCacheWorkerStub) DeleteUserPlatformQuotaCache(ctx context.Contex
return nil
}
-func (b *billingCacheWorkerStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error {
+func (b *billingCacheWorkerStub) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error {
return nil
}
+func (b *billingCacheWorkerStub) PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error) {
+ return nil, nil
+}
+
+func (b *billingCacheWorkerStub) ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error {
+ return nil
+}
+
+func (b *billingCacheWorkerStub) BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) {
+ return nil, nil
+}
+
func TestBillingCacheServiceQueueHighLoad(t *testing.T) {
cache := &billingCacheWorkerStub{}
svc := NewBillingCacheService(cache, nil, nil, nil, nil, nil, &config.Config{}, nil)
diff --git a/backend/internal/service/billing_cache_service_user_platform_quota_test.go b/backend/internal/service/billing_cache_service_user_platform_quota_test.go
index 57697ddb..a82c0e00 100644
--- a/backend/internal/service/billing_cache_service_user_platform_quota_test.go
+++ b/backend/internal/service/billing_cache_service_user_platform_quota_test.go
@@ -20,14 +20,15 @@ type fakeIncrCache struct {
}
type incrCall struct {
- userID int64
- platform string
- cost float64
- ttl time.Duration
+ userID int64
+ platform string
+ cost float64
+ ttl time.Duration
+ markDirty bool
}
-func (f *fakeIncrCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error {
- f.calls = append(f.calls, incrCall{userID, platform, cost, ttl})
+func (f *fakeIncrCache) IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error {
+ f.calls = append(f.calls, incrCall{userID, platform, cost, ttl, markDirty})
return nil
}
@@ -49,10 +50,10 @@ func TestIncrementUserPlatformQuotaUsage_SyncCallsCache(t *testing.T) {
if len(fake.calls) != 2 {
t.Fatalf("expected 2 incr calls, got %d", len(fake.calls))
}
- if fake.calls[0] != (incrCall{101, "anthropic", 0.25, 120 * time.Second}) {
+ if fake.calls[0] != (incrCall{userID: 101, platform: "anthropic", cost: 0.25, ttl: 120 * time.Second, markDirty: false}) {
t.Errorf("call[0] = %+v", fake.calls[0])
}
- if fake.calls[1] != (incrCall{101, "openai", 0.50, 120 * time.Second}) {
+ if fake.calls[1] != (incrCall{userID: 101, platform: "openai", cost: 0.50, ttl: 120 * time.Second, markDirty: false}) {
t.Errorf("call[1] = %+v", fake.calls[1])
}
}
@@ -88,13 +89,23 @@ func (f *fakeQuotaRepo) ResetExpiredWindow(_ context.Context, _ int64, _ string,
return nil
}
-// fakeFullCache 同时支持 Get + Set + Incr + Delete。
+func (f *fakeQuotaRepo) BatchSnapshotUsage(_ context.Context, _ []UserPlatformQuotaSnapshot, _ time.Time) error {
+ return nil
+}
+
+// fakeFullCache 同时支持 Get + Set + Incr + Delete + Pop/Readd/BatchGet(脏集读写)。
// mu 保护 entry 和 deleteCalls,防止异步 goroutine 与主 goroutine 之间的 data race。
type fakeFullCache struct {
BillingCache
mu sync.Mutex
entry *UserPlatformQuotaCacheEntry
deleteCalls int
+ setCalls int // SetUserPlatformQuotaCache 调用次数
+ lastSetTTL time.Duration // 最近一次 Set 的 ttl
+ getErr error // 非 nil 时 Get 先返回 (nil,false,getErr)
+ setErr error // 非 nil 时 Set 返回该 err(setCalls 仍+1)
+ // dirty 模拟脏集,供 flusher 测试使用。
+ dirty map[UserPlatformQuotaKey]struct{}
}
// getDeleteCalls 线程安全地读取 deleteCalls。
@@ -111,19 +122,41 @@ func (f *fakeFullCache) getEntry() *UserPlatformQuotaCacheEntry {
return f.entry
}
+// getSetCalls 线程安全地读取 setCalls。
+func (f *fakeFullCache) getSetCalls() int {
+ f.mu.Lock()
+ defer f.mu.Unlock()
+ return f.setCalls
+}
+
+// getLastSetTTL 线程安全地读取 lastSetTTL。
+func (f *fakeFullCache) getLastSetTTL() time.Duration {
+ f.mu.Lock()
+ defer f.mu.Unlock()
+ return f.lastSetTTL
+}
+
func (f *fakeFullCache) GetUserPlatformQuotaCache(_ context.Context, _ int64, _ string) (*UserPlatformQuotaCacheEntry, bool, error) {
f.mu.Lock()
defer f.mu.Unlock()
+ if f.getErr != nil {
+ return nil, false, f.getErr
+ }
if f.entry == nil {
return nil, false, nil
}
return f.entry, true, nil
}
-func (f *fakeFullCache) SetUserPlatformQuotaCache(_ context.Context, _ int64, _ string, e *UserPlatformQuotaCacheEntry, _ time.Duration) error {
+func (f *fakeFullCache) SetUserPlatformQuotaCache(_ context.Context, _ int64, _ string, e *UserPlatformQuotaCacheEntry, ttl time.Duration) error {
f.mu.Lock()
defer f.mu.Unlock()
+ f.setCalls++
+ if f.setErr != nil {
+ return f.setErr
+ }
f.entry = e
+ f.lastSetTTL = ttl
return nil
}
@@ -135,6 +168,48 @@ func (f *fakeFullCache) DeleteUserPlatformQuotaCache(_ context.Context, _ int64,
return nil
}
+func (f *fakeFullCache) PopDirtyUserPlatformQuotaKeys(_ context.Context, n int) ([]UserPlatformQuotaKey, error) {
+ f.mu.Lock()
+ defer f.mu.Unlock()
+ if len(f.dirty) == 0 {
+ return nil, nil
+ }
+ keys := make([]UserPlatformQuotaKey, 0, n)
+ for k := range f.dirty {
+ if len(keys) >= n {
+ break
+ }
+ keys = append(keys, k)
+ delete(f.dirty, k)
+ }
+ return keys, nil
+}
+
+func (f *fakeFullCache) ReaddDirtyUserPlatformQuotaKeys(_ context.Context, keys []UserPlatformQuotaKey) error {
+ f.mu.Lock()
+ defer f.mu.Unlock()
+ if f.dirty == nil {
+ f.dirty = make(map[UserPlatformQuotaKey]struct{})
+ }
+ for _, k := range keys {
+ f.dirty[k] = struct{}{}
+ }
+ return nil
+}
+
+// BatchGetUserPlatformQuotaCache 对每个 key 返回 f.entry(MISS → nil),
+// 保持与输入 keys 顺序/长度对齐。注意此处所有 key 共享同一个 entry,
+// 仅用于测试场景。
+func (f *fakeFullCache) BatchGetUserPlatformQuotaCache(_ context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) {
+ f.mu.Lock()
+ defer f.mu.Unlock()
+ results := make([]*UserPlatformQuotaCacheEntry, len(keys))
+ for i := range keys {
+ results[i] = f.entry
+ }
+ return results, nil
+}
+
func newServiceForPreflight(t *testing.T, repo UserPlatformQuotaRepository, cache BillingCache) *BillingCacheService {
t.Helper()
cfg := &config.Config{}
@@ -593,3 +668,165 @@ func TestMonthlyQuotaWindowExpired_BoundaryTable(t *testing.T) {
})
}
}
+
+// TestCheckUserPlatformQuotaEligibility_NoRow_WritesSentinel 验证:
+// cache MISS + DB 无行时,回填 sentinel entry(三 limit 全 nil,三 window_start 全 non-nil),
+// TTL = UserPlatformQuotaSentinelTTLSeconds,函数返回 nil(fail-open)。
+func TestCheckUserPlatformQuotaEligibility_NoRow_WritesSentinel(t *testing.T) {
+ repo := &fakeQuotaRepo{rec: nil} // DB 无行
+ cache := &fakeFullCache{} // entry=nil → Get 返回 MISS
+ svc := newServiceForPreflight(t, repo, cache)
+ svc.cfg.Billing.UserPlatformQuotaSentinelTTLSeconds = 3600
+
+ if err := svc.checkUserPlatformQuotaEligibility(context.Background(), 1, "anthropic"); err != nil {
+ t.Fatalf("expected nil (fail-open), got %v", err)
+ }
+ if cache.getSetCalls() != 1 {
+ t.Fatalf("expected 1 SetUserPlatformQuotaCache call for sentinel, got %d", cache.getSetCalls())
+ }
+ sentinel := cache.getEntry()
+ if sentinel == nil {
+ t.Fatal("expected sentinel entry backfilled")
+ }
+ if sentinel.DailyLimitUSD != nil || sentinel.WeeklyLimitUSD != nil || sentinel.MonthlyLimitUSD != nil {
+ t.Errorf("sentinel must have all-nil limits")
+ }
+ if sentinel.DailyWindowStart == nil || sentinel.WeeklyWindowStart == nil || sentinel.MonthlyWindowStart == nil {
+ t.Errorf("sentinel must have non-nil window_start to avoid refresh churn")
+ }
+ if sentinel.SchemaVersion != UserPlatformQuotaCacheSchemaV1 {
+ t.Errorf("sentinel schema = %d, want V1", sentinel.SchemaVersion)
+ }
+ if cache.getLastSetTTL() != 3600*time.Second {
+ t.Errorf("sentinel ttl = %v, want 3600s", cache.getLastSetTTL())
+ }
+}
+
+// TestCheckUserPlatformQuotaEligibility_RedisGetError_NoSentinelBackfill 验证:
+// Redis GET 故障(cacheErr!=nil)+ DB 无行时,不应回填 sentinel(与 "Redis 故障时不回填" 一致),且 fail-open。
+func TestCheckUserPlatformQuotaEligibility_RedisGetError_NoSentinelBackfill(t *testing.T) {
+ repo := &fakeQuotaRepo{rec: nil}
+ cache := &fakeFullCache{getErr: errors.New("redis get down")}
+ svc := newServiceForPreflight(t, repo, cache)
+ svc.cfg.Billing.UserPlatformQuotaSentinelTTLSeconds = 3600
+
+ if err := svc.checkUserPlatformQuotaEligibility(context.Background(), 1, "anthropic"); err != nil {
+ t.Fatalf("redis 故障应 fail-open, got %v", err)
+ }
+ if cache.getSetCalls() != 0 {
+ t.Errorf("redis-get-error 时不应回填 sentinel, got %d set calls", cache.getSetCalls())
+ }
+}
+
+// TestCheckUserPlatformQuotaEligibility_NoRow_SentinelSetFailsFailOpen 验证:
+// sentinel SET 失败时 fail-open(返回 nil)且计 metric。
+func TestCheckUserPlatformQuotaEligibility_NoRow_SentinelSetFailsFailOpen(t *testing.T) {
+ before := userPlatformQuotaSentinelSetCacheErrorTotal.Load()
+ repo := &fakeQuotaRepo{rec: nil}
+ cache := &fakeFullCache{setErr: errors.New("redis set timeout")}
+ svc := newServiceForPreflight(t, repo, cache)
+ svc.cfg.Billing.UserPlatformQuotaSentinelTTLSeconds = 3600
+
+ if err := svc.checkUserPlatformQuotaEligibility(context.Background(), 1, "anthropic"); err != nil {
+ t.Fatalf("sentinel set 失败应 fail-open, got %v", err)
+ }
+ if cache.getSetCalls() != 1 {
+ t.Errorf("应尝试 set sentinel 恰好一次, got %d", cache.getSetCalls())
+ }
+ if got := userPlatformQuotaSentinelSetCacheErrorTotal.Load() - before; got != 1 {
+ t.Errorf("set 失败应使 metric +1, got delta %d", got)
+ }
+}
+
+// TestCheckUserPlatformQuotaEligibility_SentinelCrossDay_NoRefresh 验证:
+// 命中 sentinel(三 limit 全 nil)且跨窗口(daily/weekly 过期)时,不应触发 refresh SetCache
+// (否则会把短 sentinel TTL 误升级为 quota cache 默认 86400s)。
+func TestCheckUserPlatformQuotaEligibility_SentinelCrossDay_NoRefresh(t *testing.T) {
+ yesterday := timezone.StartOfDay(time.Now().AddDate(0, 0, -1))
+ lastWeek := timezone.StartOfWeek(time.Now().AddDate(0, 0, -7))
+ monthAgoOK := time.Now().AddDate(0, 0, -5) // <30d, monthly 不过期
+ sentinel := &UserPlatformQuotaCacheEntry{
+ SchemaVersion: UserPlatformQuotaCacheSchemaV1,
+ DailyWindowStart: &yesterday, // 跨日 → daily windowExpired = true
+ WeeklyWindowStart: &lastWeek, // 跨周 → weekly windowExpired = true
+ MonthlyWindowStart: &monthAgoOK,
+ // limits 全 nil → sentinel
+ }
+ cache := &fakeFullCache{entry: sentinel} // entry 非 nil → Get HIT
+ svc := newServiceForPreflight(t, &fakeQuotaRepo{}, cache)
+
+ if err := svc.checkUserPlatformQuotaEligibility(context.Background(), 1, "anthropic"); err != nil {
+ t.Fatalf("sentinel = no limit, expected nil, got %v", err)
+ }
+ if cache.getSetCalls() != 0 {
+ t.Errorf("sentinel cross-window must NOT trigger refresh SetCache, got %d calls", cache.getSetCalls())
+ }
+}
+
+// ── TestHasUserPlatformQuotaLimit ────────────────────────────────────────────
+
+func TestHasUserPlatformQuotaLimit(t *testing.T) {
+ daily := 5.0
+
+ tests := []struct {
+ name string
+ setup func() *BillingCacheService
+ want bool
+ }{
+ {
+ name: "has_limit",
+ setup: func() *BillingCacheService {
+ entry := &UserPlatformQuotaCacheEntry{DailyLimitUSD: &daily}
+ svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{entry: entry})
+ return svc
+ },
+ want: true,
+ },
+ {
+ name: "sentinel_no_limit",
+ setup: func() *BillingCacheService {
+ entry := &UserPlatformQuotaCacheEntry{} // 三个 limit 字段全 nil
+ svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{entry: entry})
+ return svc
+ },
+ want: false,
+ },
+ {
+ name: "cache_miss",
+ setup: func() *BillingCacheService {
+ // entry==nil → GetUserPlatformQuotaCache 返回 (nil,false,nil)
+ svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{})
+ return svc
+ },
+ want: true, // fail-safe
+ },
+ {
+ name: "redis_err",
+ setup: func() *BillingCacheService {
+ svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{getErr: errors.New("redis down")})
+ return svc
+ },
+ want: true, // fail-safe
+ },
+ {
+ name: "simple_mode",
+ setup: func() *BillingCacheService {
+ entry := &UserPlatformQuotaCacheEntry{DailyLimitUSD: &daily}
+ svc := newServiceForPreflight(t, &fakeQuotaRepo{}, &fakeFullCache{entry: entry})
+ svc.cfg.RunMode = config.RunModeSimple
+ return svc
+ },
+ want: false, // simple 模式始终跳过
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ svc := tt.setup()
+ got := svc.HasUserPlatformQuotaLimit(context.Background(), 1, "anthropic")
+ if got != tt.want {
+ t.Errorf("HasUserPlatformQuotaLimit() = %v, want %v", got, tt.want)
+ }
+ })
+ }
+}
diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go
index 373502cf..6b1438e8 100644
--- a/backend/internal/service/billing_service.go
+++ b/backend/internal/service/billing_service.go
@@ -21,6 +21,12 @@ type APIKeyRateLimitCacheData struct {
Window7d int64 `json:"window_7d"`
}
+// UserPlatformQuotaKey 标识一个 user×platform,用于脏集出入与批量读。
+type UserPlatformQuotaKey struct {
+ UserID int64
+ Platform string
+}
+
// UserPlatformQuotaCacheEntry Redis hash 反序列化结果。
//
// SchemaVersion 用于向后兼容:
@@ -72,7 +78,13 @@ type BillingCache interface {
SetUserPlatformQuotaCache(ctx context.Context, userID int64, platform string, entry *UserPlatformQuotaCacheEntry, ttl time.Duration) error
DeleteUserPlatformQuotaCache(ctx context.Context, userID int64, platform string) error
// IncrUserPlatformQuotaUsageCache 在缓存命中时累加用量;缓存未命中(key 不存在)静默返回 nil。
- IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration) error
+ // markDirty=true 时将该 key 的 member 写入 Redis 脏集,供 flusher 批量回写 DB。
+ IncrUserPlatformQuotaUsageCache(ctx context.Context, userID int64, platform string, cost float64, ttl time.Duration, markDirty bool) error
+
+ // 脏集读写,供 flusher 使用。
+ PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error)
+ ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error
+ BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error)
}
// ModelPricing 模型价格配置(per-token价格,与LiteLLM格式一致)
@@ -516,6 +528,7 @@ func (s *BillingService) computeTokenBreakdown(
inputPrice := pricing.InputPricePerToken
outputPrice := pricing.OutputPricePerToken
cacheReadPrice := pricing.CacheReadPricePerToken
+ cacheCreationMultiplier := 1.0
tierMultiplier := 1.0
if usePriorityServiceTierPricing(serviceTier, pricing) {
@@ -535,6 +548,13 @@ func (s *BillingService) computeTokenBreakdown(
if applyLongCtx && s.shouldApplySessionLongContextPricing(tokens, pricing) {
inputPrice *= pricing.LongContextInputMultiplier
outputPrice *= pricing.LongContextOutputMultiplier
+ // 缓存读取本质上是输入侧的复用,应与 input 一同应用长上下文倍率;
+ // 否则 cache hit 越多,少计的费用越多(见 #2293)。
+ cacheReadPrice *= pricing.LongContextInputMultiplier
+ // 缓存创建(cache_write)也是输入侧操作,三档价格(标准 / 5m / 1h)
+ // 都通过 computeCacheCreationCost 直接读取 pricing.*,不会经过这里
+ // 的倍率修改,因此显式向下传一个倍率,避免长上下文场景下被漏乘。
+ cacheCreationMultiplier = pricing.LongContextInputMultiplier
}
bd := &CostBreakdown{}
@@ -557,7 +577,7 @@ func (s *BillingService) computeTokenBreakdown(
}
// 缓存创建费用
- bd.CacheCreationCost = s.computeCacheCreationCost(pricing, tokens)
+ bd.CacheCreationCost = s.computeCacheCreationCost(pricing, tokens, cacheCreationMultiplier)
bd.CacheReadCost = float64(tokens.CacheReadTokens) * cacheReadPrice
@@ -577,16 +597,17 @@ func (s *BillingService) computeTokenBreakdown(
}
// computeCacheCreationCost 计算缓存创建费用(支持 5m/1h 分类或标准计费)。
-func (s *BillingService) computeCacheCreationCost(pricing *ModelPricing, tokens UsageTokens) float64 {
+// multiplier 用于长上下文等场景下的整体价格缩放(普通调用传 1.0 即可)。
+func (s *BillingService) computeCacheCreationCost(pricing *ModelPricing, tokens UsageTokens, multiplier float64) float64 {
if pricing.SupportsCacheBreakdown && (pricing.CacheCreation5mPrice > 0 || pricing.CacheCreation1hPrice > 0) {
if tokens.CacheCreation5mTokens == 0 && tokens.CacheCreation1hTokens == 0 && tokens.CacheCreationTokens > 0 {
// API 未返回 ephemeral 明细,回退到全部按 5m 单价计费
- return float64(tokens.CacheCreationTokens) * pricing.CacheCreation5mPrice
+ return float64(tokens.CacheCreationTokens) * pricing.CacheCreation5mPrice * multiplier
}
- return float64(tokens.CacheCreation5mTokens)*pricing.CacheCreation5mPrice +
- float64(tokens.CacheCreation1hTokens)*pricing.CacheCreation1hPrice
+ return float64(tokens.CacheCreation5mTokens)*pricing.CacheCreation5mPrice*multiplier +
+ float64(tokens.CacheCreation1hTokens)*pricing.CacheCreation1hPrice*multiplier
}
- return float64(tokens.CacheCreationTokens) * pricing.CacheCreationPricePerToken
+ return float64(tokens.CacheCreationTokens) * pricing.CacheCreationPricePerToken * multiplier
}
// calculatePerRequestCost 按次/图片计费
diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go
index df3e3a0a..0ab1f50d 100644
--- a/backend/internal/service/billing_service_test.go
+++ b/backend/internal/service/billing_service_test.go
@@ -197,6 +197,138 @@ func TestCalculateCost_OpenAIGPT54LongContextAppliesWholeSessionMultipliers(t *t
require.InDelta(t, expectedInput+expectedOutput, cost.ActualCost, 1e-10)
}
+// 回归测试 #2293:长上下文计费触发时,cache_read_tokens 也应应用 LongContextInputMultiplier。
+// 修复前:CacheReadCost = tokens * 0.25e-6 (漏乘倍率,少计费用)。
+// 修复后:CacheReadCost = tokens * 0.25e-6 * LongContextInputMultiplier(=2.0)。
+func TestCalculateCost_OpenAIGPT54LongContextAppliesMultiplierToCacheRead(t *testing.T) {
+ svc := newTestBillingService()
+
+ // InputTokens + CacheReadTokens = 1000 + 300000 = 301000 > 272000 阈值
+ tokens := UsageTokens{
+ InputTokens: 1000,
+ CacheReadTokens: 300000,
+ OutputTokens: 1000,
+ }
+
+ cost, err := svc.CalculateCost("gpt-5.4-2026-03-05", tokens, 1.0)
+ require.NoError(t, err)
+
+ expectedInput := float64(tokens.InputTokens) * 2.5e-6 * 2.0
+ expectedOutput := float64(tokens.OutputTokens) * 15e-6 * 1.5
+ expectedCacheRead := float64(tokens.CacheReadTokens) * 0.25e-6 * 2.0
+
+ require.InDelta(t, expectedInput, cost.InputCost, 1e-10)
+ require.InDelta(t, expectedOutput, cost.OutputCost, 1e-10)
+ require.InDelta(t, expectedCacheRead, cost.CacheReadCost, 1e-10,
+ "cache_read_cost should be scaled by LongContextInputMultiplier when long-context pricing applies (issue #2293)")
+
+ expectedTotal := expectedInput + expectedOutput + expectedCacheRead
+ require.InDelta(t, expectedTotal, cost.TotalCost, 1e-10)
+ require.InDelta(t, expectedTotal, cost.ActualCost, 1e-10)
+}
+
+// 阴性测试:未触发长上下文时,cache_read_price 不应被错误地乘以倍率。
+func TestCalculateCost_OpenAIGPT54NoLongContextKeepsCacheReadAtBasePrice(t *testing.T) {
+ svc := newTestBillingService()
+
+ // InputTokens + CacheReadTokens = 1000 + 100000 = 101000 < 272000 阈值,不触发长上下文
+ tokens := UsageTokens{
+ InputTokens: 1000,
+ CacheReadTokens: 100000,
+ OutputTokens: 1000,
+ }
+
+ cost, err := svc.CalculateCost("gpt-5.4-2026-03-05", tokens, 1.0)
+ require.NoError(t, err)
+
+ expectedCacheRead := float64(tokens.CacheReadTokens) * 0.25e-6
+ require.InDelta(t, expectedCacheRead, cost.CacheReadCost, 1e-10,
+ "cache_read_cost should remain at base price when below long-context threshold")
+}
+
+// 回归测试 #2816 follow-up:长上下文计费触发时,cache_creation_tokens 也应应用
+// LongContextInputMultiplier。computeCacheCreationCost 直接读取 pricing.* 价格,
+// 不经过 computeTokenBreakdown 内的 inputPrice / cacheReadPrice 倍率修改,因此
+// 修复前 cache_creation 部分会按基础价计算,少计费用约 50%(默认倍率 2.0)。
+func TestCalculateCost_OpenAIGPT54LongContextAppliesMultiplierToCacheCreation(t *testing.T) {
+ svc := newTestBillingService()
+
+ // InputTokens + CacheReadTokens = 1000 + 300000 = 301000 > 272000 阈值
+ tokens := UsageTokens{
+ InputTokens: 1000,
+ CacheReadTokens: 300000,
+ CacheCreationTokens: 10000,
+ OutputTokens: 1000,
+ }
+
+ cost, err := svc.CalculateCost("gpt-5.4-2026-03-05", tokens, 1.0)
+ require.NoError(t, err)
+
+ // gpt-5.4 fallback: CacheCreationPricePerToken = 2.5e-6, LongContextInputMultiplier = 2.0
+ expectedCacheCreation := float64(tokens.CacheCreationTokens) * 2.5e-6 * 2.0
+ require.InDelta(t, expectedCacheCreation, cost.CacheCreationCost, 1e-10,
+ "cache_creation_cost should be scaled by LongContextInputMultiplier when long-context pricing applies")
+}
+
+// 阴性测试:未触发长上下文时,cache_creation_price 不应被错误地乘以倍率。
+func TestCalculateCost_OpenAIGPT54NoLongContextKeepsCacheCreationAtBasePrice(t *testing.T) {
+ svc := newTestBillingService()
+
+ // InputTokens + CacheReadTokens = 1000 + 100000 = 101000 < 272000 阈值,不触发长上下文
+ tokens := UsageTokens{
+ InputTokens: 1000,
+ CacheReadTokens: 100000,
+ CacheCreationTokens: 10000,
+ OutputTokens: 1000,
+ }
+
+ cost, err := svc.CalculateCost("gpt-5.4-2026-03-05", tokens, 1.0)
+ require.NoError(t, err)
+
+ expectedCacheCreation := float64(tokens.CacheCreationTokens) * 2.5e-6
+ require.InDelta(t, expectedCacheCreation, cost.CacheCreationCost, 1e-10,
+ "cache_creation_cost should remain at base price when below long-context threshold")
+}
+
+// 覆盖 5m / 1h ephemeral 分类计费路径:长上下文触发时两档价格都应被倍率缩放。
+// 使用手工构造的 pricing(参考 TestCalculateCost_SupportsCacheBreakdown 的写法)
+// 以便同时控制 SupportsCacheBreakdown + 长上下文阈值。
+func TestCalculateCost_LongContextAppliesMultiplierToCacheCreation5mAnd1h(t *testing.T) {
+ svc := &BillingService{
+ cfg: &config.Config{},
+ fallbackPrices: map[string]*ModelPricing{
+ "claude-sonnet-4": {
+ InputPricePerToken: 3e-6,
+ OutputPricePerToken: 15e-6,
+ CacheReadPricePerToken: 0.3e-6,
+ SupportsCacheBreakdown: true,
+ CacheCreation5mPrice: 4e-6,
+ CacheCreation1hPrice: 5e-6,
+ LongContextInputThreshold: 272000,
+ LongContextInputMultiplier: 2.0,
+ LongContextOutputMultiplier: 1.5,
+ },
+ },
+ }
+
+ // InputTokens + CacheReadTokens = 1000 + 300000 = 301000 > 272000 阈值
+ tokens := UsageTokens{
+ InputTokens: 1000,
+ CacheReadTokens: 300000,
+ CacheCreation5mTokens: 8000,
+ CacheCreation1hTokens: 4000,
+ OutputTokens: 1000,
+ }
+
+ cost, err := svc.CalculateCost("claude-sonnet-4", tokens, 1.0)
+ require.NoError(t, err)
+
+ expected5m := float64(tokens.CacheCreation5mTokens) * 4e-6 * 2.0
+ expected1h := float64(tokens.CacheCreation1hTokens) * 5e-6 * 2.0
+ require.InDelta(t, expected5m+expected1h, cost.CacheCreationCost, 1e-10,
+ "both 5m and 1h cache_creation prices should be scaled by LongContextInputMultiplier")
+}
+
func TestGetFallbackPricing_FamilyMatching(t *testing.T) {
svc := newTestBillingService()
diff --git a/backend/internal/service/claude_code_validator.go b/backend/internal/service/claude_code_validator.go
index 4e8ced67..2c5ded6f 100644
--- a/backend/internal/service/claude_code_validator.go
+++ b/backend/internal/service/claude_code_validator.go
@@ -56,7 +56,7 @@ func NewClaudeCodeValidator() *ClaudeCodeValidator {
// 采用与 claude-relay-service 完全一致的验证策略:
//
// Step 1: User-Agent 检查 (必需) - 必须是 claude-cli/x.x.x
-// Step 2: 对于非 messages 路径,只要 UA 匹配就通过
+// Step 2: 对于非 messages 路径和 /messages/count_tokens,只要 UA 匹配就通过
// Step 3: 检查 max_tokens=1 + haiku 探测请求绕过(UA 已验证)
// Step 4: 对于 messages 路径,进行严格验证:
// - System prompt 相似度检查
@@ -71,12 +71,17 @@ func (v *ClaudeCodeValidator) Validate(r *http.Request, body map[string]any) boo
return false
}
- // Step 2: 非 messages 路径,只要 UA 匹配就通过
+ // Step 2: 非 messages 路径只要 UA 匹配就通过
path := r.URL.Path
if !strings.Contains(path, "messages") {
return true
}
+ // count_tokens 是 Claude Code 官方辅助请求,通常不携带完整 messages system prompt。
+ if isMessagesCountTokensPath(path) {
+ return true
+ }
+
// Step 3: 检查 max_tokens=1 + haiku 探测请求绕过
// 这类请求用于 Claude Code 验证 API 连通性,不携带 system prompt
if isMaxTokensOneHaiku, ok := IsMaxTokensOneHaikuRequestFromContext(r.Context()); ok && isMaxTokensOneHaiku {
@@ -128,6 +133,10 @@ func (v *ClaudeCodeValidator) Validate(r *http.Request, body map[string]any) boo
return true
}
+func isMessagesCountTokensPath(path string) bool {
+ return strings.HasSuffix(path, "/messages/count_tokens")
+}
+
// hasClaudeCodeSystemPrompt 检查请求是否包含 Claude Code 系统提示词
// 使用字符串相似度匹配(Dice coefficient)
func (v *ClaudeCodeValidator) hasClaudeCodeSystemPrompt(body map[string]any) bool {
diff --git a/backend/internal/service/claude_code_validator_test.go b/backend/internal/service/claude_code_validator_test.go
index f87c56e8..a4b30505 100644
--- a/backend/internal/service/claude_code_validator_test.go
+++ b/backend/internal/service/claude_code_validator_test.go
@@ -48,6 +48,97 @@ func TestClaudeCodeValidator_MessagesWithoutProbeStillNeedStrictValidation(t *te
require.False(t, ok)
}
+func TestClaudeCodeValidator_CountTokensPathUAOnly(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages/count_tokens", nil)
+ req.Header.Set("User-Agent", "claude-cli/2.1.156 (Claude Code)")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ })
+ require.True(t, ok)
+}
+
+func TestClaudeCodeValidator_CountTokensPathRequiresUA(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages/count_tokens", nil)
+ req.Header.Set("User-Agent", "curl/8.0.0")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ })
+ require.False(t, ok)
+}
+
+func TestClaudeCodeValidator_MessagesPathFullValid(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages", nil)
+ req.Header.Set("User-Agent", "claude-cli/2.1.156 (Claude Code)")
+ req.Header.Set("X-App", "claude-code")
+ req.Header.Set("anthropic-beta", "claude-code-20250219")
+ req.Header.Set("anthropic-version", "2023-06-01")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ "stream": true,
+ "system": []any{
+ map[string]any{
+ "type": "text",
+ "text": "You are Claude Code, Anthropic's official CLI for Claude.",
+ },
+ },
+ "metadata": map[string]any{
+ "user_id": "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
+ },
+ })
+ require.True(t, ok)
+}
+
+func TestClaudeCodeValidator_MessagesPathRejectsNonClaudeCodeUA(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages", nil)
+ req.Header.Set("User-Agent", "curl/8.0.0")
+ req.Header.Set("X-App", "claude-code")
+ req.Header.Set("anthropic-beta", "claude-code-20250219")
+ req.Header.Set("anthropic-version", "2023-06-01")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ "stream": true,
+ "system": []any{
+ map[string]any{
+ "type": "text",
+ "text": "You are Claude Code, Anthropic's official CLI for Claude.",
+ },
+ },
+ "metadata": map[string]any{
+ "user_id": "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
+ },
+ })
+ require.False(t, ok)
+}
+
+func TestClaudeCodeValidator_MessagesPathWithoutSystemPromptStillRejected(t *testing.T) {
+ validator := NewClaudeCodeValidator()
+ req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/messages", nil)
+ req.Header.Set("User-Agent", "claude-cli/2.1.156 (Claude Code)")
+ req.Header.Set("X-App", "claude-code")
+ req.Header.Set("anthropic-beta", "claude-code-20250219")
+ req.Header.Set("anthropic-version", "2023-06-01")
+
+ ok := validator.Validate(req, map[string]any{
+ "model": "claude-opus-4-8",
+ "stream": true,
+ "messages": []any{
+ map[string]any{"role": "user", "content": "hello"},
+ },
+ "metadata": map[string]any{
+ "user_id": "user_aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa_account__session_aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa",
+ },
+ })
+ require.False(t, ok)
+}
+
func TestClaudeCodeValidator_NonMessagesPathUAOnly(t *testing.T) {
validator := NewClaudeCodeValidator()
req := httptest.NewRequest(http.MethodPost, "http://example.com/v1/models", nil)
diff --git a/backend/internal/service/content_moderation.go b/backend/internal/service/content_moderation.go
index a5a84d7b..42b909c9 100644
--- a/backend/internal/service/content_moderation.go
+++ b/backend/internal/service/content_moderation.go
@@ -211,6 +211,20 @@ type ContentModerationAPIKeyStatus struct {
Configured bool `json:"configured"`
}
+type ContentModerationAPIKeyLoad struct {
+ Index int `json:"index"`
+ KeyHash string `json:"key_hash"`
+ Masked string `json:"masked"`
+ Status string `json:"status"`
+ Active int64 `json:"active"`
+ Total int64 `json:"total"`
+ Success int64 `json:"success"`
+ Errors int64 `json:"errors"`
+ AvgLatencyMS int64 `json:"avg_latency_ms"`
+ LastLatencyMS int `json:"last_latency_ms"`
+ LastHTTPStatus int `json:"last_http_status"`
+}
+
type TestContentModerationAPIKeysInput struct {
APIKeys []string `json:"api_keys"`
BaseURL string `json:"base_url"`
@@ -399,25 +413,35 @@ type ContentModerationCleanupResult struct {
}
type ContentModerationRuntimeStatus struct {
- Enabled bool `json:"enabled"`
- RiskControlEnabled bool `json:"risk_control_enabled"`
- Mode string `json:"mode"`
- WorkerCount int `json:"worker_count"`
- MaxWorkers int `json:"max_workers"`
- ActiveWorkers int `json:"active_workers"`
- IdleWorkers int `json:"idle_workers"`
- QueueSize int `json:"queue_size"`
- QueueLength int `json:"queue_length"`
- QueueUsagePercent float64 `json:"queue_usage_percent"`
- Enqueued int64 `json:"enqueued"`
- Dropped int64 `json:"dropped"`
- Processed int64 `json:"processed"`
- Errors int64 `json:"errors"`
- APIKeyStatuses []ContentModerationAPIKeyStatus `json:"api_key_statuses"`
- FlaggedHashCount int64 `json:"flagged_hash_count"`
- LastCleanupAt *time.Time `json:"last_cleanup_at,omitempty"`
- LastCleanupDeletedHit int64 `json:"last_cleanup_deleted_hit"`
- LastCleanupDeletedNonHit int64 `json:"last_cleanup_deleted_non_hit"`
+ Enabled bool `json:"enabled"`
+ RiskControlEnabled bool `json:"risk_control_enabled"`
+ Mode string `json:"mode"`
+ WorkerCount int `json:"worker_count"`
+ MaxWorkers int `json:"max_workers"`
+ ActiveWorkers int `json:"active_workers"`
+ IdleWorkers int `json:"idle_workers"`
+ QueueSize int `json:"queue_size"`
+ QueueLength int `json:"queue_length"`
+ QueueUsagePercent float64 `json:"queue_usage_percent"`
+ Enqueued int64 `json:"enqueued"`
+ Dropped int64 `json:"dropped"`
+ Processed int64 `json:"processed"`
+ Errors int64 `json:"errors"`
+ PreBlockActive int `json:"pre_block_active"`
+ PreBlockChecked int64 `json:"pre_block_checked"`
+ PreBlockAllowed int64 `json:"pre_block_allowed"`
+ PreBlockBlocked int64 `json:"pre_block_blocked"`
+ PreBlockErrors int64 `json:"pre_block_errors"`
+ PreBlockAvgLatencyMS int64 `json:"pre_block_avg_latency_ms"`
+ PreBlockAPIKeyActive int64 `json:"pre_block_api_key_active"`
+ PreBlockAPIKeyAvailableCount int64 `json:"pre_block_api_key_available_count"`
+ PreBlockAPIKeyTotalCalls int64 `json:"pre_block_api_key_total_calls"`
+ PreBlockAPIKeyLoads []ContentModerationAPIKeyLoad `json:"pre_block_api_key_loads"`
+ APIKeyStatuses []ContentModerationAPIKeyStatus `json:"api_key_statuses"`
+ FlaggedHashCount int64 `json:"flagged_hash_count"`
+ LastCleanupAt *time.Time `json:"last_cleanup_at,omitempty"`
+ LastCleanupDeletedHit int64 `json:"last_cleanup_deleted_hit"`
+ LastCleanupDeletedNonHit int64 `json:"last_cleanup_deleted_non_hit"`
}
type ContentModerationUnbanUserResult struct {
@@ -466,6 +490,12 @@ type ContentModerationService struct {
asyncDropped atomic.Int64
asyncProcessed atomic.Int64
asyncErrors atomic.Int64
+ preBlockActive atomic.Int64
+ preBlockChecked atomic.Int64
+ preBlockAllowed atomic.Int64
+ preBlockBlocked atomic.Int64
+ preBlockErrors atomic.Int64
+ preBlockLatencyTotalMS atomic.Int64
lastCleanupUnix atomic.Int64
lastCleanupDeletedHit atomic.Int64
lastCleanupDeletedNonHit atomic.Int64
@@ -474,10 +504,14 @@ type ContentModerationService struct {
}
type contentModerationTask struct {
- input ContentModerationCheckInput
- content ContentModerationInput
- inputHash string
- enqueuedAt time.Time
+ input ContentModerationCheckInput
+ content ContentModerationInput
+ inputHash string
+ log *ContentModerationLog
+ config *ContentModerationConfig
+ recordHash bool
+ applySideEffects bool
+ enqueuedAt time.Time
}
type contentModerationKeyHealth struct {
@@ -491,6 +525,11 @@ type contentModerationKeyHealth struct {
LastLatencyMS int
LastHTTPStatus int
LastTested bool
+ SyncActive int64
+ SyncTotal int64
+ SyncSuccess int64
+ SyncErrors int64
+ SyncLatencyMS int64
}
func NewContentModerationService(
@@ -827,9 +866,11 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer
"protocol", input.Protocol,
"text_runes", len([]rune(content.Text)),
"image_count", len(content.Images))
+ hashText := content.Hash()
if cfg.Mode == ContentModerationModePreBlock {
if cfg.KeywordBlockingMode != ContentModerationKeywordModeAPIOnly && len(cfg.BlockedKeywords) > 0 {
if keyword, hit := matchBlockedKeyword(content.Text, cfg.BlockedKeywords); hit {
+ s.recordPreBlockSyncMetric(0, ContentModerationActionKeywordBlock)
slog.Info("content_moderation.keyword_block",
"user_id", input.UserID,
"api_key_id", input.APIKeyID,
@@ -840,8 +881,7 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer
"keyword", keyword)
scores := map[string]float64{contentModerationKeywordCategory: 1.0}
log := s.buildLog(input, cfg, ContentModerationActionKeywordBlock, true, contentModerationKeywordCategory, 1.0, scores, content.ExcerptText(), nil, nil, "")
- s.applyFlaggedSideEffects(ctx, cfg, log)
- _ = s.repo.CreateLog(ctx, log)
+ s.enqueueRecord(input, cfg, log, hashText, false, true)
return &ContentModerationDecision{
Allowed: false,
Blocked: true,
@@ -856,6 +896,7 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer
}
}
if cfg.KeywordBlockingMode == ContentModerationKeywordModeKeywordOnly {
+ s.recordPreBlockSyncMetric(0, ContentModerationActionAllow)
slog.Info("content_moderation.skip_api_keyword_only",
"user_id", input.UserID,
"api_key_id", input.APIKeyID,
@@ -865,13 +906,15 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer
return allow, nil
}
}
- hashText := content.Hash()
if cfg.PreHashCheckEnabled && s.hashCache != nil {
matched, err := s.hashCache.HasFlaggedInputHash(ctx, hashText)
if err != nil {
slog.Warn("content_moderation.hash_check_failed", "user_id", input.UserID, "endpoint", input.Endpoint, "error", err)
}
if matched {
+ if cfg.Mode == ContentModerationModePreBlock {
+ s.recordPreBlockSyncMetric(0, ContentModerationActionHashBlock)
+ }
slog.Info("content_moderation.hash_block",
"user_id", input.UserID,
"api_key_id", input.APIKeyID,
@@ -883,6 +926,9 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer
if message != "" {
message = fmt.Sprintf("%s(hash: %s)", message, hashText)
}
+ scores := map[string]float64{"hash": 1.0}
+ log := s.buildLog(input, cfg, ContentModerationActionHashBlock, true, "hash", 1.0, scores, content.ExcerptText(), nil, nil, "")
+ s.enqueueRecord(input, cfg, log, hashText, false, false)
return &ContentModerationDecision{
Allowed: false,
Blocked: true,
@@ -895,6 +941,9 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer
}
}
if !cfg.shouldSample(hashText) {
+ if cfg.Mode == ContentModerationModePreBlock {
+ s.recordPreBlockSyncMetric(0, ContentModerationActionAllow)
+ }
slog.Info("content_moderation.skip_sample_rate",
"user_id", input.UserID,
"api_key_id", input.APIKeyID,
@@ -905,6 +954,9 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer
return allow, nil
}
if len(cfg.apiKeys()) == 0 {
+ if cfg.Mode == ContentModerationModePreBlock {
+ s.recordPreBlockSyncMetric(0, ContentModerationActionError)
+ }
slog.Warn("content_moderation.skip_no_audit_api_keys",
"user_id", input.UserID,
"api_key_id", input.APIKeyID,
@@ -930,10 +982,18 @@ func (s *ContentModerationService) Check(ctx context.Context, input ContentModer
func (s *ContentModerationService) checkSync(ctx context.Context, input ContentModerationCheckInput, cfg *ContentModerationConfig, content ContentModerationInput, hashText string, queueDelay *int, allowBlock bool) *ContentModerationDecision {
allow := &ContentModerationDecision{Allowed: true, Action: ContentModerationActionAllow}
+ trackPreBlock := queueDelay == nil && allowBlock && cfg != nil && cfg.Mode == ContentModerationModePreBlock
+ if trackPreBlock {
+ s.preBlockActive.Add(1)
+ defer s.preBlockActive.Add(-1)
+ }
start := time.Now()
- result, err := s.callModeration(ctx, cfg, content.ModerationInput())
+ result, err := s.callModeration(ctx, cfg, content.ModerationInput(), trackPreBlock)
latency := int(time.Since(start).Milliseconds())
if err != nil {
+ if trackPreBlock {
+ s.recordPreBlockSyncMetric(latency, ContentModerationActionError)
+ }
slog.Warn("content_moderation.audit_api_failed",
"user_id", input.UserID,
"api_key_id", input.APIKeyID,
@@ -962,6 +1022,9 @@ func (s *ContentModerationService) checkSync(ctx context.Context, input ContentM
action = ContentModerationActionBlock
blocked = true
}
+ if trackPreBlock {
+ s.recordPreBlockSyncMetric(latency, action)
+ }
slog.Info("content_moderation.audit_result",
"user_id", input.UserID,
"api_key_id", input.APIKeyID,
@@ -980,13 +1043,11 @@ func (s *ContentModerationService) checkSync(ctx context.Context, input ContentM
"queue_delay_ms", queueDelay)
if flagged || cfg.RecordNonHits {
log := s.buildLog(input, cfg, action, flagged, highestCategory, highestScore, result.CategoryScores, content.ExcerptText(), &latency, queueDelay, "")
- if flagged && s.hashCache != nil {
- if err := s.hashCache.RecordFlaggedInputHash(ctx, hashText); err != nil {
- slog.Warn("content_moderation.record_hash_failed", "user_id", input.UserID, "endpoint", input.Endpoint, "error", err)
- }
+ if queueDelay == nil && cfg.Mode == ContentModerationModePreBlock {
+ s.enqueueRecord(input, cfg, log, hashText, flagged, flagged)
+ } else {
+ s.persistContentModerationLog(ctx, cfg, log, hashText, flagged, flagged)
}
- s.applyFlaggedSideEffects(ctx, cfg, log)
- _ = s.repo.CreateLog(ctx, log)
}
if blocked {
return &ContentModerationDecision{
@@ -1012,6 +1073,25 @@ func (s *ContentModerationService) checkSync(ctx context.Context, input ContentM
}
}
+func (s *ContentModerationService) recordPreBlockSyncMetric(latencyMS int, action string) {
+ if s == nil {
+ return
+ }
+ s.preBlockChecked.Add(1)
+ if latencyMS < 0 {
+ latencyMS = 0
+ }
+ s.preBlockLatencyTotalMS.Add(int64(latencyMS))
+ switch action {
+ case ContentModerationActionBlock, ContentModerationActionHashBlock, ContentModerationActionKeywordBlock:
+ s.preBlockBlocked.Add(1)
+ case ContentModerationActionError:
+ s.preBlockErrors.Add(1)
+ default:
+ s.preBlockAllowed.Add(1)
+ }
+}
+
func (s *ContentModerationService) enqueueAsync(input ContentModerationCheckInput, cfg *ContentModerationConfig, content ContentModerationInput, hashText string) {
if s == nil || s.asyncQueue == nil {
return
@@ -1040,11 +1120,49 @@ func (s *ContentModerationService) enqueueAsync(input ContentModerationCheckInpu
}
}
+func (s *ContentModerationService) enqueueRecord(input ContentModerationCheckInput, cfg *ContentModerationConfig, log *ContentModerationLog, inputHash string, recordHash bool, applySideEffects bool) {
+ if s == nil || s.asyncQueue == nil || log == nil {
+ return
+ }
+ queueSize := defaultContentModerationQueueSize
+ if cfg != nil && cfg.QueueSize > 0 {
+ queueSize = cfg.QueueSize
+ }
+ if len(s.asyncQueue) >= queueSize {
+ slog.Warn("content_moderation.record_queue_full",
+ "user_id", input.UserID,
+ "endpoint", input.Endpoint,
+ "action", log.Action,
+ "queue_size", queueSize)
+ s.asyncDropped.Add(1)
+ return
+ }
+ task := contentModerationTask{
+ input: input,
+ inputHash: inputHash,
+ log: log,
+ config: cloneContentModerationConfig(cfg),
+ recordHash: recordHash,
+ applySideEffects: applySideEffects,
+ enqueuedAt: time.Now(),
+ }
+ select {
+ case s.asyncQueue <- task:
+ s.asyncEnqueued.Add(1)
+ default:
+ slog.Warn("content_moderation.record_queue_full",
+ "user_id", input.UserID,
+ "endpoint", input.Endpoint,
+ "action", log.Action)
+ s.asyncDropped.Add(1)
+ }
+}
+
func (s *ContentModerationService) worker(id int) {
for {
ctx, cancel := context.WithTimeout(context.Background(), maxContentModerationTimeoutMS*time.Millisecond+10*time.Second)
cfg, err := s.loadConfig(ctx)
- if err != nil || !cfg.Enabled || cfg.Mode == ContentModerationModeOff || len(cfg.apiKeys()) == 0 || id >= cfg.WorkerCount {
+ if err != nil || id >= cfg.WorkerCount {
cancel()
time.Sleep(time.Second)
continue
@@ -1061,6 +1179,22 @@ func (s *ContentModerationService) worker(id int) {
slog.Error("content_moderation.worker_panic", "worker_id", id, "recover", r)
}
}()
+ if task.log != nil {
+ s.asyncActive.Add(1)
+ defer s.asyncActive.Add(-1)
+ queueDelay := int(time.Since(task.enqueuedAt).Milliseconds())
+ task.log.QueueDelayMS = &queueDelay
+ taskCfg := task.config
+ if taskCfg == nil {
+ taskCfg = cfg
+ }
+ s.persistContentModerationLog(ctx, taskCfg, task.log, task.inputHash, task.recordHash, task.applySideEffects)
+ s.asyncProcessed.Add(1)
+ return
+ }
+ if !cfg.Enabled || cfg.Mode == ContentModerationModeOff || len(cfg.apiKeys()) == 0 {
+ return
+ }
if !cfg.includesGroup(task.input.GroupID) {
return
}
@@ -1186,6 +1320,15 @@ func (s *ContentModerationService) GetStatus(ctx context.Context) (*ContentModer
if active > cfg.WorkerCount {
active = cfg.WorkerCount
}
+ preBlockActive := int(s.preBlockActive.Load())
+ if preBlockActive < 0 {
+ preBlockActive = 0
+ }
+ preBlockChecked := s.preBlockChecked.Load()
+ preBlockAvgLatency := int64(0)
+ if preBlockChecked > 0 {
+ preBlockAvgLatency = s.preBlockLatencyTotalMS.Load() / preBlockChecked
+ }
queueLength := 0
if s.asyncQueue != nil {
queueLength = len(s.asyncQueue)
@@ -1208,25 +1351,35 @@ func (s *ContentModerationService) GetStatus(ctx context.Context) (*ContentModer
lastCleanupAt = &t
}
return &ContentModerationRuntimeStatus{
- Enabled: cfg.Enabled,
- RiskControlEnabled: riskEnabled,
- Mode: cfg.Mode,
- WorkerCount: cfg.WorkerCount,
- MaxWorkers: maxContentModerationWorkerCount,
- ActiveWorkers: active,
- IdleWorkers: cfg.WorkerCount - active,
- QueueSize: cfg.QueueSize,
- QueueLength: queueLength,
- QueueUsagePercent: queueUsage,
- Enqueued: s.asyncEnqueued.Load(),
- Dropped: s.asyncDropped.Load(),
- Processed: s.asyncProcessed.Load(),
- Errors: s.asyncErrors.Load(),
- APIKeyStatuses: s.apiKeyStatuses(cfg.apiKeys()),
- FlaggedHashCount: flaggedHashCount,
- LastCleanupAt: lastCleanupAt,
- LastCleanupDeletedHit: s.lastCleanupDeletedHit.Load(),
- LastCleanupDeletedNonHit: s.lastCleanupDeletedNonHit.Load(),
+ Enabled: cfg.Enabled,
+ RiskControlEnabled: riskEnabled,
+ Mode: cfg.Mode,
+ WorkerCount: cfg.WorkerCount,
+ MaxWorkers: maxContentModerationWorkerCount,
+ ActiveWorkers: active,
+ IdleWorkers: cfg.WorkerCount - active,
+ QueueSize: cfg.QueueSize,
+ QueueLength: queueLength,
+ QueueUsagePercent: queueUsage,
+ Enqueued: s.asyncEnqueued.Load(),
+ Dropped: s.asyncDropped.Load(),
+ Processed: s.asyncProcessed.Load(),
+ Errors: s.asyncErrors.Load(),
+ PreBlockActive: preBlockActive,
+ PreBlockChecked: preBlockChecked,
+ PreBlockAllowed: s.preBlockAllowed.Load(),
+ PreBlockBlocked: s.preBlockBlocked.Load(),
+ PreBlockErrors: s.preBlockErrors.Load(),
+ PreBlockAvgLatencyMS: preBlockAvgLatency,
+ PreBlockAPIKeyActive: s.preBlockAPIKeyActive(cfg.apiKeys()),
+ PreBlockAPIKeyAvailableCount: s.preBlockAPIKeyAvailableCount(cfg.apiKeys()),
+ PreBlockAPIKeyTotalCalls: s.preBlockAPIKeyTotalCalls(cfg.apiKeys()),
+ PreBlockAPIKeyLoads: s.preBlockAPIKeyLoads(cfg.apiKeys()),
+ APIKeyStatuses: s.apiKeyStatuses(cfg.apiKeys()),
+ FlaggedHashCount: flaggedHashCount,
+ LastCleanupAt: lastCleanupAt,
+ LastCleanupDeletedHit: s.lastCleanupDeletedHit.Load(),
+ LastCleanupDeletedNonHit: s.lastCleanupDeletedNonHit.Load(),
}, nil
}
@@ -1325,7 +1478,7 @@ func (s *ContentModerationService) validateConfig(ctx context.Context, cfg *Cont
return nil
}
-func (s *ContentModerationService) callModeration(ctx context.Context, cfg *ContentModerationConfig, input any) (*moderationAPIResult, error) {
+func (s *ContentModerationService) callModeration(ctx context.Context, cfg *ContentModerationConfig, input any, trackKeyLoad ...bool) (*moderationAPIResult, error) {
attempts := cfg.RetryCount + 1
if attempts <= 0 {
attempts = 1
@@ -1333,6 +1486,7 @@ func (s *ContentModerationService) callModeration(ctx context.Context, cfg *Cont
if attempts > maxContentModerationRetryCount+1 {
attempts = maxContentModerationRetryCount + 1
}
+ trackLoad := len(trackKeyLoad) > 0 && trackKeyLoad[0]
var lastErr error
for attempt := 0; attempt < attempts; attempt++ {
key, ok := s.nextUsableAPIKey(cfg)
@@ -1340,14 +1494,23 @@ func (s *ContentModerationService) callModeration(ctx context.Context, cfg *Cont
lastErr = errors.New("no moderation api key available")
break
}
+ if trackLoad {
+ s.beginModerationAPIKeyCall(key)
+ }
start := time.Now()
httpStatus := 0
result, err := s.callModerationOnceWithInput(ctx, cfg, key, input, &httpStatus)
latency := int(time.Since(start).Milliseconds())
if err == nil {
+ if trackLoad {
+ s.finishModerationAPIKeyCall(key, latency, true)
+ }
s.markAPIKeySuccess(key, latency, httpStatus)
return result, nil
}
+ if trackLoad {
+ s.finishModerationAPIKeyCall(key, latency, false)
+ }
s.markAPIKeyError(key, err.Error(), latency, httpStatus)
lastErr = err
if httpStatus == http.StatusBadRequest {
@@ -1452,10 +1615,32 @@ func (s *ContentModerationService) buildLog(input ContentModerationCheckInput, c
}
}
-func (s *ContentModerationService) applyFlaggedSideEffects(ctx context.Context, cfg *ContentModerationConfig, log *ContentModerationLog) {
- if s == nil || cfg == nil || log == nil || !log.Flagged || log.UserID == nil || *log.UserID <= 0 {
+func (s *ContentModerationService) persistContentModerationLog(ctx context.Context, cfg *ContentModerationConfig, log *ContentModerationLog, hashText string, recordHash bool, applySideEffects bool) {
+ if s == nil || log == nil {
return
}
+ if recordHash && s.hashCache != nil {
+ if err := s.hashCache.RecordFlaggedInputHash(ctx, hashText); err != nil {
+ slog.Warn("content_moderation.record_hash_failed", "user_id", contentModerationEmailUserID(log), "endpoint", log.Endpoint, "error", err)
+ }
+ }
+ autoBanJustApplied := false
+ if applySideEffects {
+ autoBanJustApplied = s.applyFlaggedAccountSideEffects(ctx, cfg, log)
+ s.sendFlaggedNotificationSideEffects(ctx, cfg, log, autoBanJustApplied)
+ }
+ if s.repo != nil {
+ if err := s.repo.CreateLog(ctx, log); err != nil {
+ slog.Warn("content_moderation.create_log_failed", "user_id", contentModerationEmailUserID(log), "endpoint", log.Endpoint, "action", log.Action, "error", err)
+ return
+ }
+ }
+}
+
+func (s *ContentModerationService) applyFlaggedAccountSideEffects(ctx context.Context, cfg *ContentModerationConfig, log *ContentModerationLog) bool {
+ if s == nil || cfg == nil || log == nil || !log.Flagged || log.UserID == nil || *log.UserID <= 0 {
+ return false
+ }
count := 1
if s.repo != nil && cfg.ViolationWindowHours > 0 {
since := time.Now().Add(-time.Duration(cfg.ViolationWindowHours) * time.Hour)
@@ -1469,13 +1654,18 @@ func (s *ContentModerationService) applyFlaggedSideEffects(ctx context.Context,
user, err := s.userRepo.GetByID(ctx, *log.UserID)
if err != nil {
slog.Warn("content_moderation.ban_get_user_failed", "user_id", *log.UserID, "error", err)
- return
+ return false
+ }
+ if user.IsAdmin() {
+ slog.Warn("content_moderation.autoban_skipped_admin", "user_id", *log.UserID, "role", user.Role, "count", count, "threshold", cfg.BanThreshold)
+ // TODO: Disable the triggering API key instead when API key mutation is available here.
+ return false
}
if user.Status != StatusDisabled {
user.Status = StatusDisabled
if err := s.userRepo.Update(ctx, user); err != nil {
slog.Warn("content_moderation.ban_update_user_failed", "user_id", *log.UserID, "error", err)
- return
+ return false
}
if s.authCacheInvalidator != nil {
s.authCacheInvalidator.InvalidateAuthCacheByUserID(ctx, *log.UserID)
@@ -1484,7 +1674,13 @@ func (s *ContentModerationService) applyFlaggedSideEffects(ctx context.Context,
}
log.AutoBanned = true
}
+ return autoBanJustApplied
+}
+func (s *ContentModerationService) sendFlaggedNotificationSideEffects(ctx context.Context, cfg *ContentModerationConfig, log *ContentModerationLog, autoBanJustApplied bool) {
+ if s == nil || cfg == nil || log == nil || !log.Flagged {
+ return
+ }
if s.emailService == nil || strings.TrimSpace(log.UserEmail) == "" {
return
}
@@ -1642,6 +1838,22 @@ func defaultContentModerationConfig() *ContentModerationConfig {
}
}
+func cloneContentModerationConfig(cfg *ContentModerationConfig) *ContentModerationConfig {
+ if cfg == nil {
+ return nil
+ }
+ clone := *cfg
+ clone.APIKeys = append([]string(nil), cfg.APIKeys...)
+ clone.GroupIDs = append([]int64(nil), cfg.GroupIDs...)
+ clone.BlockedKeywords = append([]string(nil), cfg.BlockedKeywords...)
+ clone.Thresholds = cloneFloatMap(cfg.Thresholds)
+ clone.ModelFilter = ContentModerationModelFilter{
+ Type: cfg.ModelFilter.Type,
+ Models: append([]string(nil), cfg.ModelFilter.Models...),
+ }
+ return &clone
+}
+
func (cfg *ContentModerationConfig) normalize() {
if cfg.APIKey != "" {
cfg.APIKeys = normalizeModerationAPIKeys(append(cfg.APIKeys, cfg.APIKey))
@@ -1807,6 +2019,40 @@ func (s *ContentModerationService) isAPIKeyFrozen(key string, now time.Time) boo
return state != nil && state.FrozenUntil.After(now)
}
+func (s *ContentModerationService) beginModerationAPIKeyCall(key string) {
+ hash := moderationAPIKeyHash(key)
+ if hash == "" || s == nil {
+ return
+ }
+ s.keyHealthMu.Lock()
+ defer s.keyHealthMu.Unlock()
+ state := s.ensureAPIKeyHealthLocked(hash, maskSecretTail(key))
+ state.SyncActive++
+}
+
+func (s *ContentModerationService) finishModerationAPIKeyCall(key string, latencyMS int, success bool) {
+ hash := moderationAPIKeyHash(key)
+ if hash == "" || s == nil {
+ return
+ }
+ if latencyMS < 0 {
+ latencyMS = 0
+ }
+ s.keyHealthMu.Lock()
+ defer s.keyHealthMu.Unlock()
+ state := s.ensureAPIKeyHealthLocked(hash, maskSecretTail(key))
+ if state.SyncActive > 0 {
+ state.SyncActive--
+ }
+ state.SyncTotal++
+ state.SyncLatencyMS += int64(latencyMS)
+ if success {
+ state.SyncSuccess++
+ return
+ }
+ state.SyncErrors++
+}
+
func (s *ContentModerationService) markAPIKeySuccess(key string, latencyMS int, httpStatus int) {
hash := moderationAPIKeyHash(key)
if hash == "" || s == nil {
@@ -1926,6 +2172,71 @@ func (s *ContentModerationService) apiKeyStatuses(keys []string) []ContentModera
return out
}
+func (s *ContentModerationService) preBlockAPIKeyLoads(keys []string) []ContentModerationAPIKeyLoad {
+ out := make([]ContentModerationAPIKeyLoad, 0, len(keys))
+ for idx, key := range keys {
+ out = append(out, s.preBlockAPIKeyLoadForHash(idx, moderationAPIKeyHash(key), maskSecretTail(key)))
+ }
+ return out
+}
+
+func (s *ContentModerationService) preBlockAPIKeyActive(keys []string) int64 {
+ var total int64
+ for _, item := range s.preBlockAPIKeyLoads(keys) {
+ total += item.Active
+ }
+ return total
+}
+
+func (s *ContentModerationService) preBlockAPIKeyAvailableCount(keys []string) int64 {
+ now := time.Now()
+ var count int64
+ for _, key := range keys {
+ if !s.isAPIKeyFrozen(key, now) {
+ count++
+ }
+ }
+ return count
+}
+
+func (s *ContentModerationService) preBlockAPIKeyTotalCalls(keys []string) int64 {
+ var total int64
+ for _, item := range s.preBlockAPIKeyLoads(keys) {
+ total += item.Total
+ }
+ return total
+}
+
+func (s *ContentModerationService) preBlockAPIKeyLoadForHash(index int, hash string, masked string) ContentModerationAPIKeyLoad {
+ load := ContentModerationAPIKeyLoad{
+ Index: index,
+ KeyHash: hash,
+ Masked: masked,
+ Status: "unknown",
+ }
+ status := s.apiKeyStatusForHash(index, hash, masked, true)
+ load.Status = status.Status
+ load.LastLatencyMS = status.LastLatencyMS
+ load.LastHTTPStatus = status.LastHTTPStatus
+ if hash == "" || s == nil {
+ return load
+ }
+ s.keyHealthMu.Lock()
+ defer s.keyHealthMu.Unlock()
+ state := s.keyHealth[hash]
+ if state == nil {
+ return load
+ }
+ load.Active = state.SyncActive
+ load.Total = state.SyncTotal
+ load.Success = state.SyncSuccess
+ load.Errors = state.SyncErrors
+ if state.SyncTotal > 0 {
+ load.AvgLatencyMS = state.SyncLatencyMS / state.SyncTotal
+ }
+ return load
+}
+
func (s *ContentModerationService) apiKeyStatusForHash(index int, hash string, masked string, configured bool) ContentModerationAPIKeyStatus {
status := ContentModerationAPIKeyStatus{
Index: index,
diff --git a/backend/internal/service/content_moderation_test.go b/backend/internal/service/content_moderation_test.go
index 20fce3ec..9cfdc1e4 100644
--- a/backend/internal/service/content_moderation_test.go
+++ b/backend/internal/service/content_moderation_test.go
@@ -1,11 +1,15 @@
package service
import (
+ "bytes"
"context"
"encoding/json"
+ "fmt"
+ "log/slog"
"net/http"
"net/http/httptest"
"strings"
+ "sync"
"testing"
"time"
@@ -73,10 +77,13 @@ func (r *contentModerationTestSettingRepo) Delete(ctx context.Context, key strin
}
type contentModerationTestRepo struct {
+ mu sync.Mutex
logs []ContentModerationLog
}
func (r *contentModerationTestRepo) CreateLog(ctx context.Context, log *ContentModerationLog) error {
+ r.mu.Lock()
+ defer r.mu.Unlock()
if log != nil {
r.logs = append(r.logs, *log)
}
@@ -88,14 +95,55 @@ func (r *contentModerationTestRepo) ListLogs(ctx context.Context, filter Content
}
func (r *contentModerationTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time) (int, error) {
- return 0, nil
+ r.mu.Lock()
+ defer r.mu.Unlock()
+ count := 0
+ for _, log := range r.logs {
+ if log.UserID == nil || *log.UserID != userID || !log.Flagged || log.Action == ContentModerationActionHashBlock {
+ continue
+ }
+ if log.CreatedAt.IsZero() || log.CreatedAt.Before(since) {
+ continue
+ }
+ count++
+ }
+ return count, nil
}
func (r *contentModerationTestRepo) CleanupExpiredLogs(ctx context.Context, hitBefore time.Time, nonHitBefore time.Time) (*ContentModerationCleanupResult, error) {
return &ContentModerationCleanupResult{}, nil
}
+func (r *contentModerationTestRepo) snapshotLogs() []ContentModerationLog {
+ r.mu.Lock()
+ defer r.mu.Unlock()
+ out := make([]ContentModerationLog, len(r.logs))
+ copy(out, r.logs)
+ return out
+}
+
+func requireContentModerationLogCount(t *testing.T, repo *contentModerationTestRepo, want int) []ContentModerationLog {
+ t.Helper()
+ var logs []ContentModerationLog
+ require.Eventually(t, func() bool {
+ logs = repo.snapshotLogs()
+ return len(logs) == want
+ }, time.Second, 10*time.Millisecond)
+ return logs
+}
+
+func requireRecordedHashCount(t *testing.T, cache *contentModerationTestHashCache, want int) []string {
+ t.Helper()
+ var hashes []string
+ require.Eventually(t, func() bool {
+ hashes = cache.snapshotRecorded()
+ return len(hashes) == want
+ }, time.Second, 10*time.Millisecond)
+ return hashes
+}
+
type contentModerationTestHashCache struct {
+ mu sync.Mutex
hashes map[string]struct{}
recorded []string
checked []string
@@ -231,6 +279,10 @@ func (r *contentModerationTestUserRepo) DisableTotp(ctx context.Context, userID
panic("unexpected DisableTotp call")
}
+func (r *contentModerationTestUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return r.GetByID(ctx, id)
+}
+
type contentModerationTestAuthCacheInvalidator struct {
userIDs []int64
}
@@ -246,6 +298,8 @@ func (i *contentModerationTestAuthCacheInvalidator) InvalidateAuthCacheByGroupID
}
func (c *contentModerationTestHashCache) RecordFlaggedInputHash(ctx context.Context, inputHash string) error {
+ c.mu.Lock()
+ defer c.mu.Unlock()
if c.hashes == nil {
c.hashes = map[string]struct{}{}
}
@@ -255,6 +309,8 @@ func (c *contentModerationTestHashCache) RecordFlaggedInputHash(ctx context.Cont
}
func (c *contentModerationTestHashCache) HasFlaggedInputHash(ctx context.Context, inputHash string) (bool, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
c.checked = append(c.checked, inputHash)
if c.hasResultUsed {
return c.hasResult, nil
@@ -264,6 +320,8 @@ func (c *contentModerationTestHashCache) HasFlaggedInputHash(ctx context.Context
}
func (c *contentModerationTestHashCache) DeleteFlaggedInputHash(ctx context.Context, inputHash string) (bool, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
c.deleted = append(c.deleted, inputHash)
if c.hashes == nil {
return false, nil
@@ -276,15 +334,50 @@ func (c *contentModerationTestHashCache) DeleteFlaggedInputHash(ctx context.Cont
}
func (c *contentModerationTestHashCache) ClearFlaggedInputHashes(ctx context.Context) (int64, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
deleted := int64(len(c.hashes))
c.hashes = map[string]struct{}{}
return deleted, nil
}
func (c *contentModerationTestHashCache) CountFlaggedInputHashes(ctx context.Context) (int64, error) {
+ c.mu.Lock()
+ defer c.mu.Unlock()
return int64(len(c.hashes)), nil
}
+func (c *contentModerationTestHashCache) snapshotRecorded() []string {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ out := make([]string, len(c.recorded))
+ copy(out, c.recorded)
+ return out
+}
+
+func (c *contentModerationTestHashCache) snapshotChecked() []string {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ out := make([]string, len(c.checked))
+ copy(out, c.checked)
+ return out
+}
+
+func (c *contentModerationTestHashCache) hasHash(inputHash string) bool {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ _, ok := c.hashes[inputHash]
+ return ok
+}
+
+func (c *contentModerationTestHashCache) snapshotDeleted() []string {
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ out := make([]string, len(c.deleted))
+ copy(out, c.deleted)
+ return out
+}
+
func TestBuildContentModerationLog_RedactsInputExcerpt(t *testing.T) {
svc := &ContentModerationService{}
cfg := defaultContentModerationConfig()
@@ -381,10 +474,10 @@ func TestContentModerationCheck_PreBlockKeywordHitSkipsUpstreamCall(t *testing.T
require.True(t, decision.Blocked)
require.Equal(t, ContentModerationActionKeywordBlock, decision.Action)
require.False(t, upstreamCalled, "keyword block must short-circuit upstream moderation call")
- require.Len(t, repo.logs, 1)
- require.True(t, repo.logs[0].Flagged)
- require.Equal(t, ContentModerationActionKeywordBlock, repo.logs[0].Action)
- require.Equal(t, contentModerationKeywordCategory, repo.logs[0].HighestCategory)
+ logs := requireContentModerationLogCount(t, repo, 1)
+ require.True(t, logs[0].Flagged)
+ require.Equal(t, ContentModerationActionKeywordBlock, logs[0].Action)
+ require.Equal(t, contentModerationKeywordCategory, logs[0].HighestCategory)
}
func TestContentModerationCheck_KeywordsIgnoredInObserveMode(t *testing.T) {
@@ -474,7 +567,7 @@ func TestContentModerationCheck_KeywordOnlyStrategySkipsAPIOnMiss(t *testing.T)
require.NoError(t, err)
require.True(t, decision.Allowed, "keyword-only must allow misses without calling the API")
require.False(t, upstreamCalled, "keyword-only must not call the upstream moderation API")
- require.Len(t, repo.logs, 0)
+ require.Len(t, repo.snapshotLogs(), 0)
}
func TestContentModerationCheck_APIOnlyStrategyIgnoresKeywordList(t *testing.T) {
@@ -545,7 +638,7 @@ func TestContentModerationCheck_ModelFilterAllAuditsEveryModel(t *testing.T) {
require.True(t, decision.Blocked)
require.Equal(t, ContentModerationActionKeywordBlock, decision.Action)
}
- require.Len(t, repo.logs, 2)
+ requireContentModerationLogCount(t, repo, 2)
}
func TestContentModerationCheck_ModelFilterIncludeOnlyAuditsListedModels(t *testing.T) {
@@ -571,8 +664,8 @@ func TestContentModerationCheck_ModelFilterIncludeOnlyAuditsListedModels(t *test
require.True(t, decision.Allowed)
require.False(t, decision.Blocked)
require.Equal(t, ContentModerationActionAllow, decision.Action)
- require.Len(t, repo.logs, 1)
- require.Equal(t, "gpt-5.5", repo.logs[0].Model)
+ logs := requireContentModerationLogCount(t, repo, 1)
+ require.Equal(t, "gpt-5.5", logs[0].Model)
}
func TestContentModerationCheck_ModelFilterExcludeSkipsListedModels(t *testing.T) {
@@ -598,8 +691,8 @@ func TestContentModerationCheck_ModelFilterExcludeSkipsListedModels(t *testing.T
require.True(t, decision.Allowed)
require.False(t, decision.Blocked)
require.Equal(t, ContentModerationActionAllow, decision.Action)
- require.Len(t, repo.logs, 1)
- require.Equal(t, "gpt-5.5", repo.logs[0].Model)
+ logs := requireContentModerationLogCount(t, repo, 1)
+ require.Equal(t, "gpt-5.5", logs[0].Model)
}
func TestContentModerationLoadConfig_LegacyConfigDefaultsModelFilterToAll(t *testing.T) {
@@ -639,8 +732,8 @@ func TestContentModerationCheck_ModelFilterUsesRequestedModelNotBodyModel(t *tes
require.NoError(t, err)
require.True(t, decision.Blocked)
require.Equal(t, ContentModerationActionKeywordBlock, decision.Action)
- require.Len(t, repo.logs, 1)
- require.Equal(t, "gpt-5.5", repo.logs[0].Model)
+ logs := requireContentModerationLogCount(t, repo, 1)
+ require.Equal(t, "gpt-5.5", logs[0].Model)
}
func defaultContentModerationModelFilterTestConfig() *ContentModerationConfig {
@@ -939,11 +1032,11 @@ func TestContentModerationCheck_OpenAIResponsesRecordsNonHitForCodexPayload(t *t
require.NoError(t, err)
require.False(t, decision.Blocked)
- require.Len(t, repo.logs, 1)
- require.False(t, repo.logs[0].Flagged)
- require.Equal(t, ContentModerationActionAllow, repo.logs[0].Action)
- require.Equal(t, "/responses", repo.logs[0].Endpoint)
- require.Equal(t, "last user prompt", repo.logs[0].InputExcerpt)
+ logs := requireContentModerationLogCount(t, repo, 1)
+ require.False(t, logs[0].Flagged)
+ require.Equal(t, ContentModerationActionAllow, logs[0].Action)
+ require.Equal(t, "/responses", logs[0].Endpoint)
+ require.Equal(t, "last user prompt", logs[0].InputExcerpt)
require.Equal(t, "last user prompt", moderationRequest.Input)
}
@@ -1007,14 +1100,164 @@ func TestContentModerationCheck_PreBlockBlocksCodexResponsesLatestUserInput(t *t
require.Equal(t, ContentModerationActionBlock, decision.Action)
require.Equal(t, http.StatusUnavailableForLegalReasons, decision.StatusCode)
require.Equal(t, "内容审计测试阻断", decision.Message)
- require.Len(t, repo.logs, 1)
- require.True(t, repo.logs[0].Flagged)
- require.Equal(t, ContentModerationActionBlock, repo.logs[0].Action)
- require.Equal(t, ContentModerationModePreBlock, repo.logs[0].Mode)
- require.Equal(t, "latest blocked prompt", repo.logs[0].InputExcerpt)
+ logs := requireContentModerationLogCount(t, repo, 1)
+ require.True(t, logs[0].Flagged)
+ require.Equal(t, ContentModerationActionBlock, logs[0].Action)
+ require.Equal(t, ContentModerationModePreBlock, logs[0].Mode)
+ require.Equal(t, "latest blocked prompt", logs[0].InputExcerpt)
require.Equal(t, "latest blocked prompt", moderationRequest.Input)
}
+func TestContentModerationStatusTracksPreBlockSyncMetrics(t *testing.T) {
+ var requestCount int
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ requestCount++
+ score := 0.01
+ if requestCount == 1 {
+ score = 0.9
+ }
+ time.Sleep(5 * time.Millisecond)
+ _ = json.NewEncoder(w).Encode(moderationAPIResponse{
+ Results: []moderationAPIResult{{
+ CategoryScores: map[string]float64{"sexual": score},
+ }},
+ })
+ }))
+ defer server.Close()
+
+ cfg := defaultContentModerationConfig()
+ cfg.Enabled = true
+ cfg.Mode = ContentModerationModePreBlock
+ cfg.BaseURL = server.URL
+ cfg.APIKeys = []string{"sk-test"}
+ rawCfg, err := json.Marshal(cfg)
+ require.NoError(t, err)
+
+ svc := NewContentModerationService(
+ &contentModerationTestSettingRepo{values: map[string]string{
+ SettingKeyRiskControlEnabled: "true",
+ SettingKeyContentModerationConfig: string(rawCfg),
+ }},
+ &contentModerationTestRepo{},
+ &contentModerationTestHashCache{},
+ nil,
+ nil,
+ nil,
+ nil,
+ )
+
+ for _, prompt := range []string{"blocked prompt", "clean prompt"} {
+ _, err := svc.Check(context.Background(), ContentModerationCheckInput{
+ UserID: 1001,
+ Protocol: ContentModerationProtocolOpenAIChat,
+ Body: []byte(fmt.Sprintf(`{"messages":[{"role":"user","content":%q}]}`, prompt)),
+ })
+ require.NoError(t, err)
+ }
+
+ status, err := svc.GetStatus(context.Background())
+ require.NoError(t, err)
+ require.Equal(t, int64(2), status.PreBlockChecked)
+ require.Equal(t, int64(1), status.PreBlockAllowed)
+ require.Equal(t, int64(1), status.PreBlockBlocked)
+ require.Equal(t, int64(0), status.PreBlockErrors)
+ require.Equal(t, 0, status.PreBlockActive)
+ require.GreaterOrEqual(t, status.PreBlockAvgLatencyMS, int64(1))
+}
+
+func TestContentModerationStatusTracksPreBlockAPIKeyLoad(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ _ = json.NewEncoder(w).Encode(moderationAPIResponse{
+ Results: []moderationAPIResult{{
+ CategoryScores: map[string]float64{"sexual": 0.01},
+ }},
+ })
+ }))
+ defer server.Close()
+
+ cfg := defaultContentModerationConfig()
+ cfg.Enabled = true
+ cfg.Mode = ContentModerationModePreBlock
+ cfg.BaseURL = server.URL
+ cfg.APIKeys = []string{"sk-one", "sk-two"}
+ rawCfg, err := json.Marshal(cfg)
+ require.NoError(t, err)
+
+ svc := NewContentModerationService(
+ &contentModerationTestSettingRepo{values: map[string]string{
+ SettingKeyRiskControlEnabled: "true",
+ SettingKeyContentModerationConfig: string(rawCfg),
+ }},
+ &contentModerationTestRepo{},
+ &contentModerationTestHashCache{},
+ nil,
+ nil,
+ nil,
+ nil,
+ )
+
+ for idx := 0; idx < 4; idx++ {
+ _, err := svc.Check(context.Background(), ContentModerationCheckInput{
+ UserID: 1001,
+ Protocol: ContentModerationProtocolOpenAIChat,
+ Body: []byte(fmt.Sprintf(`{"messages":[{"role":"user","content":"prompt %d"}]}`, idx)),
+ })
+ require.NoError(t, err)
+ }
+
+ status, err := svc.GetStatus(context.Background())
+ require.NoError(t, err)
+ require.Len(t, status.PreBlockAPIKeyLoads, 2)
+ require.Equal(t, int64(4), status.PreBlockAPIKeyTotalCalls)
+ require.Equal(t, int64(2), status.PreBlockAPIKeyAvailableCount)
+ require.Equal(t, int64(0), status.PreBlockAPIKeyActive)
+ require.Equal(t, int64(0), status.PreBlockAPIKeyLoads[0].Active)
+ require.Equal(t, int64(2), status.PreBlockAPIKeyLoads[0].Total)
+ require.Equal(t, int64(2), status.PreBlockAPIKeyLoads[0].Success)
+ require.Equal(t, int64(0), status.PreBlockAPIKeyLoads[0].Errors)
+ require.Equal(t, int64(2), status.PreBlockAPIKeyLoads[1].Total)
+ require.Equal(t, int64(2), status.PreBlockAPIKeyLoads[1].Success)
+}
+
+func TestContentModerationStatusTracksPreBlockLocalBlocks(t *testing.T) {
+ cfg := defaultContentModerationConfig()
+ cfg.Enabled = true
+ cfg.Mode = ContentModerationModePreBlock
+ cfg.KeywordBlockingMode = ContentModerationKeywordModeKeywordOnly
+ cfg.BlockedKeywords = []string{"blocked"}
+ rawCfg, err := json.Marshal(cfg)
+ require.NoError(t, err)
+
+ svc := NewContentModerationService(
+ &contentModerationTestSettingRepo{values: map[string]string{
+ SettingKeyRiskControlEnabled: "true",
+ SettingKeyContentModerationConfig: string(rawCfg),
+ }},
+ &contentModerationTestRepo{},
+ &contentModerationTestHashCache{},
+ nil,
+ nil,
+ nil,
+ nil,
+ )
+
+ for _, prompt := range []string{"blocked prompt", "clean prompt"} {
+ _, err := svc.Check(context.Background(), ContentModerationCheckInput{
+ UserID: 1001,
+ Protocol: ContentModerationProtocolOpenAIChat,
+ Body: []byte(fmt.Sprintf(`{"messages":[{"role":"user","content":%q}]}`, prompt)),
+ })
+ require.NoError(t, err)
+ }
+
+ status, err := svc.GetStatus(context.Background())
+ require.NoError(t, err)
+ require.Equal(t, int64(2), status.PreBlockChecked)
+ require.Equal(t, int64(1), status.PreBlockAllowed)
+ require.Equal(t, int64(1), status.PreBlockBlocked)
+ require.Equal(t, int64(0), status.PreBlockErrors)
+}
+
func TestBuildContentModerationTestAuditResult_UsesConfiguredThresholdsOnly(t *testing.T) {
result := buildContentModerationTestAuditResult(&moderationAPIResult{
Flagged: true,
@@ -1137,6 +1380,8 @@ func TestContentModerationCheck_PreHashUsesRedisHashCache(t *testing.T) {
cfg.APIKeys = []string{"sk-test"}
cfg.BlockStatus = http.StatusConflict
cfg.BlockMessage = "命中历史风险输入"
+ cfg.AutoBanEnabled = true
+ cfg.BanThreshold = 1
rawCfg, err := json.Marshal(cfg)
require.NoError(t, err)
@@ -1145,20 +1390,23 @@ func TestContentModerationCheck_PreHashUsesRedisHashCache(t *testing.T) {
content.Normalize()
hashCache.hashes[content.Hash()] = struct{}{}
+ repo := &contentModerationTestRepo{}
+ userRepo := &contentModerationTestUserRepo{user: &User{ID: 1001, Status: StatusActive}}
svc := NewContentModerationService(
&contentModerationTestSettingRepo{values: map[string]string{
SettingKeyRiskControlEnabled: "true",
SettingKeyContentModerationConfig: string(rawCfg),
}},
- &contentModerationTestRepo{},
+ repo,
hashCache,
nil,
- nil,
+ userRepo,
nil,
nil,
)
decision, err := svc.Check(context.Background(), ContentModerationCheckInput{
+ UserID: 1001,
Protocol: ContentModerationProtocolOpenAIChat,
Body: []byte(`{"messages":[{"role":"user","content":"blocked prompt"}]}`),
})
@@ -1169,7 +1417,161 @@ func TestContentModerationCheck_PreHashUsesRedisHashCache(t *testing.T) {
require.Equal(t, content.Hash(), decision.InputHash)
require.Contains(t, decision.Message, "命中历史风险输入")
require.Contains(t, decision.Message, content.Hash())
- require.Len(t, hashCache.checked, 1)
+ require.Len(t, hashCache.snapshotChecked(), 1)
+ logs := requireContentModerationLogCount(t, repo, 1)
+ require.True(t, logs[0].Flagged)
+ require.Equal(t, ContentModerationActionHashBlock, logs[0].Action)
+ require.Equal(t, 1.0, logs[0].CategoryScores["hash"])
+ require.Equal(t, ContentModerationModePreBlock, logs[0].Mode)
+ require.Zero(t, logs[0].ViolationCount)
+ require.False(t, logs[0].AutoBanned)
+ require.Empty(t, userRepo.updated)
+}
+
+func TestContentModerationCheck_HashBlockLogsDoNotIncreaseNextViolationCount(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ _ = json.NewEncoder(w).Encode(moderationAPIResponse{
+ Results: []moderationAPIResult{{
+ CategoryScores: map[string]float64{"sexual": 0.9},
+ }},
+ })
+ }))
+ defer server.Close()
+
+ cfg := defaultContentModerationConfig()
+ cfg.Enabled = true
+ cfg.Mode = ContentModerationModePreBlock
+ cfg.BaseURL = server.URL
+ cfg.APIKeys = []string{"sk-test"}
+ cfg.AutoBanEnabled = false
+ rawCfg, err := json.Marshal(cfg)
+ require.NoError(t, err)
+
+ userID := int64(1001)
+ repo := &contentModerationTestRepo{}
+ hashLog := &ContentModerationLog{
+ UserID: &userID,
+ Action: ContentModerationActionHashBlock,
+ Flagged: true,
+ HighestCategory: "hash",
+ HighestScore: 1,
+ CreatedAt: time.Now(),
+ }
+ require.NoError(t, repo.CreateLog(context.Background(), hashLog))
+
+ svc := NewContentModerationService(
+ &contentModerationTestSettingRepo{values: map[string]string{
+ SettingKeyRiskControlEnabled: "true",
+ SettingKeyContentModerationConfig: string(rawCfg),
+ }},
+ repo,
+ &contentModerationTestHashCache{},
+ nil,
+ nil,
+ nil,
+ nil,
+ )
+
+ decision, err := svc.Check(context.Background(), ContentModerationCheckInput{
+ UserID: userID,
+ Protocol: ContentModerationProtocolOpenAIChat,
+ Body: []byte(`{"messages":[{"role":"user","content":"new blocked prompt"}]}`),
+ })
+
+ require.NoError(t, err)
+ require.True(t, decision.Blocked)
+ logs := requireContentModerationLogCount(t, repo, 2)
+ require.Equal(t, ContentModerationActionHashBlock, logs[0].Action)
+ require.Equal(t, ContentModerationActionBlock, logs[1].Action)
+ require.Equal(t, 1, logs[1].ViolationCount)
+}
+
+func TestContentModerationAutoBanSkipsAdminAccount(t *testing.T) {
+ var slogOutput bytes.Buffer
+ previousLogger := slog.Default()
+ slog.SetDefault(slog.New(slog.NewTextHandler(&slogOutput, nil)))
+ t.Cleanup(func() {
+ slog.SetDefault(previousLogger)
+ })
+
+ cfg := defaultContentModerationConfig()
+ cfg.BanThreshold = 2
+ cfg.ViolationWindowHours = 24
+
+ userID := int64(1001)
+ repo := &contentModerationTestRepo{}
+ require.NoError(t, repo.CreateLog(context.Background(), newContentModerationFlaggedLog(userID)))
+ userRepo := &contentModerationTestUserRepo{user: &User{ID: userID, Role: RoleAdmin, Status: StatusActive}}
+ invalidator := &contentModerationTestAuthCacheInvalidator{}
+ svc := NewContentModerationService(nil, repo, nil, nil, userRepo, invalidator, nil)
+
+ svc.persistContentModerationLog(context.Background(), cfg, newContentModerationFlaggedLog(userID), "", false, true)
+
+ logs := requireContentModerationLogCount(t, repo, 2)
+ require.Equal(t, 2, logs[1].ViolationCount)
+ require.False(t, logs[1].AutoBanned)
+ require.Equal(t, StatusActive, userRepo.user.Status)
+ require.Empty(t, userRepo.updated)
+ require.Empty(t, invalidator.userIDs)
+ require.Contains(t, slogOutput.String(), "content_moderation.autoban_skipped_admin")
+ require.Contains(t, slogOutput.String(), "user_id=1001")
+ require.Contains(t, slogOutput.String(), "role=admin")
+ require.Contains(t, slogOutput.String(), "count=2")
+ require.Contains(t, slogOutput.String(), "threshold=2")
+}
+
+func TestContentModerationAutoBanDisablesRegularUserAtThreshold(t *testing.T) {
+ cfg := defaultContentModerationConfig()
+ cfg.BanThreshold = 2
+ cfg.ViolationWindowHours = 24
+
+ userID := int64(1001)
+ repo := &contentModerationTestRepo{}
+ require.NoError(t, repo.CreateLog(context.Background(), newContentModerationFlaggedLog(userID)))
+ userRepo := &contentModerationTestUserRepo{user: &User{ID: userID, Role: RoleUser, Status: StatusActive}}
+ invalidator := &contentModerationTestAuthCacheInvalidator{}
+ svc := NewContentModerationService(nil, repo, nil, nil, userRepo, invalidator, nil)
+
+ svc.persistContentModerationLog(context.Background(), cfg, newContentModerationFlaggedLog(userID), "", false, true)
+
+ logs := requireContentModerationLogCount(t, repo, 2)
+ require.Equal(t, 2, logs[1].ViolationCount)
+ require.True(t, logs[1].AutoBanned)
+ require.Len(t, userRepo.updated, 1)
+ require.Equal(t, StatusDisabled, userRepo.user.Status)
+ require.Equal(t, []int64{userID}, invalidator.userIDs)
+}
+
+func TestContentModerationAdminBelowBanThresholdRecordsViolationOnly(t *testing.T) {
+ cfg := defaultContentModerationConfig()
+ cfg.BanThreshold = 2
+ cfg.ViolationWindowHours = 24
+
+ userID := int64(1001)
+ repo := &contentModerationTestRepo{}
+ userRepo := &contentModerationTestUserRepo{user: &User{ID: userID, Role: RoleAdmin, Status: StatusActive}}
+ invalidator := &contentModerationTestAuthCacheInvalidator{}
+ svc := NewContentModerationService(nil, repo, nil, nil, userRepo, invalidator, nil)
+
+ svc.persistContentModerationLog(context.Background(), cfg, newContentModerationFlaggedLog(userID), "", false, true)
+
+ logs := requireContentModerationLogCount(t, repo, 1)
+ require.Equal(t, 1, logs[0].ViolationCount)
+ require.False(t, logs[0].AutoBanned)
+ require.Equal(t, StatusActive, userRepo.user.Status)
+ require.Empty(t, userRepo.updated)
+ require.Empty(t, invalidator.userIDs)
+}
+
+func newContentModerationFlaggedLog(userID int64) *ContentModerationLog {
+ return &ContentModerationLog{
+ UserID: &userID,
+ Action: ContentModerationActionBlock,
+ Flagged: true,
+ HighestCategory: "sexual",
+ HighestScore: 0.9,
+ CreatedAt: time.Now(),
+ }
}
func TestContentModerationCheck_PreBlockFlaggedWritesRedisHashCache(t *testing.T) {
@@ -1219,8 +1621,8 @@ func TestContentModerationCheck_PreBlockFlaggedWritesRedisHashCache(t *testing.T
require.True(t, decision.Blocked)
require.Equal(t, ContentModerationActionBlock, decision.Action)
require.Equal(t, 1, requestCount)
- require.Len(t, hashCache.recorded, 1)
- require.Len(t, repo.logs, 1)
+ recorded := requireRecordedHashCount(t, hashCache, 1)
+ requireContentModerationLogCount(t, repo, 1)
decision, err = svc.Check(context.Background(), ContentModerationCheckInput{
Protocol: ContentModerationProtocolOpenAIChat,
@@ -1229,9 +1631,11 @@ func TestContentModerationCheck_PreBlockFlaggedWritesRedisHashCache(t *testing.T
require.NoError(t, err)
require.True(t, decision.Blocked)
require.Equal(t, ContentModerationActionHashBlock, decision.Action)
- require.Equal(t, hashCache.recorded[0], decision.InputHash)
+ require.Equal(t, recorded[0], decision.InputHash)
require.Equal(t, 1, requestCount)
- require.Len(t, repo.logs, 1)
+ logs := requireContentModerationLogCount(t, repo, 2)
+ require.Equal(t, ContentModerationActionBlock, logs[0].Action)
+ require.Equal(t, ContentModerationActionHashBlock, logs[1].Action)
}
func TestContentModerationDeleteFlaggedInputHash_NormalizesAndDeletes(t *testing.T) {
@@ -1246,8 +1650,8 @@ func TestContentModerationDeleteFlaggedInputHash_NormalizesAndDeletes(t *testing
require.NoError(t, err)
require.Equal(t, existingHash, result.InputHash)
require.True(t, result.Deleted)
- require.NotContains(t, hashCache.hashes, existingHash)
- require.Equal(t, []string{existingHash}, hashCache.deleted)
+ require.False(t, hashCache.hasHash(existingHash))
+ require.Equal(t, []string{existingHash}, hashCache.snapshotDeleted())
result, err = svc.DeleteFlaggedInputHash(context.Background(), existingHash)
@@ -1327,8 +1731,8 @@ func TestContentModerationCheck_AsyncFlaggedWritesRedisHashCache(t *testing.T) {
}, cfg, ContentModerationInput{Text: "bad prompt"}, strings.Repeat("b", 64), contentModerationIntPtr(25), false)
require.False(t, decision.Blocked)
- require.Len(t, hashCache.recorded, 1)
- require.Len(t, repo.logs, 1)
+ requireRecordedHashCount(t, hashCache, 1)
+ requireContentModerationLogCount(t, repo, 1)
}
func TestBuildContentModerationAccountDisabledEmailBody_ContainsBanDetails(t *testing.T) {
diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go
index 59c34eaa..11245d00 100644
--- a/backend/internal/service/domain_constants.go
+++ b/backend/internal/service/domain_constants.go
@@ -431,6 +431,9 @@ const (
// 当客户端 UA 被识别为浏览器(Chrome/Firefox/Safari/Edge 等)时,转发给 OpenAI 上游前会替换为此值,
// 用于避免 Cloudflare 对浏览器型 UA 的质询拦截。
SettingKeyOpenAICodexUserAgent = "openai_codex_user_agent"
+ // SettingKeyOpenAIAllowClaudeCodeCodexPlugin 全局开关:是否额外放行 Claude Code 的 Codex 插件(默认 false)。
+ // 仅在账号 codex_cli_only 开启时生效;开启后无需逐账号配置 codex_cli_only_allowed_clients。
+ SettingKeyOpenAIAllowClaudeCodeCodexPlugin = "openai_allow_claude_code_codex_plugin"
// 余额不足提醒
SettingKeyBalanceLowNotifyEnabled = "balance_low_notify_enabled" // 全局开关
@@ -460,3 +463,7 @@ func SettingKeyAuthSourcePlatformQuotas(source string) string {
// AdminAPIKeyPrefix is the prefix for admin API keys (distinct from user "sk-" keys).
const AdminAPIKeyPrefix = "admin-"
+
+// SettingKeyAllowUserViewErrorRequests controls whether end users can view
+// their own failed requests on the usage page. Default false (opt-in).
+const SettingKeyAllowUserViewErrorRequests = "allow_user_view_error_requests"
diff --git a/backend/internal/service/error_policy_test.go b/backend/internal/service/error_policy_test.go
index 297a954c..2aa7a421 100644
--- a/backend/internal/service/error_policy_test.go
+++ b/backend/internal/service/error_policy_test.go
@@ -389,6 +389,60 @@ func TestApplyErrorPolicy(t *testing.T) {
}
}
+func TestApplyErrorPolicy_GeminiRateLimitBypassesCustomSkip(t *testing.T) {
+ repo := &stubAntigravityAccountRepo{}
+ cache := &stubSmartRetryCache{}
+ rlSvc := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
+ svc := &AntigravityGatewayService{
+ rateLimitService: rlSvc,
+ accountRepo: repo,
+ cache: cache,
+ }
+
+ account := &Account{
+ ID: 31,
+ Type: AccountTypeAPIKey,
+ Platform: PlatformAntigravity,
+ Credentials: map[string]any{
+ "custom_error_codes_enabled": true,
+ "custom_error_codes": []any{float64(500)},
+ },
+ }
+ body := []byte(`{
+ "error": {
+ "status": "RESOURCE_EXHAUSTED",
+ "details": [
+ {"@type": "type.googleapis.com/google.rpc.ErrorInfo", "metadata": {"model": "gemini-3-flash"}, "reason": "RATE_LIMIT_EXCEEDED"},
+ {"@type": "type.googleapis.com/google.rpc.RetryInfo", "retryDelay": "15s"}
+ ]
+ }
+ }`)
+ p := antigravityRetryLoopParams{
+ ctx: context.Background(),
+ prefix: "[test]",
+ account: account,
+ accountRepo: repo,
+ groupID: 42,
+ sessionHash: "gemini:sticky",
+ handleError: func(context.Context, string, *Account, int, http.Header, []byte, string, int64, string, bool) *handleModelRateLimitResult {
+ t.Fatal("model rate limit should be handled before custom error fallback")
+ return nil
+ },
+ }
+
+ handled, outStatus, retErr := svc.applyErrorPolicy(p, http.StatusTooManyRequests, http.Header{}, body)
+
+ require.True(t, handled)
+ require.Equal(t, http.StatusTooManyRequests, outStatus)
+ require.NoError(t, retErr)
+ require.Len(t, repo.modelRateLimitCalls, 2)
+ require.Equal(t, "gemini-3-flash", repo.modelRateLimitCalls[0].modelKey)
+ require.Equal(t, antigravityGeminiModelRateLimitKey, repo.modelRateLimitCalls[1].modelKey)
+ require.Len(t, cache.deleteCalls, 1)
+ require.Equal(t, int64(42), cache.deleteCalls[0].groupID)
+ require.Equal(t, "gemini:sticky", cache.deleteCalls[0].sessionHash)
+}
+
// ---------------------------------------------------------------------------
// errorPolicyRepoStub — minimal AccountRepository stub for error policy tests
// ---------------------------------------------------------------------------
diff --git a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go
index 5cb03f30..e2da89b5 100644
--- a/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go
+++ b/backend/internal/service/gateway_anthropic_apikey_passthrough_test.go
@@ -112,7 +112,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardStreamPreservesBodyAnd
body := []byte(`{"model":"claude-3-7-sonnet-20250219","stream":true,"system":[{"type":"text","text":"x-anthropic-billing-header keep"}],"messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: "claude-3-7-sonnet-20250219",
Stream: true,
}
@@ -202,7 +202,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardCountTokensPreservesBo
body := []byte(`{"model":"claude-3-5-sonnet-latest","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}],"thinking":{"type":"enabled"}}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: "claude-3-5-sonnet-latest",
}
@@ -344,7 +344,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ModelMappingEdgeCases(t *test
body := []byte(`{"model":"` + tt.model + `","messages":[{"role":"user","content":[{"type":"text","text":"hello"}]}]}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: tt.model,
}
@@ -429,7 +429,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ModelMappingPreservesOtherFie
// 包含复杂字段的请求体:system、thinking、messages
body := []byte(`{"model":"claude-sonnet-4-20250514","system":[{"type":"text","text":"You are a helpful assistant."}],"messages":[{"role":"user","content":[{"type":"text","text":"hello world"}]}],"thinking":{"type":"enabled","budget_tokens":5000},"max_tokens":1024}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: "claude-sonnet-4-20250514",
}
@@ -476,6 +476,66 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ModelMappingPreservesOtherFie
require.Equal(t, int64(1024), gjson.GetBytes(sentBody, "max_tokens").Int(), "max_tokens 不应被修改")
}
+func TestGatewayService_AnthropicAPIKeyPassthrough_CountTokensFiltersGenerationFields(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil)
+
+ body := []byte(`{"model":"claude-sonnet-4-20250514","system":[{"type":"text","text":"sys"}],"messages":[{"role":"user","content":"hello"}],"tools":[{"name":"tool","input_schema":{"type":"object"}}],"temperature":0.7,"top_p":0.9,"top_k":40,"stream":true,"stop_sequences":["END"],"max_tokens":1024,"thinking":{"type":"enabled","budget_tokens":5000}}`)
+ parsed := &ParsedRequest{
+ Body: NewRequestBodyRef(body),
+ Model: "claude-sonnet-4-20250514",
+ }
+
+ upstreamRespBody := `{"input_tokens":42}`
+ upstream := &anthropicHTTPUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(upstreamRespBody)),
+ },
+ }
+
+ svc := &GatewayService{
+ cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}},
+ httpUpstream: upstream,
+ rateLimitService: &RateLimitService{},
+ }
+
+ account := &Account{
+ ID: 302,
+ Name: "count-token-filter-test",
+ Platform: PlatformAnthropic,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "upstream-key",
+ "base_url": "https://api.anthropic.com",
+ },
+ Extra: map[string]any{"anthropic_passthrough": true},
+ Status: StatusActive,
+ Schedulable: true,
+ }
+
+ err := svc.ForwardCountTokens(context.Background(), c, account, parsed)
+ require.NoError(t, err)
+
+ sentBody := upstream.lastBody
+ require.False(t, gjson.GetBytes(sentBody, "temperature").Exists())
+ require.False(t, gjson.GetBytes(sentBody, "top_p").Exists())
+ require.False(t, gjson.GetBytes(sentBody, "top_k").Exists())
+ require.False(t, gjson.GetBytes(sentBody, "stream").Exists())
+ require.False(t, gjson.GetBytes(sentBody, "stop_sequences").Exists())
+ require.Equal(t, "claude-sonnet-4-20250514", gjson.GetBytes(sentBody, "model").String())
+ require.Equal(t, "sys", gjson.GetBytes(sentBody, "system.0.text").String())
+ require.Equal(t, "hello", gjson.GetBytes(sentBody, "messages.0.content").String())
+ require.Equal(t, "tool", gjson.GetBytes(sentBody, "tools.0.name").String())
+ require.Equal(t, int64(1024), gjson.GetBytes(sentBody, "max_tokens").Int())
+ require.Equal(t, "enabled", gjson.GetBytes(sentBody, "thinking.type").String())
+}
+
// TestGatewayService_AnthropicAPIKeyPassthrough_EmptyModelSkipsMapping
// 确保空模型名不会触发映射逻辑
func TestGatewayService_AnthropicAPIKeyPassthrough_EmptyModelSkipsMapping(t *testing.T) {
@@ -487,7 +547,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_EmptyModelSkipsMapping(t *tes
body := []byte(`{"messages":[{"role":"user","content":"hello"}]}`)
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
Model: "", // 空模型
}
@@ -576,7 +636,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_CountTokens404PassthroughNotE
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil)
body := []byte(`{"model":"claude-sonnet-4-5-20250929","messages":[{"role":"user","content":"hi"}]}`)
- parsed := &ParsedRequest{Body: body, Model: "claude-sonnet-4-5-20250929"}
+ parsed := &ParsedRequest{Body: NewRequestBodyRef(body), Model: "claude-sonnet-4-5-20250929"}
upstream := &anthropicHTTPUpstreamRecorder{
resp: &http.Response{
@@ -653,7 +713,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_BuildRequestRejectsInvalidBas
},
}
- _, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(context.Background(), c, account, []byte(`{}`), "k")
+ _, _, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(context.Background(), c, account, []byte(`{}`), "k")
require.Error(t, err)
}
@@ -678,7 +738,7 @@ func TestGatewayService_AnthropicOAuth_NotAffectedByAPIKeyPassthroughToggle(t *t
require.False(t, account.IsAnthropicAPIKeyPassthroughEnabled())
- req, err := svc.buildUpstreamRequest(context.Background(), c, account, []byte(`{"model":"claude-3-7-sonnet-20250219"}`), "oauth-token", "oauth", "claude-3-7-sonnet-20250219", true, false)
+ req, _, err := svc.buildUpstreamRequest(context.Background(), c, account, []byte(`{"model":"claude-3-7-sonnet-20250219"}`), "oauth-token", "oauth", "claude-3-7-sonnet-20250219", true, false)
require.NoError(t, err)
require.Equal(t, "Bearer oauth-token", getHeaderRaw(req.Header, "authorization"))
require.Contains(t, getHeaderRaw(req.Header, "anthropic-beta"), claude.BetaOAuth, "OAuth 链路仍应按原逻辑补齐 oauth beta")
@@ -707,7 +767,7 @@ func TestGatewayService_AnthropicOAuth_ForwardPreservesBillingHeaderSystemBlock(
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
- parsed, err := ParseGatewayRequest([]byte(tt.body), PlatformAnthropic)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), PlatformAnthropic)
require.NoError(t, err)
upstream := &anthropicHTTPUpstreamRecorder{
diff --git a/backend/internal/service/gateway_anthropic_vertex_service_account_test.go b/backend/internal/service/gateway_anthropic_vertex_service_account_test.go
index aa779805..be8c5867 100644
--- a/backend/internal/service/gateway_anthropic_vertex_service_account_test.go
+++ b/backend/internal/service/gateway_anthropic_vertex_service_account_test.go
@@ -35,7 +35,7 @@ func TestGatewayService_BuildAnthropicVertexServiceAccountRequest(t *testing.T)
body := []byte(`{"model":"claude-sonnet-4-5","stream":false,"max_tokens":32,"messages":[{"role":"user","content":"hello"}]}`)
svc := &GatewayService{}
- req, err := svc.buildUpstreamRequest(
+ req, _, err := svc.buildUpstreamRequest(
context.Background(),
c,
account,
@@ -66,3 +66,67 @@ func readRequestBodyForTest(t *testing.T, req *http.Request) []byte {
require.NoError(t, err)
return body
}
+
+// Vertex 路径回归保护:同样需要
+// body↔beta header 能力维度对称。客户端 header 不带 context-management beta
+// 但 body 带 context_management 字段 → Vertex builder 必须 strip 字段,与 Anthropic
+// 直连 / Bedrock 路径保持一致。
+func TestGatewayService_BuildAnthropicVertexServiceAccount_StripsContextManagementWhenBetaMissing(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
+ // 客户端 header 只带 interleaved-thinking,不带 context-management-2025-06-27
+ c.Request.Header.Set("Anthropic-Beta", "interleaved-thinking-2025-05-14")
+
+ account := &Account{
+ ID: 302, Platform: PlatformAnthropic, Type: AccountTypeServiceAccount,
+ Credentials: map[string]any{"project_id": "vertex-proj", "location": "us-east5"},
+ }
+ // body 带了 context_management 字段(客户端透传 / normalize 补齐 / mimicry 注入等场景都可能导致)
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015","keep":"all"}]},"messages":[{"role":"user","content":"hi"}]}`)
+
+ svc := &GatewayService{}
+ req, _, err := svc.buildUpstreamRequest(
+ context.Background(), c, account, body,
+ "vertex-token", "service_account", "claude-haiku-4-5@20251001", false, false,
+ )
+ require.NoError(t, err)
+
+ got := readRequestBodyForTest(t, req)
+ require.False(t, gjson.GetBytes(got, "context_management").Exists(),
+ "Vertex 路径下客户端 header 缺 context-management beta 时,必须 strip body 同名字段")
+ // header 对称断言:覆盖未来某人在 Vertex builder 里加入与 sanitize 不一致的 header 处理。
+ outBeta := getHeaderRaw(req.Header, "anthropic-beta")
+ require.False(t, anthropicBetaTokensContains(outBeta, "context-management-2025-06-27"),
+ "与 body 对称:outgoing anthropic-beta header 也不含 context-management beta")
+}
+
+// Vertex 路径反面:客户端 header 含 context-management beta 时保留字段。
+func TestGatewayService_BuildAnthropicVertexServiceAccount_PreservesContextManagementWhenBetaPresent(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
+ c.Request.Header.Set("Anthropic-Beta", "interleaved-thinking-2025-05-14,context-management-2025-06-27")
+
+ account := &Account{
+ ID: 303, Platform: PlatformAnthropic, Type: AccountTypeServiceAccount,
+ Credentials: map[string]any{"project_id": "vertex-proj", "location": "us-east5"},
+ }
+ body := []byte(`{"model":"claude-sonnet-4-6","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
+
+ svc := &GatewayService{}
+ req, _, err := svc.buildUpstreamRequest(
+ context.Background(), c, account, body,
+ "vertex-token", "service_account", "claude-sonnet-4-6@20260218", false, false,
+ )
+ require.NoError(t, err)
+
+ got := readRequestBodyForTest(t, req)
+ require.True(t, gjson.GetBytes(got, "context_management").Exists(),
+ "Vertex + 客户端 header 包含 context-management beta 时字段必须保留")
+ outBeta := getHeaderRaw(req.Header, "anthropic-beta")
+ require.True(t, anthropicBetaTokensContains(outBeta, "context-management-2025-06-27"),
+ "与 body 对称:outgoing anthropic-beta header 同步含 context-management beta")
+}
diff --git a/backend/internal/service/gateway_context_management_test.go b/backend/internal/service/gateway_context_management_test.go
new file mode 100644
index 00000000..51b12809
--- /dev/null
+++ b/backend/internal/service/gateway_context_management_test.go
@@ -0,0 +1,667 @@
+//go:build unit
+
+package service
+
+import (
+ "context"
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "regexp"
+ "strings"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/claude"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
+)
+
+// ============================================================================
+// 背景
+// ============================================================================
+//
+// Anthropic 上游对 body.context_management 字段实施 Pydantic schema 校验:
+// 当且仅当 anthropic-beta header 含 context-management-2025-06-27 时接受。
+// 否则报:
+// "context_management: Extra inputs are not permitted"
+//
+// 本仓采用能力维度对称约束(与 Bedrock 路径的 sanitizeBedrockFieldsForBetaTokens
+// 对称):在所有 Anthropic 直连出口,按最终 anthropic-beta header 是否含上述 token
+// 决定 body 是否保留同名字段。
+//
+// 本文件覆盖:
+// 1) sanitizeAnthropicBodyForBetaTokens 纯函数
+// 2) anthropicBetaTokensContains 解析辅助函数
+// 3) computeFinalAnthropicBeta / computeFinalCountTokensAnthropicBeta 各路径
+// 4) normalizeClaudeOAuthRequestBody 的 context_management 补齐行为(不再按 model 短路)
+
+// ============================================================================
+// anthropicBetaTokensContains
+// ============================================================================
+
+func TestAnthropicBetaTokensContains_EmptyInputs(t *testing.T) {
+ require.False(t, anthropicBetaTokensContains("", "context-management-2025-06-27"))
+ require.False(t, anthropicBetaTokensContains("oauth-2025-04-20", ""))
+}
+
+func TestAnthropicBetaTokensContains_SingleToken(t *testing.T) {
+ require.True(t, anthropicBetaTokensContains("context-management-2025-06-27", "context-management-2025-06-27"))
+}
+
+func TestAnthropicBetaTokensContains_MultiTokenComma(t *testing.T) {
+ header := "oauth-2025-04-20,context-management-2025-06-27,interleaved-thinking-2025-05-14"
+ require.True(t, anthropicBetaTokensContains(header, "context-management-2025-06-27"))
+ require.True(t, anthropicBetaTokensContains(header, "oauth-2025-04-20"))
+ require.False(t, anthropicBetaTokensContains(header, "fast-mode-2026-02-01"))
+}
+
+func TestAnthropicBetaTokensContains_ToleratesWhitespace(t *testing.T) {
+ header := "oauth-2025-04-20 , context-management-2025-06-27 , interleaved-thinking-2025-05-14"
+ require.True(t, anthropicBetaTokensContains(header, "context-management-2025-06-27"))
+}
+
+func TestAnthropicBetaTokensContains_SubstringNotMatched(t *testing.T) {
+ // 严格 token 比较,不应被子串误匹配
+ require.False(t, anthropicBetaTokensContains("context-management-2025-06-27-rev2", "context-management-2025-06-27"),
+ "必须按 token 边界匹配,不允许 prefix 子串误命中")
+}
+
+// ============================================================================
+// sanitizeAnthropicBodyForBetaTokens
+// ============================================================================
+
+func TestSanitizeAnthropicBodyForBetaTokens_NoFieldNoChange(t *testing.T) {
+ body := []byte(`{"model":"claude-haiku-4-5","messages":[]}`)
+ out, changed := sanitizeAnthropicBodyForBetaTokens(body, "oauth-2025-04-20")
+ require.False(t, changed)
+ require.Equal(t, string(body), string(out))
+}
+
+func TestSanitizeAnthropicBodyForBetaTokens_FieldKeptWhenBetaPresent(t *testing.T) {
+ body := []byte(`{"model":"claude-opus-4-7","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
+ out, changed := sanitizeAnthropicBodyForBetaTokens(body,
+ "oauth-2025-04-20,context-management-2025-06-27,interleaved-thinking-2025-05-14")
+ require.False(t, changed)
+ require.True(t, gjson.GetBytes(out, "context_management").Exists())
+ require.Equal(t, "clear_thinking_20251015",
+ gjson.GetBytes(out, "context_management.edits.0.type").String())
+}
+
+func TestSanitizeAnthropicBodyForBetaTokens_FieldStrippedWhenBetaMissing(t *testing.T) {
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
+ out, changed := sanitizeAnthropicBodyForBetaTokens(body, "oauth-2025-04-20,interleaved-thinking-2025-05-14")
+ require.True(t, changed)
+ require.False(t, gjson.GetBytes(out, "context_management").Exists(),
+ "header 不含 context-management beta 时必须 strip 同名字段")
+}
+
+func TestSanitizeAnthropicBodyForBetaTokens_FieldStrippedWhenBetaEmpty(t *testing.T) {
+ body := []byte(`{"context_management":{"edits":[]},"messages":[]}`)
+ out, changed := sanitizeAnthropicBodyForBetaTokens(body, "")
+ require.True(t, changed)
+ require.False(t, gjson.GetBytes(out, "context_management").Exists())
+}
+
+func TestSanitizeAnthropicBodyForBetaTokens_EmptyBody(t *testing.T) {
+ out, changed := sanitizeAnthropicBodyForBetaTokens([]byte{}, "")
+ require.False(t, changed)
+ require.Empty(t, out)
+
+ out, changed = sanitizeAnthropicBodyForBetaTokens(nil, "")
+ require.False(t, changed)
+ require.Empty(t, out)
+}
+
+// ★ 关键回归断言:能力维度 sanitize 解决了 "真 CC + haiku" 路径的过度删除问题。
+// 真实 Claude Code CLI 2.1.87+ 客户端 header 含 context-management beta;
+// 即使 model 是 haiku,sanitize 也不应剥离功能字段。
+func TestSanitizeAnthropicBodyForBetaTokens_HaikuRealCCClientPreservesField(t *testing.T) {
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015","keep":"all"}]},"messages":[]}`)
+ // 真 Claude Code CLI 2.1.87+ 客户端 header 含 context-management beta
+ clientBeta := "claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14,context-management-2025-06-27"
+ out, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta)
+ require.False(t, changed,
+ "真 CC 客户端 header 含 context-management beta 时,haiku body 字段必须保留(功能不丢)")
+ require.True(t, gjson.GetBytes(out, "context_management").Exists())
+}
+
+// ============================================================================
+// computeFinalAnthropicBeta — 关键路径
+// ============================================================================
+
+func newTestGatewayServiceForBeta(injectBetaForAPIKey bool) *GatewayService {
+ cfg := &config.Config{}
+ cfg.Gateway.InjectBetaForAPIKey = injectBetaForAPIKey
+ return &GatewayService{cfg: cfg}
+}
+
+func TestComputeFinalAnthropicBeta_OAuthMimic_NonHaiku_IncludesContextManagement(t *testing.T) {
+ s := newTestGatewayServiceForBeta(false)
+ final, ok := s.computeFinalAnthropicBeta("oauth", true, "claude-sonnet-4-6", http.Header{}, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement),
+ "OAuth mimic non-haiku 必须注入完整 CC mimicry beta,含 context-management-2025-06-27")
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaOAuth))
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaClaudeCode))
+}
+
+func TestComputeFinalAnthropicBeta_OAuthMimic_Haiku_ExcludesContextManagement(t *testing.T) {
+ s := newTestGatewayServiceForBeta(false)
+ final, ok := s.computeFinalAnthropicBeta("oauth", true, "claude-haiku-4-5", http.Header{}, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.False(t, anthropicBetaTokensContains(final, claude.BetaContextManagement),
+ "OAuth mimic haiku 仅注入 oauth + interleaved-thinking,不含 context-management")
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaOAuth))
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaInterleavedThinking))
+}
+
+func TestComputeFinalAnthropicBeta_OAuthMimic_IgnoresClientBeta(t *testing.T) {
+ // mimic 路径下原代码白名单透传被跳过,client beta 应被忽略
+ s := newTestGatewayServiceForBeta(false)
+ hdr := http.Header{}
+ hdr.Set("anthropic-beta", "custom-experimental-beta")
+ final, ok := s.computeFinalAnthropicBeta("oauth", true, "claude-sonnet-4-6", hdr, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.False(t, strings.Contains(final, "custom-experimental-beta"),
+ "mimic 路径必须忽略客户端 anthropic-beta header")
+}
+
+func TestComputeFinalAnthropicBeta_OAuthTransparent_NonHaiku_PreservesClientContextManagement(t *testing.T) {
+ // 真 CC 客户端透传:客户端 header 中的 context-management beta 必须保留
+ s := newTestGatewayServiceForBeta(false)
+ hdr := http.Header{}
+ hdr.Set("anthropic-beta", "claude-code-20250219,oauth-2025-04-20,context-management-2025-06-27")
+ final, ok := s.computeFinalAnthropicBeta("oauth", false, "claude-sonnet-4-6", hdr, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement))
+}
+
+func TestComputeFinalAnthropicBeta_OAuthTransparent_Haiku_RealCCPreservesContextManagement(t *testing.T) {
+ // haiku 透传 + 客户端带 context-management beta → 必须保留
+ // (能力维度核心场景:避免 model-name 误删客户端透传的功能 beta)
+ s := newTestGatewayServiceForBeta(false)
+ hdr := http.Header{}
+ hdr.Set("anthropic-beta", "claude-code-20250219,oauth-2025-04-20,context-management-2025-06-27,interleaved-thinking-2025-05-14")
+ final, ok := s.computeFinalAnthropicBeta("oauth", false, "claude-haiku-4-5", hdr, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement),
+ "真 CC + haiku + 客户端带 context-management beta → 透传必须保留")
+}
+
+func TestComputeFinalAnthropicBeta_APIKey_PassesClientBetaThroughDropSet(t *testing.T) {
+ s := newTestGatewayServiceForBeta(false)
+ hdr := http.Header{}
+ hdr.Set("anthropic-beta", "oauth-2025-04-20,custom-beta")
+ final, ok := s.computeFinalAnthropicBeta("apikey", false, "claude-sonnet-4-6", hdr, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.True(t, anthropicBetaTokensContains(final, "oauth-2025-04-20"))
+ require.True(t, anthropicBetaTokensContains(final, "custom-beta"))
+}
+
+func TestComputeFinalAnthropicBeta_APIKey_NoClientBetaInjectOff_ShouldNotSet(t *testing.T) {
+ s := newTestGatewayServiceForBeta(false)
+ final, ok := s.computeFinalAnthropicBeta("apikey", false, "claude-sonnet-4-6", http.Header{}, []byte(`{}`), nil)
+ require.False(t, ok, "API-key + 客户端未传 + InjectBetaForAPIKey 关 → 不应主动设置 anthropic-beta")
+ require.Equal(t, "", final)
+}
+
+// ============================================================================
+// computeFinalCountTokensAnthropicBeta
+// ============================================================================
+
+func TestComputeFinalCountTokensAnthropicBeta_OAuthMimic_AlwaysIncludesContextManagement(t *testing.T) {
+ // count_tokens 路径下 mimic 不按 haiku 排除:始终注入完整 mimicry beta
+ s := newTestGatewayServiceForBeta(false)
+ final, ok := s.computeFinalCountTokensAnthropicBeta("oauth", true, "claude-haiku-4-5", http.Header{}, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement),
+ "count_tokens + mimic 即使 haiku 也注入 context-management beta(与 messages 不同)")
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaTokenCounting),
+ "count_tokens 路径必须含 token-counting beta")
+}
+
+// 重构等价性回归:
+// 原 main buildCountTokensRequest 在 count_tokens mimic 分支上不跳过白名单透传
+// (与 messages mimic 不同),incomingBeta 取自客户端透传。重构后必须从 clientHeaders
+// 拿同一个值并 merge,否则会丢失客户端 beta。
+func TestComputeFinalCountTokensAnthropicBeta_OAuthMimic_PreservesClientBeta(t *testing.T) {
+ s := newTestGatewayServiceForBeta(false)
+ hdr := http.Header{}
+ hdr.Set("anthropic-beta", "custom-experimental-beta,context-1m-2025-08-07")
+ final, ok := s.computeFinalCountTokensAnthropicBeta("oauth", true, "claude-haiku-4-5", hdr, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.True(t, anthropicBetaTokensContains(final, "custom-experimental-beta"),
+ "count_tokens mimic 不同于 messages mimic:原代码会保留客户端透传的 beta")
+ require.True(t, anthropicBetaTokensContains(final, "context-1m-2025-08-07"),
+ "客户端透传的其他 beta token 同样需要保留")
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement),
+ "同时 FullClaudeCodeMimicryBetas 不打折扣")
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaTokenCounting),
+ "同时补齐 token-counting beta")
+}
+
+// messages mimic 路径反向验证:原代码会跳过白名单透传,
+// 客户端 beta 不会进入 mimic 计算。重构后 messages computeFinalAnthropicBeta
+// mimic 分支依然不该使用 clientBeta。
+func TestComputeFinalAnthropicBeta_OAuthMimic_IgnoresClientBetaExplicit(t *testing.T) {
+ s := newTestGatewayServiceForBeta(false)
+ hdr := http.Header{}
+ hdr.Set("anthropic-beta", "custom-experimental-beta")
+ final, ok := s.computeFinalAnthropicBeta("oauth", true, "claude-sonnet-4-6", hdr, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.False(t, anthropicBetaTokensContains(final, "custom-experimental-beta"),
+ "messages mimic 原代码跳过白名单透传 → 客户端 beta 不进入计算。"+
+ "与 count_tokens mimic 是不同的设计,不能合并为同一函数。")
+}
+
+func TestComputeFinalCountTokensAnthropicBeta_OAuthTransparent_NoClientBetaInjectsDefault(t *testing.T) {
+ // 真 CC 客户端透传 + 客户端未传 anthropic-beta → 用 CountTokensBetaHeader 兜底
+ s := newTestGatewayServiceForBeta(false)
+ final, ok := s.computeFinalCountTokensAnthropicBeta("oauth", false, "claude-haiku-4-5", http.Header{}, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.Equal(t, claude.CountTokensBetaHeader, final)
+ // CountTokensBetaHeader 不含 context-management beta
+ require.False(t, anthropicBetaTokensContains(final, claude.BetaContextManagement))
+}
+
+func TestComputeFinalCountTokensAnthropicBeta_OAuthTransparent_AppendsBetaTokenCounting(t *testing.T) {
+ s := newTestGatewayServiceForBeta(false)
+ hdr := http.Header{}
+ hdr.Set("anthropic-beta", "oauth-2025-04-20,context-management-2025-06-27")
+ final, ok := s.computeFinalCountTokensAnthropicBeta("oauth", false, "claude-sonnet-4-6", hdr, []byte(`{}`), nil)
+ require.True(t, ok)
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaTokenCounting),
+ "客户端未带 token-counting beta 时必须补齐")
+ require.True(t, anthropicBetaTokensContains(final, claude.BetaContextManagement),
+ "客户端带的 context-management beta 必须保留")
+}
+
+// ============================================================================
+// normalizeClaudeOAuthRequestBody — 回归:context_management 补齐恢复原行为
+// ============================================================================
+//
+// 重构后该函数不再按 model 名短路:thinking=enabled/adaptive 时补齐 context_management,
+// 与 model 无关。strip 责任移交 sanitizeAnthropicBodyForBetaTokens(在
+// buildUpstreamRequest 层按最终 beta header 执行)。
+
+func TestNormalizeClaudeOAuthRequestBody_InjectsContextManagement_ThinkingEnabled(t *testing.T) {
+ body := []byte(`{"model":"claude-sonnet-4-6","thinking":{"type":"enabled","budget_tokens":1000},"messages":[]}`)
+ out, _ := normalizeClaudeOAuthRequestBody(body, "claude-sonnet-4-6", claudeOAuthNormalizeOptions{})
+ require.True(t, gjson.GetBytes(out, "context_management").Exists())
+ require.Equal(t, "clear_thinking_20251015",
+ gjson.GetBytes(out, "context_management.edits.0.type").String())
+}
+
+func TestNormalizeClaudeOAuthRequestBody_InjectsContextManagement_ThinkingAdaptive(t *testing.T) {
+ body := []byte(`{"model":"claude-opus-4-7","thinking":{"type":"adaptive"},"messages":[]}`)
+ out, _ := normalizeClaudeOAuthRequestBody(body, "claude-opus-4-7", claudeOAuthNormalizeOptions{})
+ require.True(t, gjson.GetBytes(out, "context_management").Exists())
+}
+
+func TestNormalizeClaudeOAuthRequestBody_HaikuStillInjects_StripDeferredToSanitize(t *testing.T) {
+ // haiku + thinking=enabled:normalize 阶段仍按 CLI mimicry 行为补齐字段;
+ // strip 由 buildUpstreamRequest 层的 sanitize 兜底(如果 final beta 不含 token)。
+ body := []byte(`{"model":"claude-haiku-4-5","thinking":{"type":"enabled","budget_tokens":1000},"messages":[]}`)
+ out, _ := normalizeClaudeOAuthRequestBody(body, "claude-haiku-4-5", claudeOAuthNormalizeOptions{})
+ require.True(t, gjson.GetBytes(out, "context_management").Exists(),
+ "normalize 不再按 model 名短路;strip 责任移交 sanitize 层")
+}
+
+func TestNormalizeClaudeOAuthRequestBody_PreservesClientContextManagement(t *testing.T) {
+ body := []byte(`{"model":"claude-opus-4-7","context_management":{"edits":[{"type":"custom_strategy"}]},"thinking":{"type":"enabled","budget_tokens":1000},"messages":[]}`)
+ out, _ := normalizeClaudeOAuthRequestBody(body, "claude-opus-4-7", claudeOAuthNormalizeOptions{})
+ require.Equal(t, "custom_strategy",
+ gjson.GetBytes(out, "context_management.edits.0.type").String(),
+ "客户端透传的 context_management 内容必须原样保留")
+}
+
+func TestNormalizeClaudeOAuthRequestBody_NoThinking_NoInject(t *testing.T) {
+ body := []byte(`{"model":"claude-sonnet-4-6","messages":[]}`)
+ out, _ := normalizeClaudeOAuthRequestBody(body, "claude-sonnet-4-6", claudeOAuthNormalizeOptions{})
+ require.False(t, gjson.GetBytes(out, "context_management").Exists())
+}
+
+// ============================================================================
+// passthrough 集成测试:buildUpstreamRequest-
+// AnthropicAPIKeyPassthrough 与 buildCountTokensRequestAnthropicAPIKeyPassthrough
+// 路径上 sanitize 是否生效。
+// ============================================================================
+
+// passthrough 集成测试不设 base_url,避开 validateUpstreamBaseURL 对 cfg.Security 的依赖。
+// targetURL 会走默认 claudeAPIURL,sanitize 逻辑与 baseURL 是否存在无关。
+func newAnthropicAPIKeyPassthroughAccountForBetaTest() *Account {
+ return &Account{
+ ID: 501,
+ Name: "anthropic-apikey-passthrough-ctxmgmt-test",
+ Platform: PlatformAnthropic,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{
+ "api_key": "upstream-key",
+ },
+ Extra: map[string]any{"anthropic_passthrough": true},
+ Status: StatusActive,
+ Schedulable: true,
+ }
+}
+
+func readUpstreamBodyForTest(t *testing.T, req *http.Request) []byte {
+ t.Helper()
+ require.NotNil(t, req.Body)
+ b, err := io.ReadAll(req.Body)
+ require.NoError(t, err)
+ return b
+}
+
+func TestBuildUpstreamRequestAnthropicAPIKeyPassthrough_StripsContextManagementWhenClientHeaderMissingBeta(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
+ // 客户端仅带 oauth beta,不带 context-management-2025-06-27
+ c.Request.Header.Set("Anthropic-Beta", "oauth-2025-04-20")
+
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, _, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(
+ context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token",
+ )
+ require.NoError(t, err)
+ require.False(t, gjson.GetBytes(readUpstreamBodyForTest(t, req), "context_management").Exists(),
+ "API-key passthrough + 客户端未带 context-management beta → strip body 字段")
+}
+
+func TestBuildUpstreamRequestAnthropicAPIKeyPassthrough_PreservesContextManagementWhenClientHeaderHasBeta(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
+ c.Request.Header.Set("Anthropic-Beta", "oauth-2025-04-20,context-management-2025-06-27")
+
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, _, err := svc.buildUpstreamRequestAnthropicAPIKeyPassthrough(
+ context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token",
+ )
+ require.NoError(t, err)
+ require.True(t, gjson.GetBytes(readUpstreamBodyForTest(t, req), "context_management").Exists(),
+ "API-key passthrough + 客户端带 context-management beta → 字段保留(不过度删除)")
+}
+
+func TestBuildCountTokensRequestAnthropicAPIKeyPassthrough_StripsContextManagementWhenClientHeaderMissingBeta(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil)
+ c.Request.Header.Set("Anthropic-Beta", "oauth-2025-04-20,token-counting-2024-11-01")
+
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, err := svc.buildCountTokensRequestAnthropicAPIKeyPassthrough(
+ context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token",
+ )
+ require.NoError(t, err)
+ require.False(t, gjson.GetBytes(readUpstreamBodyForTest(t, req), "context_management").Exists(),
+ "count_tokens passthrough + 客户端未带 context-management beta → strip")
+}
+
+// ============================================================================
+// 集成测试:buildUpstreamRequest
+// 全路径验证上游 outgoing body 与 anthropic-beta header 严格对称。
+// 这个测试能挡住未来某人忘调 sanitize / 将 sanitize 挪到 CCH 之后 等 regression。
+// ============================================================================
+
+func TestBuildUpstreamRequest_OAuthMimicHaiku_StripsContextManagementEndToEnd(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
+
+ account := &Account{ID: 401, Platform: PlatformAnthropic, Type: AccountTypeOAuth,
+ Credentials: map[string]any{"access_token": "oauth-tok"},
+ Status: StatusActive,
+ Schedulable: true,
+ }
+ // haiku + mimic CC → final beta = HaikuBetaHeader(不含 context-management)→
+ // body 必须 strip。
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, _, err := svc.buildUpstreamRequest(
+ context.Background(), c, account, body,
+ "oauth-tok", "oauth", "claude-haiku-4-5", false, true, // mimicClaudeCode=true
+ )
+ require.NoError(t, err)
+
+ outBody := readUpstreamBodyForTest(t, req)
+ outBeta := getHeaderRaw(req.Header, "anthropic-beta")
+
+ require.False(t, gjson.GetBytes(outBody, "context_management").Exists(),
+ "OAuth mimic + haiku 端到端:outgoing body 不应含 context_management")
+ require.False(t, anthropicBetaTokensContains(outBeta, claude.BetaContextManagement),
+ "对称约束:outgoing anthropic-beta header 也不带 context-management beta")
+}
+
+func TestBuildUpstreamRequest_OAuthMimicNonHaiku_PreservesContextManagementEndToEnd(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
+
+ account := &Account{ID: 402, Platform: PlatformAnthropic, Type: AccountTypeOAuth,
+ Credentials: map[string]any{"access_token": "oauth-tok"},
+ Status: StatusActive,
+ Schedulable: true,
+ }
+ // sonnet + mimic CC → final beta = FullClaudeCodeMimicryBetas(含 context-management)→
+ // body 保留。
+ body := []byte(`{"model":"claude-sonnet-4-6","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, _, err := svc.buildUpstreamRequest(
+ context.Background(), c, account, body,
+ "oauth-tok", "oauth", "claude-sonnet-4-6", false, true,
+ )
+ require.NoError(t, err)
+
+ outBody := readUpstreamBodyForTest(t, req)
+ outBeta := getHeaderRaw(req.Header, "anthropic-beta")
+
+ require.True(t, gjson.GetBytes(outBody, "context_management").Exists(),
+ "OAuth mimic + non-haiku:outgoing body 必须保留 context_management。")
+ require.True(t, anthropicBetaTokensContains(outBeta, claude.BetaContextManagement),
+ "对称约束:outgoing anthropic-beta header 同时含 context-management beta")
+}
+
+func TestBuildUpstreamRequest_OAuthTransparentHaikuWithRealCCBeta_PreservesField(t *testing.T) {
+ // 端到端验证:真 CC 客户端 + haiku + 客户端 header 带 context-management beta
+ // → final beta 透传 → 不应该过度删除 body 字段
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
+ c.Request.Header.Set("Anthropic-Beta",
+ "claude-code-20250219,oauth-2025-04-20,interleaved-thinking-2025-05-14,context-management-2025-06-27")
+
+ account := &Account{ID: 403, Platform: PlatformAnthropic, Type: AccountTypeOAuth,
+ Credentials: map[string]any{"access_token": "oauth-tok"},
+ Status: StatusActive, Schedulable: true,
+ }
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015","keep":"all"}]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, _, err := svc.buildUpstreamRequest(
+ context.Background(), c, account, body,
+ "oauth-tok", "oauth", "claude-haiku-4-5", false, false, // mimicClaudeCode=false(真 CC)
+ )
+ require.NoError(t, err)
+
+ outBody := readUpstreamBodyForTest(t, req)
+ outBeta := getHeaderRaw(req.Header, "anthropic-beta")
+
+ require.True(t, anthropicBetaTokensContains(outBeta, claude.BetaContextManagement),
+ "真 CC 透传路径:客户端 header 中的 context-management beta 必须保留")
+ require.True(t, gjson.GetBytes(outBody, "context_management").Exists(),
+ "回归保护:真 CC + haiku + 客户端带 beta token 时,clear_thinking_20251015 功能不能静默失效")
+}
+
+// CCH 顺序语义测试:sanitize 必须在 signBillingHeaderCCH 之前,
+// 否则签名的 hash 与最终发送的 body 不一致,被 Anthropic 判 third-party。
+//
+// 该测试不走 buildUpstreamRequest 完整路径(需要 mock SettingService 成本高),
+// 而是直接验证两个顺序产生的 cch 不同,证明二者不可交换。
+// 测试名本身是语义约束的文档化 marker。
+func TestSanitizeMustBeBeforeCCHSigning_HashConsistency(t *testing.T) {
+ // 构造 body:含 context_management + cch=00000 占位符
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"system":[{"type":"text","text":"x-anthropic-billing-header: cc_version=2.1.92; cch=00000;"}],"messages":[]}`)
+
+ // 最终发送场景:final beta 不含 context-management beta → sanitize 会 strip
+ finalBeta := "oauth-2025-04-20,interleaved-thinking-2025-05-14"
+
+ extractCCH := func(t *testing.T, b []byte) string {
+ t.Helper()
+ m := regexp.MustCompile(`\bcch=([0-9a-fA-F]{5})\b`).FindSubmatch(b)
+ require.NotNil(t, m, "body 里找不到 cch=<5hex> :%s", string(b))
+ return string(m[1])
+ }
+
+ // === 正确顺序:sanitize → signBillingHeaderCCH ===
+ // 1. strip context_management
+ sanitizedFirst, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBeta)
+ require.True(t, changed)
+ require.False(t, gjson.GetBytes(sanitizedFirst, "context_management").Exists())
+ // 2. 基于“strip 后的 body”算 hash
+ correctFinal := signBillingHeaderCCH(sanitizedFirst)
+ correctCCH := extractCCH(t, correctFinal)
+ require.NotEqual(t, "00000", correctCCH, "placeholder 应被替换")
+
+ // === 错误顺序:signBillingHeaderCCH → sanitize(未来 regression 场景)===
+ // 1. 先基于“含 context_management 的 body”算 hash → cch=H_with
+ signedFirst := signBillingHeaderCCH(body)
+ wrongCCH := extractCCH(t, signedFirst)
+ require.NotEqual(t, "00000", wrongCCH)
+ // 2. 后 strip context_management → body 变化但 cch 仍是 H_with
+ wrongFinal, _ := sanitizeAnthropicBodyForBetaTokens(signedFirst, finalBeta)
+ wrongFinalCCH := extractCCH(t, wrongFinal)
+
+ // === 关键断言 ===
+ // 上游验证逻辑:将 outgoing body 的 cch 还原为 00000、重算 hash、与 cch 字段比较。
+ // 模拟上游验证:用发送 body 算出“期望的 cch”,与发送 body 里的 cch 字段比。
+ recomputeExpected := func(b []byte, currentCCH string) string {
+ t.Helper()
+ // 把 cch= 还原为 cch=00000
+ re := regexp.MustCompile(`(\bcch=)` + currentCCH + `(\b)`)
+ restored := re.ReplaceAll(b, []byte("${1}00000${2}"))
+ return extractCCH(t, signBillingHeaderCCH(restored))
+ }
+
+ // 正确顺序:发送 body 的 cch == 重算 hash → 上游验证过
+ require.Equal(t, correctCCH, recomputeExpected(correctFinal, correctCCH),
+ "正确顺序:final body 里的 cch 与重算 hash 一致 → 上游验证通过")
+
+ // 错误顺序:发送 body 的 cch 是“含 ctx 算的”,但最终 body 不含 ctx → 重算 hash 不同
+ require.NotEqual(t, wrongFinalCCH, recomputeExpected(wrongFinal, wrongFinalCCH),
+ "错误顺序:final body 里的 cch 是基于含 ctx 的 body 算的,"+
+ "但发送 body 已 strip ctx → 上游重算 hash 与 cch 不一致 → 被判 third-party。"+
+ "这是 buildUpstreamRequest / buildCountTokensRequest 里 sanitize 必须在 "+
+ "signBillingHeaderCCH 之前的原因。")
+}
+
+// count_tokens 主路径 E2E 集成测试
+func TestBuildCountTokensRequest_OAuthMimicHaiku_PreservesContextManagementEndToEnd(t *testing.T) {
+ // count_tokens 路径下 mimic 不按 haiku 排除,始终注入 BetaContextManagement
+ // → sanitize 看到最终 beta header 含 context-management beta → 字段保留。
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil)
+
+ account := &Account{ID: 411, Platform: PlatformAnthropic, Type: AccountTypeOAuth,
+ Credentials: map[string]any{"access_token": "oauth-tok"},
+ Status: StatusActive, Schedulable: true,
+ }
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, _, err := svc.buildCountTokensRequest(
+ context.Background(), c, account, body,
+ "oauth-tok", "oauth", "claude-haiku-4-5", true, // mimicClaudeCode=true
+ )
+ require.NoError(t, err)
+
+ outBody := readUpstreamBodyForTest(t, req)
+ outBeta := getHeaderRaw(req.Header, "anthropic-beta")
+
+ require.True(t, anthropicBetaTokensContains(outBeta, claude.BetaContextManagement),
+ "count_tokens mimic 始终注入 context-management beta")
+ require.True(t, gjson.GetBytes(outBody, "context_management").Exists(),
+ "对称约束:final beta 含 token 时 body 字段保留")
+ require.True(t, anthropicBetaTokensContains(outBeta, claude.BetaTokenCounting),
+ "count_tokens 路径必须含 token-counting beta")
+}
+
+func TestBuildCountTokensRequest_APIKeyHaiku_StripsContextManagementEndToEnd(t *testing.T) {
+ // API-key + haiku + 客户端 header 不带 context-management beta → final beta 不含 → strip
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil)
+ c.Request.Header.Set("Anthropic-Beta", "interleaved-thinking-2025-05-14")
+
+ account := &Account{ID: 412, Platform: PlatformAnthropic, Type: AccountTypeAPIKey,
+ Credentials: map[string]any{"api_key": "sk-ant-xxx"},
+ Status: StatusActive, Schedulable: true,
+ }
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, _, err := svc.buildCountTokensRequest(
+ context.Background(), c, account, body,
+ "sk-ant-xxx", "apikey", "claude-haiku-4-5", false,
+ )
+ require.NoError(t, err)
+
+ outBody := readUpstreamBodyForTest(t, req)
+ require.False(t, gjson.GetBytes(outBody, "context_management").Exists(),
+ "count_tokens API-key + 客户端未带 beta token → body strip")
+}
+
+// count_tokens passthrough preserve 测试
+func TestBuildCountTokensRequestAnthropicAPIKeyPassthrough_PreservesContextManagementWhenClientHeaderHasBeta(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages/count_tokens", nil)
+ c.Request.Header.Set("Anthropic-Beta", "oauth-2025-04-20,context-management-2025-06-27,token-counting-2024-11-01")
+
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[{"type":"clear_thinking_20251015"}]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, err := svc.buildCountTokensRequestAnthropicAPIKeyPassthrough(
+ context.Background(), c, newAnthropicAPIKeyPassthroughAccountForBetaTest(), body, "token",
+ )
+ require.NoError(t, err)
+ require.True(t, gjson.GetBytes(readUpstreamBodyForTest(t, req), "context_management").Exists(),
+ "count_tokens passthrough + 客户端带 context-management beta → 字段保留")
+}
+
+func TestBuildUpstreamRequest_APIKeyHaikuWithContextManagement_StripsField(t *testing.T) {
+ // API-key + haiku + body 带 context_management + 客户端 header 未带 context-management beta
+ // → final beta 不含 → body 字段被 strip
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
+ c.Request.Header.Set("Anthropic-Beta", "interleaved-thinking-2025-05-14")
+
+ account := &Account{ID: 404, Platform: PlatformAnthropic, Type: AccountTypeAPIKey,
+ Credentials: map[string]any{"api_key": "sk-ant-xxx"},
+ Status: StatusActive, Schedulable: true,
+ }
+ body := []byte(`{"model":"claude-haiku-4-5","context_management":{"edits":[]},"messages":[]}`)
+ svc := &GatewayService{cfg: &config.Config{}}
+ req, _, err := svc.buildUpstreamRequest(
+ context.Background(), c, account, body,
+ "sk-ant-xxx", "apikey", "claude-haiku-4-5", false, false,
+ )
+ require.NoError(t, err)
+
+ outBody := readUpstreamBodyForTest(t, req)
+ require.False(t, gjson.GetBytes(outBody, "context_management").Exists(),
+ "API-key + haiku + 客户端未带 beta token → body 字段必须被 strip")
+}
diff --git a/backend/internal/service/gateway_forward_as_chat_completions.go b/backend/internal/service/gateway_forward_as_chat_completions.go
index 7ac77f77..1df450d6 100644
--- a/backend/internal/service/gateway_forward_as_chat_completions.go
+++ b/backend/internal/service/gateway_forward_as_chat_completions.go
@@ -119,7 +119,7 @@ func (s *GatewayService) ForwardAsChatCompletions(
// 10. Build upstream request
upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, reqStream)
- upstreamReq, err := s.buildUpstreamRequest(upstreamCtx, c, account, anthropicBody, token, tokenType, mappedModel, reqStream, shouldMimicClaudeCode)
+ upstreamReq, _, err := s.buildUpstreamRequest(upstreamCtx, c, account, anthropicBody, token, tokenType, mappedModel, reqStream, shouldMimicClaudeCode)
releaseUpstreamCtx()
if err != nil {
return nil, fmt.Errorf("build upstream request: %w", err)
@@ -148,7 +148,7 @@ func (s *GatewayService) ForwardAsChatCompletions(
// 12. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -166,7 +166,7 @@ func (s *GatewayService) ForwardAsChatCompletions(
Message: upstreamMsg,
})
if s.rateLimitService != nil {
- s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
+ s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, mappedModel)
}
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
diff --git a/backend/internal/service/gateway_forward_as_responses.go b/backend/internal/service/gateway_forward_as_responses.go
index 8f8a1e94..22951b88 100644
--- a/backend/internal/service/gateway_forward_as_responses.go
+++ b/backend/internal/service/gateway_forward_as_responses.go
@@ -116,7 +116,7 @@ func (s *GatewayService) ForwardAsResponses(
// 10. Build upstream request
upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, reqStream)
- upstreamReq, err := s.buildUpstreamRequest(upstreamCtx, c, account, anthropicBody, token, tokenType, mappedModel, reqStream, shouldMimicClaudeCode)
+ upstreamReq, _, err := s.buildUpstreamRequest(upstreamCtx, c, account, anthropicBody, token, tokenType, mappedModel, reqStream, shouldMimicClaudeCode)
releaseUpstreamCtx()
if err != nil {
return nil, fmt.Errorf("build upstream request: %w", err)
@@ -145,7 +145,7 @@ func (s *GatewayService) ForwardAsResponses(
// 12. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -163,7 +163,7 @@ func (s *GatewayService) ForwardAsResponses(
Message: upstreamMsg,
})
if s.rateLimitService != nil {
- s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
+ s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, mappedModel)
}
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go
index 72832837..7a6acaac 100644
--- a/backend/internal/service/gateway_multiplatform_test.go
+++ b/backend/internal/service/gateway_multiplatform_test.go
@@ -156,7 +156,7 @@ func (m *mockAccountRepoForPlatform) ListSchedulableUngroupedByPlatforms(ctx con
func (m *mockAccountRepoForPlatform) SetRateLimited(ctx context.Context, id int64, resetAt time.Time) error {
return nil
}
-func (m *mockAccountRepoForPlatform) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time) error {
+func (m *mockAccountRepoForPlatform) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error {
return nil
}
func (m *mockAccountRepoForPlatform) SetOverloaded(ctx context.Context, id int64, until time.Time) error {
@@ -1229,6 +1229,106 @@ func TestGatewayService_selectAccountWithMixedScheduling(t *testing.T) {
require.Equal(t, int64(2), acc.ID, "应选择优先级最高的账户(包含启用混合调度的antigravity)")
})
+ t.Run("混合调度-Gemini家族限流后跳过Antigravity账户", func(t *testing.T) {
+ resetAt := time.Now().Add(10 * time.Minute).Format(time.RFC3339)
+ repo := &mockAccountRepoForPlatform{
+ accounts: []Account{
+ {
+ ID: 1,
+ Platform: PlatformAntigravity,
+ Priority: 1,
+ Status: StatusActive,
+ Schedulable: true,
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": resetAt,
+ },
+ },
+ },
+ },
+ {
+ ID: 2,
+ Platform: PlatformAntigravity,
+ Priority: 1,
+ Status: StatusActive,
+ Schedulable: true,
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": resetAt,
+ },
+ },
+ },
+ },
+ {
+ ID: 3,
+ Platform: PlatformAntigravity,
+ Priority: 2,
+ Status: StatusActive,
+ Schedulable: true,
+ Extra: map[string]any{"mixed_scheduling": true},
+ },
+ },
+ accountsByID: map[int64]*Account{},
+ }
+ for i := range repo.accounts {
+ repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
+ }
+
+ svc := &GatewayService{
+ accountRepo: repo,
+ cache: &mockGatewayCacheForPlatform{},
+ cfg: testConfig(),
+ }
+
+ acc, err := svc.selectAccountWithMixedScheduling(ctx, nil, "", "gemini-3-pro-preview", nil, PlatformGemini)
+ require.NoError(t, err)
+ require.NotNil(t, acc)
+ require.Equal(t, int64(3), acc.ID)
+ })
+
+ t.Run("混合调度-Gemini家族限流不影响Claude调度", func(t *testing.T) {
+ resetAt := time.Now().Add(10 * time.Minute).Format(time.RFC3339)
+ repo := &mockAccountRepoForPlatform{
+ accounts: []Account{
+ {
+ ID: 1,
+ Platform: PlatformAntigravity,
+ Priority: 1,
+ Status: StatusActive,
+ Schedulable: true,
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": resetAt,
+ },
+ },
+ },
+ },
+ {ID: 2, Platform: PlatformAnthropic, Priority: 2, Status: StatusActive, Schedulable: true},
+ },
+ accountsByID: map[int64]*Account{},
+ }
+ for i := range repo.accounts {
+ repo.accountsByID[repo.accounts[i].ID] = &repo.accounts[i]
+ }
+
+ svc := &GatewayService{
+ accountRepo: repo,
+ cache: &mockGatewayCacheForPlatform{},
+ cfg: testConfig(),
+ }
+
+ acc, err := svc.selectAccountWithMixedScheduling(ctx, nil, "", "claude-sonnet-4-5", nil, PlatformAnthropic)
+ require.NoError(t, err)
+ require.NotNil(t, acc)
+ require.Equal(t, int64(1), acc.ID)
+ })
+
t.Run("混合调度-路由优先选择路由账号", func(t *testing.T) {
groupID := int64(30)
requestedModel := "claude-sonnet-4-5"
diff --git a/backend/internal/service/gateway_oauth_metadata_test.go b/backend/internal/service/gateway_oauth_metadata_test.go
index ed6f1887..b172dc6e 100644
--- a/backend/internal/service/gateway_oauth_metadata_test.go
+++ b/backend/internal/service/gateway_oauth_metadata_test.go
@@ -14,8 +14,6 @@ func TestBuildOAuthMetadataUserID_FallbackWithoutAccountUUID(t *testing.T) {
Model: "claude-sonnet-4-5",
Stream: true,
MetadataUserID: "",
- System: nil,
- Messages: nil,
}
account := &Account{
diff --git a/backend/internal/service/gateway_request.go b/backend/internal/service/gateway_request.go
index 498336a4..dc59611c 100644
--- a/backend/internal/service/gateway_request.go
+++ b/backend/internal/service/gateway_request.go
@@ -12,6 +12,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/domain"
"github.com/Wei-Shaw/sub2api/internal/pkg/antigravity"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
)
@@ -50,6 +51,164 @@ type SessionContext struct {
APIKeyID int64
}
+type jsonRange struct {
+ start int // 原始请求体中的起始偏移(闭区间)
+ end int // 原始请求体中的结束偏移(开区间)
+ kind gjson.Type // JSON 值类型,用于调用方做轻量分支
+}
+
+type RequestBodyRef struct {
+ data []byte
+}
+
+func NewRequestBodyRef(data []byte) *RequestBodyRef {
+ return &RequestBodyRef{data: data}
+}
+
+func (b *RequestBodyRef) Bytes() []byte {
+ if b == nil {
+ return nil
+ }
+ return b.data
+}
+
+func (b *RequestBodyRef) Len() int {
+ if b == nil {
+ return 0
+ }
+ return len(b.data)
+}
+
+func (b *RequestBodyRef) Replace(data []byte) {
+ if b == nil {
+ return
+ }
+ b.data = data
+}
+
+func missingJSONRange() jsonRange {
+ return jsonRange{start: -1, end: -1}
+}
+
+func rangeFromResult(r gjson.Result) jsonRange {
+ if r.Raw == "" || r.Index <= 0 {
+ return missingJSONRange()
+ }
+ end := r.Index + len(r.Raw)
+ if end < r.Index {
+ return missingJSONRange()
+ }
+ return jsonRange{start: r.Index, end: end, kind: r.Type}
+}
+
+func (r jsonRange) exists() bool {
+ return r.start >= 0 && r.end >= r.start
+}
+
+// clearGatewayRequestDerivedState 清空绑定当前 body 的轻量派生字段,防止 ReplaceBody 后读到旧值。
+func clearGatewayRequestDerivedState(parsed *ParsedRequest) {
+ if parsed == nil {
+ return
+ }
+ parsed.Model = ""
+ parsed.Stream = false
+ parsed.MetadataUserID = ""
+ parsed.HasSystem = false
+ parsed.ThinkingEnabled = false
+ parsed.OutputEffort = ""
+ parsed.MaxTokens = 0
+ parsed.systemRange = missingJSONRange()
+ parsed.messagesRange = missingJSONRange()
+}
+
+func clearGatewayRequestRanges(parsed *ParsedRequest) {
+ if parsed == nil {
+ return
+ }
+ parsed.HasSystem = false
+ parsed.systemRange = missingJSONRange()
+ parsed.messagesRange = missingJSONRange()
+}
+
+func setGatewayRequestRanges(parsed *ParsedRequest, protocol string, jsonStr string) {
+ if parsed == nil {
+ return
+ }
+ switch protocol {
+ case domain.PlatformGemini:
+ if sysParts := gjson.Get(jsonStr, "systemInstruction.parts"); sysParts.Exists() && sysParts.IsArray() {
+ parsed.systemRange = rangeFromResult(sysParts)
+ }
+ if contents := gjson.Get(jsonStr, "contents"); contents.Exists() && contents.IsArray() {
+ parsed.messagesRange = rangeFromResult(contents)
+ }
+ default:
+ if sys := gjson.Get(jsonStr, "system"); sys.Exists() {
+ parsed.HasSystem = true
+ parsed.systemRange = rangeFromResult(sys)
+ }
+ if msgs := gjson.Get(jsonStr, "messages"); msgs.Exists() && msgs.IsArray() {
+ parsed.messagesRange = rangeFromResult(msgs)
+ }
+ }
+}
+
+// parseGatewayRequestCurrentBody 只做标量和 raw range 轻量解析,不恢复 system/messages 对象图。
+func parseGatewayRequestCurrentBody(parsed *ParsedRequest, protocol string) error {
+ if parsed == nil || parsed.Body == nil {
+ return fmt.Errorf("empty request body")
+ }
+
+ bodyBytes := parsed.Body.Bytes()
+ if !gjson.ValidBytes(bodyBytes) {
+ return fmt.Errorf("invalid json")
+ }
+
+ // 只在当前函数内零拷贝读取 JSON 字段;ReplaceBody 后必须重新进入本函数刷新派生状态。
+ jsonStr := *(*string)(unsafe.Pointer(&bodyBytes))
+ clearGatewayRequestDerivedState(parsed)
+ parsed.protocol = protocol
+
+ modelResult := gjson.Get(jsonStr, "model")
+ if modelResult.Exists() {
+ if modelResult.Type != gjson.String {
+ return fmt.Errorf("invalid model field type")
+ }
+ parsed.Model = modelResult.String()
+ }
+
+ streamResult := gjson.Get(jsonStr, "stream")
+ if streamResult.Exists() {
+ if streamResult.Type != gjson.True && streamResult.Type != gjson.False {
+ return fmt.Errorf("invalid stream field type")
+ }
+ parsed.Stream = streamResult.Bool()
+ }
+
+ parsed.MetadataUserID = gjson.Get(jsonStr, "metadata.user_id").String()
+
+ thinkingType := gjson.Get(jsonStr, "thinking.type").String()
+ parsed.ThinkingEnabled = thinkingType == "enabled" || thinkingType == "adaptive"
+
+ parsed.OutputEffort = strings.TrimSpace(gjson.Get(jsonStr, "output_config.effort").String())
+
+ maxTokensResult := gjson.Get(jsonStr, "max_tokens")
+ if maxTokensResult.Exists() && maxTokensResult.Type == gjson.Number {
+ f := maxTokensResult.Float()
+ if !math.IsNaN(f) && !math.IsInf(f, 0) && f == math.Trunc(f) &&
+ f <= float64(math.MaxInt) && f >= float64(math.MinInt) {
+ parsed.MaxTokens = int(f)
+ }
+ }
+
+ setGatewayRequestRanges(parsed, protocol, jsonStr)
+ return nil
+}
+
+func refreshGatewayRequestRanges(parsed *ParsedRequest, protocol string) error {
+ return parseGatewayRequestCurrentBody(parsed, protocol)
+}
+
// ParsedRequest 保存网关请求的预解析结果
//
// 性能优化说明:
@@ -63,18 +222,20 @@ type SessionContext struct {
// 2. 将解析结果 ParsedRequest 传递给 Service 层
// 3. 避免重复 json.Unmarshal,减少 CPU 和内存开销
type ParsedRequest struct {
- Body []byte // 原始请求体(保留用于转发)
+ Body *RequestBodyRef // 原始请求体引用(保留用于转发);替换内容请走 ReplaceBody
Model string // 请求的模型名称
Stream bool // 是否为流式请求
MetadataUserID string // metadata.user_id(用于会话亲和)
- System any // system 字段内容
- Messages []any // messages 数组
HasSystem bool // 是否包含 system 字段(包含 null 也视为显式传入)
ThinkingEnabled bool // 是否开启 thinking(部分平台会影响最终模型名)
OutputEffort string // output_config.effort(Claude API 的推理强度控制)
MaxTokens int // max_tokens 值(用于探测请求拦截)
SessionContext *SessionContext // 可选:请求上下文区分因子(nil 时行为不变)
+ protocol string // 当前 Body 的协议格式,用于 Body 替换后刷新 raw range
+ systemRange jsonRange // system/systemInstruction.parts 的 raw JSON 范围,绑定 Body 当前内容
+ messagesRange jsonRange // messages/contents 的 raw JSON 范围,绑定 Body 当前内容
+
// GroupID 请求所属分组 ID(来自 API Key)
GroupID *int64
@@ -129,119 +290,92 @@ func normalizeSessionUserAgentFallback(raw string) string {
// ParseGatewayRequest 解析网关请求体并返回结构化结果。
// protocol 指定请求协议格式(domain.PlatformAnthropic / domain.PlatformGemini),
// 不同协议使用不同的 system/messages 字段名。
-func ParseGatewayRequest(body []byte, protocol string) (*ParsedRequest, error) {
- // 保持与旧实现一致:请求体必须是合法 JSON。
- // 注意:gjson.GetBytes 对非法 JSON 不会报错,因此需要显式校验。
- if !gjson.ValidBytes(body) {
- return nil, fmt.Errorf("invalid json")
+func ParseGatewayRequest(body *RequestBodyRef, protocol string) (*ParsedRequest, error) {
+ parsed := &ParsedRequest{Body: body}
+ if err := parseGatewayRequestCurrentBody(parsed, protocol); err != nil {
+ return nil, err
}
-
- // 性能:
- // - gjson.GetBytes 会把匹配的 Raw/Str 安全复制成 string(对于巨大 messages 会产生额外拷贝)。
- // - 这里将 body 通过 unsafe 零拷贝视为 string,仅在本函数内使用,且 body 不会被修改。
- jsonStr := *(*string)(unsafe.Pointer(&body))
-
- parsed := &ParsedRequest{
- Body: body,
- }
-
- // --- gjson 提取简单字段(避免完整 Unmarshal) ---
-
- // model: 需要严格类型校验,非 string 返回错误
- modelResult := gjson.Get(jsonStr, "model")
- if modelResult.Exists() {
- if modelResult.Type != gjson.String {
- return nil, fmt.Errorf("invalid model field type")
- }
- parsed.Model = modelResult.String()
- }
-
- // stream: 需要严格类型校验,非 bool 返回错误
- streamResult := gjson.Get(jsonStr, "stream")
- if streamResult.Exists() {
- if streamResult.Type != gjson.True && streamResult.Type != gjson.False {
- return nil, fmt.Errorf("invalid stream field type")
- }
- parsed.Stream = streamResult.Bool()
- }
-
- // metadata.user_id: 直接路径提取,不需要严格类型校验
- parsed.MetadataUserID = gjson.Get(jsonStr, "metadata.user_id").String()
-
- // thinking.type: enabled/adaptive 都视为开启
- thinkingType := gjson.Get(jsonStr, "thinking.type").String()
- if thinkingType == "enabled" || thinkingType == "adaptive" {
- parsed.ThinkingEnabled = true
- }
-
- // output_config.effort: Claude API 的推理强度控制参数
- parsed.OutputEffort = strings.TrimSpace(gjson.Get(jsonStr, "output_config.effort").String())
-
- // max_tokens: 仅接受整数值
- maxTokensResult := gjson.Get(jsonStr, "max_tokens")
- if maxTokensResult.Exists() && maxTokensResult.Type == gjson.Number {
- f := maxTokensResult.Float()
- if !math.IsNaN(f) && !math.IsInf(f, 0) && f == math.Trunc(f) &&
- f <= float64(math.MaxInt) && f >= float64(math.MinInt) {
- parsed.MaxTokens = int(f)
- }
- }
-
- // --- system/messages 提取 ---
- // 避免把整个 body Unmarshal 到 map(会产生大量 map/接口分配)。
- // 使用 gjson 抽取目标字段的 Raw,再对该子树进行 Unmarshal。
-
- switch protocol {
- case domain.PlatformGemini:
- // Gemini 原生格式: systemInstruction.parts / contents
- if sysParts := gjson.Get(jsonStr, "systemInstruction.parts"); sysParts.Exists() && sysParts.IsArray() {
- var parts []any
- if err := json.Unmarshal(sliceRawFromBody(body, sysParts), &parts); err != nil {
- return nil, err
- }
- parsed.System = parts
- }
-
- if contents := gjson.Get(jsonStr, "contents"); contents.Exists() && contents.IsArray() {
- var msgs []any
- if err := json.Unmarshal(sliceRawFromBody(body, contents), &msgs); err != nil {
- return nil, err
- }
- parsed.Messages = msgs
- }
- default:
- // Anthropic / OpenAI 格式: system / messages
- // system 字段只要存在就视为显式提供(即使为 null),
- // 以避免客户端传 null 时被默认 system 误注入。
- if sys := gjson.Get(jsonStr, "system"); sys.Exists() {
- parsed.HasSystem = true
- switch sys.Type {
- case gjson.Null:
- parsed.System = nil
- case gjson.String:
- // 与 encoding/json 的 Unmarshal 行为一致:返回解码后的字符串。
- parsed.System = sys.String()
- default:
- var system any
- if err := json.Unmarshal(sliceRawFromBody(body, sys), &system); err != nil {
- return nil, err
- }
- parsed.System = system
- }
- }
-
- if msgs := gjson.Get(jsonStr, "messages"); msgs.Exists() && msgs.IsArray() {
- var messages []any
- if err := json.Unmarshal(sliceRawFromBody(body, msgs), &messages); err != nil {
- return nil, err
- }
- parsed.Messages = messages
- }
- }
-
return parsed, nil
}
+func (p *ParsedRequest) raw(r jsonRange) []byte {
+ if p == nil || p.Body == nil || !r.exists() {
+ return nil
+ }
+ body := p.Body.Bytes()
+ if r.end > len(body) {
+ return nil
+ }
+ return body[r.start:r.end]
+}
+
+func (p *ParsedRequest) SystemRaw() []byte {
+ return p.raw(p.systemRange)
+}
+
+func (p *ParsedRequest) MessagesRaw() []byte {
+ return p.raw(p.messagesRange)
+}
+
+func (p *ParsedRequest) DecodeSystem(dst any) error {
+ raw := p.SystemRaw()
+ if len(raw) == 0 {
+ return nil
+ }
+ return json.Unmarshal(raw, dst)
+}
+
+func (p *ParsedRequest) DecodeMessages(dst any) error {
+ raw := p.MessagesRaw()
+ if len(raw) == 0 {
+ return nil
+ }
+ return json.Unmarshal(raw, dst)
+}
+
+func (p *ParsedRequest) SystemValue() (any, bool) {
+ raw := p.SystemRaw()
+ if len(raw) == 0 {
+ return nil, false
+ }
+ var system any
+ if err := json.Unmarshal(raw, &system); err != nil {
+ return nil, false
+ }
+ return system, true
+}
+
+// CloneForBody 为单次账号尝试创建独立 body 视图,避免 failover 复用已改写的 ParsedRequest。
+func (p *ParsedRequest) CloneForBody(body []byte) (*ParsedRequest, error) {
+ if p == nil {
+ return nil, fmt.Errorf("parse request: empty request")
+ }
+ clone := *p
+ clone.Body = NewRequestBodyRef(body)
+ clone.OnUpstreamAccepted = nil
+ if err := refreshGatewayRequestRanges(&clone, clone.protocol); err != nil {
+ return nil, err
+ }
+ return &clone, nil
+}
+
+// ReplaceBody 统一刷新当前 body 和 raw range,保证后续 helper 读取的是最新请求体。
+func (p *ParsedRequest) ReplaceBody(data []byte) error {
+ if p == nil {
+ return fmt.Errorf("parse request: empty request")
+ }
+ if p.Body == nil {
+ p.Body = NewRequestBodyRef(data)
+ } else {
+ p.Body.Replace(data)
+ }
+ if err := refreshGatewayRequestRanges(p, p.protocol); err != nil {
+ clearGatewayRequestRanges(p)
+ return err
+ }
+ return nil
+}
+
// sliceRawFromBody 返回 Result.Raw 对应的原始字节切片。
// 优先使用 Result.Index 直接从 body 切片,避免对大字段(如 messages)产生额外拷贝。
// 当 Index 不可用时,退化为复制(理论上极少发生)。
@@ -665,6 +799,69 @@ func removeThinkingDependentContextStrategies(body []byte) []byte {
return body
}
+// anthropicBetaContextManagementToken 是 context_management 字段受的 beta token。
+// 与 claude.BetaContextManagement 保持一致;在本文件本地定义以避免震荡
+// claude package 的该常量含义。
+const anthropicBetaContextManagementToken = "context-management-2025-06-27"
+
+// sanitizeAnthropicBodyForBetaTokens 是对 Anthropic 直连路径上 body↔beta header
+// **能力维度**对称约束的统一实现,与 Bedrock 路径的
+// `sanitizeBedrockFieldsForBetaTokens` 对称。
+//
+// 问题场景:
+// - context_management 是 Claude Code CLI 2.1.87+ 默认携带的 beta 字段
+// (含 clear_thinking_20251015 等清理策略)
+// - 其被 Anthropic 上游接受的前提是 anthropic-beta header 含
+// `context-management-2025-06-27`
+// - 若两侧不一致上游 Pydantic schema 拒收:
+// "context_management: Extra inputs are not permitted"
+//
+// 本函数按最终发送的 anthropic-beta header 决定是否保留 body 中的
+// context_management 字段:缺 beta token → strip。这将限制完全建立在
+// "能力维度" 上,与 model 名 / token type / mimicry 子路径无关。
+//
+// 调用约束:必须在 CCH 签名之前调用,否则签名 hash 与最终 body
+// 不一致,上游会以 third-party 拒收。
+//
+// 返回 (sanitized, changed):changed 表示是否发生实际删除,供调用方决定
+// 是否重用原 body 引用。
+func sanitizeAnthropicBodyForBetaTokens(body []byte, anthropicBetaHeader string) ([]byte, bool) {
+ if len(body) == 0 {
+ return body, false
+ }
+ if !gjson.GetBytes(body, "context_management").Exists() {
+ return body, false
+ }
+ if anthropicBetaTokensContains(anthropicBetaHeader, anthropicBetaContextManagementToken) {
+ return body, false
+ }
+ if b, err := sjson.DeleteBytes(body, "context_management"); err == nil {
+ return b, true
+ } else {
+ // 不应发生:gjson 刚验证过字段存在 + body 是合法 JSON。如果 sjson 仍报错,
+ // 调用方会拿到 (body, false),但此前 computeFinalAnthropicBeta 已按“strip 后”
+ // 计算了 finalBeta——两侧会不一致。记录 warning 最小限度提醒运维。
+ logger.LegacyPrintf("service.gateway",
+ "[CtxMgmtSanitize] sjson.DeleteBytes failed unexpectedly: %v (body len=%d). "+
+ "body and final anthropic-beta header may be out of sync.", err, len(body))
+ }
+ return body, false
+}
+
+// anthropicBetaTokensContains 检测逗号分隔的 anthropic-beta header 是否含指定 token。
+// 宋体空格宽容;区分大小写(Anthropic beta token 始终是小写)。
+func anthropicBetaTokensContains(header, token string) bool {
+ if header == "" || token == "" {
+ return false
+ }
+ for _, part := range strings.Split(header, ",") {
+ if strings.TrimSpace(part) == token {
+ return true
+ }
+ }
+ return false
+}
+
// FilterSignatureSensitiveBlocksForRetry is a stronger retry filter for cases where upstream errors indicate
// signature/thought_signature validation issues involving tool blocks.
//
diff --git a/backend/internal/service/gateway_request_test.go b/backend/internal/service/gateway_request_test.go
index 045dc66c..288c031c 100644
--- a/backend/internal/service/gateway_request_test.go
+++ b/backend/internal/service/gateway_request_test.go
@@ -10,24 +10,25 @@ import (
"github.com/Wei-Shaw/sub2api/internal/domain"
"github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
)
func TestParseGatewayRequest(t *testing.T) {
body := []byte(`{"model":"claude-3-7-sonnet","stream":true,"metadata":{"user_id":"session_123e4567-e89b-12d3-a456-426614174000"},"system":[{"type":"text","text":"hello","cache_control":{"type":"ephemeral"}}],"messages":[{"content":"hi"}]}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, "claude-3-7-sonnet", parsed.Model)
require.True(t, parsed.Stream)
require.Equal(t, "session_123e4567-e89b-12d3-a456-426614174000", parsed.MetadataUserID)
require.True(t, parsed.HasSystem)
- require.NotNil(t, parsed.System)
- require.Len(t, parsed.Messages, 1)
+ require.NotEmpty(t, parsed.SystemRaw())
+ require.NotEmpty(t, parsed.MessagesRaw())
require.False(t, parsed.ThinkingEnabled)
}
func TestParseGatewayRequest_ThinkingEnabled(t *testing.T) {
body := []byte(`{"model":"claude-sonnet-4-5","thinking":{"type":"enabled"},"messages":[{"content":"hi"}]}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, "claude-sonnet-4-5", parsed.Model)
require.True(t, parsed.ThinkingEnabled)
@@ -35,7 +36,7 @@ func TestParseGatewayRequest_ThinkingEnabled(t *testing.T) {
func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) {
body := []byte(`{"model":"claude-sonnet-4-5","thinking":{"type":"adaptive"},"messages":[{"content":"hi"}]}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, "claude-sonnet-4-5", parsed.Model)
require.True(t, parsed.ThinkingEnabled)
@@ -43,36 +44,36 @@ func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) {
func TestParseGatewayRequest_MaxTokens(t *testing.T) {
body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, 1, parsed.MaxTokens)
}
func TestParseGatewayRequest_MaxTokensNonIntegralIgnored(t *testing.T) {
body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1.5}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
require.Equal(t, 0, parsed.MaxTokens)
}
func TestParseGatewayRequest_SystemNull(t *testing.T) {
body := []byte(`{"model":"claude-3","system":null}`)
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.NoError(t, err)
// 显式传入 system:null 也应视为“字段已存在”,避免默认 system 被注入。
require.True(t, parsed.HasSystem)
- require.Nil(t, parsed.System)
+ require.Equal(t, []byte("null"), parsed.SystemRaw())
}
func TestParseGatewayRequest_InvalidModelType(t *testing.T) {
body := []byte(`{"model":123}`)
- _, err := ParseGatewayRequest(body, "")
+ _, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.Error(t, err)
}
func TestParseGatewayRequest_InvalidStreamType(t *testing.T) {
body := []byte(`{"stream":"true"}`)
- _, err := ParseGatewayRequest(body, "")
+ _, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
require.Error(t, err)
}
@@ -86,11 +87,11 @@ func TestParseGatewayRequest_GeminiContents(t *testing.T) {
{"role": "user", "parts": [{"text": "How are you?"}]}
]
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
- require.Len(t, parsed.Messages, 3, "should parse contents as Messages")
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), 3, "should parse contents as Messages")
require.False(t, parsed.HasSystem, "Gemini format should not set HasSystem")
- require.Nil(t, parsed.System, "no systemInstruction means nil System")
+ require.Nil(t, parsed.SystemRaw(), "no systemInstruction means nil System")
}
func TestParseGatewayRequest_GeminiSystemInstruction(t *testing.T) {
@@ -102,16 +103,13 @@ func TestParseGatewayRequest_GeminiSystemInstruction(t *testing.T) {
{"role": "user", "parts": [{"text": "Hello"}]}
]
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
- require.NotNil(t, parsed.System, "should parse systemInstruction.parts as System")
- parts, ok := parsed.System.([]any)
- require.True(t, ok)
- require.Len(t, parts, 1)
- partMap, ok := parts[0].(map[string]any)
- require.True(t, ok)
- require.Equal(t, "You are a helpful assistant.", partMap["text"])
- require.Len(t, parsed.Messages, 1)
+ system := gjson.ParseBytes(parsed.SystemRaw())
+ require.True(t, system.IsArray(), "should parse systemInstruction.parts as System")
+ require.Len(t, system.Array(), 1)
+ require.Equal(t, "You are a helpful assistant.", system.Get("0.text").String())
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), 1)
}
func TestParseGatewayRequest_GeminiWithModel(t *testing.T) {
@@ -119,10 +117,10 @@ func TestParseGatewayRequest_GeminiWithModel(t *testing.T) {
"model": "gemini-2.5-pro",
"contents": [{"role": "user", "parts": [{"text": "test"}]}]
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.Equal(t, "gemini-2.5-pro", parsed.Model)
- require.Len(t, parsed.Messages, 1)
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), 1)
}
func TestParseGatewayRequest_GeminiIgnoresAnthropicFields(t *testing.T) {
@@ -132,25 +130,25 @@ func TestParseGatewayRequest_GeminiIgnoresAnthropicFields(t *testing.T) {
"messages": [{"role": "user", "content": "ignored"}],
"contents": [{"role": "user", "parts": [{"text": "real content"}]}]
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
require.False(t, parsed.HasSystem, "Gemini protocol should not parse Anthropic system field")
- require.Nil(t, parsed.System, "no systemInstruction = nil System")
- require.Len(t, parsed.Messages, 1, "should use contents, not messages")
+ require.Nil(t, parsed.SystemRaw(), "no systemInstruction = nil System")
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), 1, "should use contents, not messages")
}
func TestParseGatewayRequest_GeminiEmptyContents(t *testing.T) {
body := []byte(`{"contents": []}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
- require.Empty(t, parsed.Messages)
+ require.Empty(t, gjson.ParseBytes(parsed.MessagesRaw()).Array())
}
func TestParseGatewayRequest_GeminiNoContents(t *testing.T) {
body := []byte(`{"model": "gemini-2.5-flash"}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformGemini)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
require.NoError(t, err)
- require.Nil(t, parsed.Messages)
+ require.Nil(t, parsed.MessagesRaw())
require.Equal(t, "gemini-2.5-flash", parsed.Model)
}
@@ -162,14 +160,13 @@ func TestParseGatewayRequest_AnthropicIgnoresGeminiFields(t *testing.T) {
"contents": [{"role": "user", "parts": [{"text": "ignored"}]}],
"systemInstruction": {"parts": [{"text": "ignored"}]}
}`)
- parsed, err := ParseGatewayRequest(body, domain.PlatformAnthropic)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic)
require.NoError(t, err)
require.True(t, parsed.HasSystem)
- require.Equal(t, "real system", parsed.System)
- require.Len(t, parsed.Messages, 1)
- msg, ok := parsed.Messages[0].(map[string]any)
- require.True(t, ok)
- require.Equal(t, "real content", msg["content"])
+ require.Equal(t, "real system", gjson.ParseBytes(parsed.SystemRaw()).String())
+ messages := gjson.ParseBytes(parsed.MessagesRaw()).Array()
+ require.Len(t, messages, 1)
+ require.Equal(t, "real content", messages[0].Get("content").String())
}
func TestFilterThinkingBlocks(t *testing.T) {
@@ -897,7 +894,7 @@ func TestParseGatewayRequest_TypeValidation(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- _, err := ParseGatewayRequest([]byte(tt.body), "")
+ _, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
if tt.wantErr {
require.Error(t, err)
if tt.errSubstr != "" {
@@ -959,7 +956,7 @@ func TestParseGatewayRequest_OptionalFieldsMissing(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- parsed, err := ParseGatewayRequest([]byte(tt.body), "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
require.NoError(t, err)
require.Equal(t, tt.wantModel, parsed.Model)
@@ -970,10 +967,10 @@ func TestParseGatewayRequest_OptionalFieldsMissing(t *testing.T) {
require.Equal(t, tt.wantMaxTokens, parsed.MaxTokens)
if tt.wantMessagesNil {
- require.Nil(t, parsed.Messages)
+ require.Nil(t, parsed.MessagesRaw())
}
if tt.wantMessagesLen > 0 {
- require.Len(t, parsed.Messages, tt.wantMessagesLen)
+ require.Len(t, gjson.ParseBytes(parsed.MessagesRaw()).Array(), tt.wantMessagesLen)
}
})
}
@@ -1023,7 +1020,7 @@ func TestParseGatewayRequest_MaxTokensBoundary(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- parsed, err := ParseGatewayRequest([]byte(tt.body), "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
if tt.wantErr {
require.Error(t, err)
return
@@ -1040,7 +1037,7 @@ func TestParseGatewayRequest_MaxTokensBoundary(t *testing.T) {
// 核心路径:先 Unmarshal 到 map[string]any,再逐字段提取。
func parseGatewayRequestOld(body []byte, protocol string) (*ParsedRequest, error) {
parsed := &ParsedRequest{
- Body: body,
+ Body: NewRequestBodyRef(body),
}
var req map[string]any
@@ -1087,25 +1084,8 @@ func parseGatewayRequestOld(body []byte, protocol string) (*ParsedRequest, error
}
}
- // system / messages(按协议分支)
- switch protocol {
- case domain.PlatformGemini:
- if sysInst, ok := req["systemInstruction"].(map[string]any); ok {
- if parts, ok := sysInst["parts"].([]any); ok {
- parsed.System = parts
- }
- }
- if contents, ok := req["contents"].([]any); ok {
- parsed.Messages = contents
- }
- default:
- if system, ok := req["system"]; ok {
- parsed.HasSystem = true
- parsed.System = system
- }
- if messages, ok := req["messages"].([]any); ok {
- parsed.Messages = messages
- }
+ if err := refreshGatewayRequestRanges(parsed, protocol); err != nil {
+ return nil, err
}
return parsed, nil
@@ -1151,7 +1131,7 @@ func BenchmarkParseGatewayRequest_New_Small(b *testing.B) {
b.SetBytes(int64(len(data)))
b.ResetTimer()
for i := 0; i < b.N; i++ {
- _, _ = ParseGatewayRequest(data, "")
+ _, _ = ParseGatewayRequest(NewRequestBodyRef(data), "")
}
}
@@ -1203,7 +1183,7 @@ func TestParseGatewayRequest_OutputEffort(t *testing.T) {
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
- parsed, err := ParseGatewayRequest([]byte(tt.body), "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(tt.body)), "")
require.NoError(t, err)
require.Equal(t, tt.wantEffort, parsed.OutputEffort)
})
@@ -1245,6 +1225,6 @@ func BenchmarkParseGatewayRequest_New_Large(b *testing.B) {
b.SetBytes(int64(len(data)))
b.ResetTimer()
for i := 0; i < b.N; i++ {
- _, _ = ParseGatewayRequest(data, "")
+ _, _ = ParseGatewayRequest(NewRequestBodyRef(data), "")
}
}
diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go
index 6106391b..812780dc 100644
--- a/backend/internal/service/gateway_service.go
+++ b/backend/internal/service/gateway_service.go
@@ -23,6 +23,7 @@ import (
"sync/atomic"
"syscall"
"time"
+ "unsafe"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/claude"
@@ -56,6 +57,8 @@ const (
defaultModelsListCacheTTL = 15 * time.Second
postUsageBillingTimeout = 15 * time.Second
debugGatewayBodyEnv = "SUB2API_DEBUG_GATEWAY_BODY"
+ // 上游错误体只需要提取错误 JSON/日志摘要,默认 512KiB 避免错误风暴叠加大请求体。
+ gatewayUpstreamErrorBodyReadLimit int64 = 512 << 10
)
const (
@@ -96,15 +99,20 @@ var (
modelsListCacheMissTotal atomic.Int64
modelsListCacheStoreTotal atomic.Int64
+ // Deprecated: flusher_enabled=true 后不再增长(仅 flag=false 降级直写路径使用);新主路径见 FlusherMetrics。remove after 2026-09。
// userPlatformQuotaDBIncrErrorTotal 统计 finalizePostUsageBilling 异步 goroutine
// 中 IncrementUsageWithReset 失败次数。Redis 已成功累加 + DB 写失败意味着
// Redis cache TTL 过期或被清后该笔 cost 会丢失(与实际消费偏差)。
// oncall 通过 GatewayUserPlatformQuotaIncrStats() 暴露给 ops 面板做阈值告警。
userPlatformQuotaDBIncrErrorTotal atomic.Int64
+ // Deprecated: flusher_enabled=true 后不再增长(仅 flag=false 降级直写路径使用);新主路径见 FlusherMetrics。remove after 2026-09。
// userPlatformQuotaDBIncrLegacyErrorTotal 统计 legacy postUsageBilling
// (applyUsageBilling 在 repo==nil 时 fallback)路径下的失败次数;
// 与 DB Incr 失败分开计数,便于区分"主路径暂时故障"vs"基础设施长期未配齐"。
userPlatformQuotaDBIncrLegacyErrorTotal atomic.Int64
+ // userPlatformQuotaSentinelSetCacheErrorTotal 统计 checkUserPlatformQuotaEligibility
+ // 在 DB 无行时回填 sentinel cache entry 写 Redis 失败的次数(phase A)。
+ userPlatformQuotaSentinelSetCacheErrorTotal atomic.Int64
)
func GatewayWindowCostPrefetchStats() (cacheHit, cacheMiss, batchSQL, fallback, errCount int64) {
@@ -127,13 +135,32 @@ func GatewayModelsListCacheStats() (cacheHit, cacheMiss, store int64) {
return modelsListCacheHitTotal.Load(), modelsListCacheMissTotal.Load(), modelsListCacheStoreTotal.Load()
}
-// GatewayUserPlatformQuotaIncrStats 返回 (mainPathErr, legacyPathErr)。
+// GatewayUserPlatformQuotaIncrStats 返回 (mainPathErr, legacyPathErr, sentinelSetErr)。
// mainPathErr:finalizePostUsageBilling 异步 goroutine 写 DB 失败累计次数;
-// legacyPathErr:postUsageBilling fallback 路径写 DB 失败累计次数。
+// legacyPathErr:postUsageBilling fallback 路径写 DB 失败累计次数;
+// sentinelSetErr:DB 无行时回填 sentinel cache entry 写 Redis 失败累计次数。
// ops 监控面板可以按"持续上升斜率"做告警阈值。
-func GatewayUserPlatformQuotaIncrStats() (mainPathErr, legacyPathErr int64) {
+func GatewayUserPlatformQuotaIncrStats() (mainPathErr, legacyPathErr, sentinelSetErr int64) {
return userPlatformQuotaDBIncrErrorTotal.Load(),
- userPlatformQuotaDBIncrLegacyErrorTotal.Load()
+ userPlatformQuotaDBIncrLegacyErrorTotal.Load(),
+ userPlatformQuotaSentinelSetCacheErrorTotal.Load()
+}
+
+// GatewayUserPlatformQuotaFlusherStats 暴露 flusher 运行指标供 ops/health 面板查询。
+func GatewayUserPlatformQuotaFlusherStats(f *UserPlatformQuotaUsageFlusher) map[string]int64 {
+ if f == nil || f.metrics == nil {
+ return nil
+ }
+ m := f.metrics
+ return map[string]int64{
+ "flush_success": m.FlushSuccessTotal.Load(),
+ "flush_error": m.FlushErrorTotal.Load(),
+ "flush_batch_size": m.FlushBatchSizeTotal.Load(),
+ "flush_latency_ms_max": m.FlushLatencyMsMax.Load(),
+ "dirty_readd": m.DirtyReaddTotal.Load(),
+ "dirty_lost": m.DirtyLostTotal.Load(),
+ "flush_fk_violation": m.FlushFKViolationTotal.Load(),
+ }
}
func openAIStreamEventIsTerminal(data string) bool {
@@ -724,31 +751,10 @@ func (s *GatewayService) GenerateSessionHash(parsed *ParsedRequest) string {
_, _ = combined.WriteString(strconv.FormatInt(parsed.SessionContext.APIKeyID, 10))
_, _ = combined.WriteString("|")
}
- if parsed.System != nil {
- systemText := s.extractTextFromSystem(parsed.System)
- if systemText != "" {
- _, _ = combined.WriteString(systemText)
- }
- }
- for _, msg := range parsed.Messages {
- if m, ok := msg.(map[string]any); ok {
- if content, exists := m["content"]; exists {
- // Anthropic: messages[].content
- if msgText := s.extractTextFromContent(content); msgText != "" {
- _, _ = combined.WriteString(msgText)
- }
- } else if parts, ok := m["parts"].([]any); ok {
- // Gemini: contents[].parts[].text
- for _, part := range parts {
- if partMap, ok := part.(map[string]any); ok {
- if text, ok := partMap["text"].(string); ok {
- _, _ = combined.WriteString(text)
- }
- }
- }
- }
- }
+ if systemText := extractTextFromSystemRaw(parsed.SystemRaw()); systemText != "" {
+ _, _ = combined.WriteString(systemText)
}
+ appendMessageTextsFromRaw(&combined, parsed.MessagesRaw())
if combined.Len() > 0 {
hash := s.hashContent(combined.String())
slog.Info("sticky.hash_source",
@@ -823,82 +829,135 @@ func (s *GatewayService) extractCacheableContent(parsed *ParsedRequest) string {
return ""
}
- var builder strings.Builder
-
- // 检查 system 中的 cacheable 内容
- if system, ok := parsed.System.([]any); ok {
- for _, part := range system {
- if partMap, ok := part.(map[string]any); ok {
- if cc, ok := partMap["cache_control"].(map[string]any); ok {
- if cc["type"] == "ephemeral" {
- if text, ok := partMap["text"].(string); ok {
- _, _ = builder.WriteString(text)
- }
- }
- }
- }
- }
+ systemText := extractCacheableTextFromSystemRaw(parsed.SystemRaw())
+ if messageText := extractCacheableTextFromMessagesRaw(parsed.MessagesRaw()); messageText != "" {
+ return messageText
}
- systemText := builder.String()
-
- // 检查 messages 中的 cacheable 内容
- for _, msg := range parsed.Messages {
- if msgMap, ok := msg.(map[string]any); ok {
- if msgContent, ok := msgMap["content"].([]any); ok {
- for _, part := range msgContent {
- if partMap, ok := part.(map[string]any); ok {
- if cc, ok := partMap["cache_control"].(map[string]any); ok {
- if cc["type"] == "ephemeral" {
- return s.extractTextFromContent(msgMap["content"])
- }
- }
- }
- }
- }
- }
- }
-
return systemText
}
-func (s *GatewayService) extractTextFromSystem(system any) string {
- switch v := system.(type) {
- case string:
- return v
- case []any:
- var texts []string
- for _, part := range v {
- if partMap, ok := part.(map[string]any); ok {
- if text, ok := partMap["text"].(string); ok {
- texts = append(texts, text)
- }
- }
+func parseRawJSONView(raw []byte) gjson.Result {
+ if len(raw) == 0 {
+ return gjson.Result{}
+ }
+ // 这里只做同步只读解析,避免 gjson.ParseBytes 为大 messages/contents 复制整段 raw。
+ return gjson.Parse(*(*string)(unsafe.Pointer(&raw)))
+}
+
+func extractTextFromSystemRaw(raw []byte) string {
+ system := parseRawJSONView(raw)
+ switch system.Type {
+ case gjson.String:
+ return system.String()
+ case gjson.JSON:
+ if !system.IsArray() {
+ return ""
}
- return strings.Join(texts, "")
+ var builder strings.Builder
+ system.ForEach(func(_, part gjson.Result) bool {
+ if text := part.Get("text").String(); text != "" {
+ _, _ = builder.WriteString(text)
+ }
+ return true
+ })
+ return builder.String()
}
return ""
}
-func (s *GatewayService) extractTextFromContent(content any) string {
- switch v := content.(type) {
- case string:
- return v
- case []any:
- var texts []string
- for _, part := range v {
- if partMap, ok := part.(map[string]any); ok {
- if partMap["type"] == "text" {
- if text, ok := partMap["text"].(string); ok {
- texts = append(texts, text)
- }
+func extractTextFromContentRaw(content gjson.Result) string {
+ switch content.Type {
+ case gjson.String:
+ return content.String()
+ case gjson.JSON:
+ if !content.IsArray() {
+ return ""
+ }
+ var builder strings.Builder
+ content.ForEach(func(_, part gjson.Result) bool {
+ if part.Get("type").String() == "text" {
+ if text := part.Get("text").String(); text != "" {
+ _, _ = builder.WriteString(text)
}
}
- }
- return strings.Join(texts, "")
+ return true
+ })
+ return builder.String()
}
return ""
}
+func appendMessageTextsFromRaw(builder *strings.Builder, raw []byte) {
+ if builder == nil || len(raw) == 0 {
+ return
+ }
+ messages := parseRawJSONView(raw)
+ if !messages.IsArray() {
+ return
+ }
+ messages.ForEach(func(_, msg gjson.Result) bool {
+ if content := msg.Get("content"); content.Exists() {
+ _, _ = builder.WriteString(extractTextFromContentRaw(content))
+ return true
+ }
+ parts := msg.Get("parts")
+ if parts.IsArray() {
+ parts.ForEach(func(_, part gjson.Result) bool {
+ if text := part.Get("text").String(); text != "" {
+ _, _ = builder.WriteString(text)
+ }
+ return true
+ })
+ }
+ return true
+ })
+}
+
+func extractCacheableTextFromSystemRaw(raw []byte) string {
+ system := parseRawJSONView(raw)
+ if !system.IsArray() {
+ return ""
+ }
+ var builder strings.Builder
+ system.ForEach(func(_, part gjson.Result) bool {
+ if part.Get("cache_control.type").String() == "ephemeral" {
+ if text := part.Get("text").String(); text != "" {
+ _, _ = builder.WriteString(text)
+ }
+ }
+ return true
+ })
+ return builder.String()
+}
+
+func extractCacheableTextFromMessagesRaw(raw []byte) string {
+ messages := parseRawJSONView(raw)
+ if !messages.IsArray() {
+ return ""
+ }
+ var text string
+ messages.ForEach(func(_, msg gjson.Result) bool {
+ content := msg.Get("content")
+ if !content.IsArray() {
+ return true
+ }
+ found := false
+ content.ForEach(func(_, part gjson.Result) bool {
+ if part.Get("cache_control.type").String() == "ephemeral" {
+ found = true
+ return false
+ }
+ return true
+ })
+ if found {
+ text = extractTextFromContentRaw(content)
+ return false
+ }
+ return true
+ })
+ return text
+}
+
func (s *GatewayService) hashContent(content string) string {
h := xxhash.Sum64String(content)
return strconv.FormatUint(h, 36)
@@ -1155,6 +1214,12 @@ func normalizeClaudeOAuthRequestBody(body []byte, modelID string, opts claudeOAu
// context_management:thinking.type 为 enabled/adaptive 时,真实 CLI 会自动
// 附带 {"edits":[{"type":"clear_thinking_20251015","keep":"all"}]}。
// 客户端显式传了就透传;否则按 CLI 行为补齐。
+ //
+ // 注:本函数不按 model 名决定是否保留 context_management。“最终 beta
+ // header 不含 context-management-2025-06-27 时 strip 字段”的能力维度
+ // 对称约束由 sanitizeAnthropicBodyForBetaTokens 在 buildUpstreamRequest /
+ // buildCountTokensRequest 层统一执行,与 Bedrock 路径的
+ // sanitizeBedrockFieldsForBetaTokens 对称。
if !gjson.GetBytes(out, "context_management").Exists() {
thinkingType := gjson.GetBytes(out, "thinking.type").String()
if thinkingType == "enabled" || thinkingType == "adaptive" {
@@ -1254,7 +1319,7 @@ func (s *GatewayService) applyClaudeCodeOAuthMimicryToBody(
systemRewritten := false
if !strings.Contains(strings.ToLower(model), "haiku") {
- body = rewriteSystemForNonClaudeCode(body, systemRaw)
+ body = rewriteSystemForNonClaudeCode(body, normalizeSystemParam(systemRaw))
systemRewritten = true
}
@@ -4370,12 +4435,12 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
// Web Search 模拟:纯 web_search 请求时,直接调用搜索 API 构造响应
- if account != nil && s.shouldEmulateWebSearch(ctx, account, parsed.GroupID, parsed.Body) {
+ if account != nil && s.shouldEmulateWebSearch(ctx, account, parsed.GroupID, parsed.Body.Bytes()) {
return s.handleWebSearchEmulation(ctx, c, account, parsed)
}
if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() {
- passthroughBody := parsed.Body
+ passthroughBody := parsed.Body.Bytes()
passthroughModel := parsed.Model
if passthroughModel != "" {
if mappedModel := account.GetMappedModel(passthroughModel); mappedModel != passthroughModel {
@@ -4386,6 +4451,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
return s.forwardAnthropicAPIKeyPassthroughWithInput(ctx, c, account, anthropicPassthroughForwardInput{
Body: passthroughBody,
+ Parsed: parsed,
RequestModel: passthroughModel,
OriginalModel: parsed.Model,
RequestStream: parsed.Stream,
@@ -4411,7 +4477,14 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
c.Set(betaPolicyFilterSetKey, filterSet)
}
- body := parsed.Body
+ body := parsed.Body.Bytes()
+ replaceBody := func(next []byte) error {
+ if err := parsed.ReplaceBody(next); err != nil {
+ return fmt.Errorf("rewrite request body: %w", err)
+ }
+ body = parsed.Body.Bytes()
+ return nil
+ }
reqModel := parsed.Model
reqStream := parsed.Stream
originalModel := reqModel
@@ -4444,7 +4517,10 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
// Parrot 的 transform_request 从不检查客户端 system 内容,直接覆盖。
systemRewritten := false
if !strings.Contains(strings.ToLower(reqModel), "haiku") {
- body = rewriteSystemForNonClaudeCode(body, parsed.System)
+ systemRaw, _ := parsed.SystemValue()
+ if err := replaceBody(rewriteSystemForNonClaudeCode(body, systemRaw)); err != nil {
+ return nil, err
+ }
systemRewritten = true
}
@@ -4466,22 +4542,34 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
}
- body, reqModel = normalizeClaudeOAuthRequestBody(body, reqModel, normalizeOpts)
+ var normalizedBody []byte
+ normalizedBody, reqModel = normalizeClaudeOAuthRequestBody(body, reqModel, normalizeOpts)
+ if err := replaceBody(normalizedBody); err != nil {
+ return nil, err
+ }
// D/E/F: 可选 messages cache 策略 + 工具名混淆 + tools[-1] 断点
// 与 forward_as_chat_completions / forward_as_responses 路径对齐,
// 原生 /v1/messages 路径也走同一套可配置字段级改写。
- body = s.rewriteMessageCacheControlIfEnabled(ctx, body)
+ if err := replaceBody(s.rewriteMessageCacheControlIfEnabled(ctx, body)); err != nil {
+ return nil, err
+ }
if rw := buildToolNameRewriteFromBody(body); rw != nil {
- body = applyToolNameRewriteToBody(body, rw)
+ if err := replaceBody(applyToolNameRewriteToBody(body, rw)); err != nil {
+ return nil, err
+ }
c.Set(toolNameRewriteKey, rw)
} else {
- body = applyToolsLastCacheBreakpoint(body)
+ if err := replaceBody(applyToolsLastCacheBreakpoint(body)); err != nil {
+ return nil, err
+ }
}
}
// 强制执行 cache_control 块数量限制(最多 4 个)
- body = enforceCacheControlLimit(body)
+ if err := replaceBody(enforceCacheControlLimit(body)); err != nil {
+ return nil, err
+ }
// 应用模型映射:
// - APIKey 账号:使用账号级别的显式映射(如果配置),否则透传原始模型名
@@ -4515,13 +4603,18 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
}
if mappedModel != reqModel {
// 替换请求体中的模型名
- body = s.replaceModelInBody(body, mappedModel)
+ if err := replaceBody(s.replaceModelInBody(body, mappedModel)); err != nil {
+ return nil, err
+ }
reqModel = mappedModel
+ parsed.Model = mappedModel
logger.LegacyPrintf("service.gateway", "Model mapping applied: %s -> %s (account: %s, source=%s)", originalModel, mappedModel, account.Name, mappingSource)
}
if s.shouldInjectAnthropicCacheTTL1h(ctx, account) {
- body = injectAnthropicCacheControlTTL1h(body)
+ if err := replaceBody(injectAnthropicCacheControlTTL1h(body)); err != nil {
+ return nil, err
+ }
}
// 获取凭证
@@ -4545,19 +4638,24 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
logger.LegacyPrintf("service.gateway", "[Forward] Using account: ID=%d Name=%s Platform=%s Type=%s TLSFingerprint=%v Proxy=%s",
account.ID, account.Name, account.Platform, account.Type, tlsProfile, proxyURL)
// Pre-filter: strip empty text blocks (including nested in tool_result) to prevent upstream 400.
- body = StripEmptyTextBlocks(body)
+ if err := replaceBody(StripEmptyTextBlocks(body)); err != nil {
+ return nil, err
+ }
// 重试循环
var resp *http.Response
+ lastWireBody := body
retryStart := time.Now()
for attempt := 1; attempt <= maxRetryAttempts; attempt++ {
// 构建上游请求(每次重试需要重新构建,因为请求体需要重新读取)
upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, reqStream)
- upstreamReq, err := s.buildUpstreamRequest(upstreamCtx, c, account, body, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
+ upstreamReq, wireBody, err := s.buildUpstreamRequest(upstreamCtx, c, account, body, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
releaseUpstreamCtx()
if err != nil {
return nil, err
}
+ // 记录本次实际发送的 wire body;只有请求成功后才写回 ParsedRequest,避免 400 retry 基于已签名 CCH 再改写。
+ lastWireBody = wireBody
// 发送请求
resp, err = s.httpUpstream.DoWithTLS(upstreamReq, proxyURL, account.ID, account.Concurrency, tlsProfile)
@@ -4589,7 +4687,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
// 优先检测thinking block签名错误(400)并重试一次
if resp.StatusCode == 400 {
- respBody, readErr := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, readErr := s.readUpstreamErrorBody(resp)
if readErr == nil {
_ = resp.Body.Close()
@@ -4635,18 +4733,24 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
filteredBody := FilterThinkingBlocksForRetry(body)
retryCtx, releaseRetryCtx := detachStreamUpstreamContext(ctx, reqStream)
- retryReq, buildErr := s.buildUpstreamRequest(retryCtx, c, account, filteredBody, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
+ retryReq, retryWireBody, buildErr := s.buildUpstreamRequest(retryCtx, c, account, filteredBody, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
releaseRetryCtx()
if buildErr == nil {
retryResp, retryErr := s.httpUpstream.DoWithTLS(retryReq, proxyURL, account.ID, account.Concurrency, tlsProfile)
if retryErr == nil {
if retryResp.StatusCode < 400 {
+ // 重试请求被上游接受后同步 ParsedRequest,保证 usage/日志看到真实请求体。
+ lastWireBody = retryWireBody
+ if err := replaceBody(retryWireBody); err != nil {
+ _ = retryResp.Body.Close()
+ return nil, err
+ }
logger.LegacyPrintf("service.gateway", "Account %d: thinking block retry succeeded (blocks downgraded)", account.ID)
resp = retryResp
break
}
- retryRespBody, retryReadErr := io.ReadAll(io.LimitReader(retryResp.Body, 2<<20))
+ retryRespBody, retryReadErr := s.readUpstreamErrorBody(retryResp)
_ = retryResp.Body.Close()
if retryReadErr == nil && retryResp.StatusCode == 400 && s.isSignatureErrorPattern(ctx, account, retryRespBody) {
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
@@ -4670,11 +4774,19 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
logger.LegacyPrintf("service.gateway", "Account %d: signature retry still failing and looks tool-related, retrying with tool blocks downgraded", account.ID)
filteredBody2 := FilterSignatureSensitiveBlocksForRetry(body)
retryCtx2, releaseRetryCtx2 := detachStreamUpstreamContext(ctx, reqStream)
- retryReq2, buildErr2 := s.buildUpstreamRequest(retryCtx2, c, account, filteredBody2, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
+ retryReq2, retryWireBody2, buildErr2 := s.buildUpstreamRequest(retryCtx2, c, account, filteredBody2, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
releaseRetryCtx2()
if buildErr2 == nil {
retryResp2, retryErr2 := s.httpUpstream.DoWithTLS(retryReq2, proxyURL, account.ID, account.Concurrency, tlsProfile)
if retryErr2 == nil {
+ if retryResp2.StatusCode < 400 {
+ // 二阶段工具块降级成功时也必须更新当前 body。
+ lastWireBody = retryWireBody2
+ if err := replaceBody(retryWireBody2); err != nil {
+ _ = retryResp2.Body.Close()
+ return nil, err
+ }
+ }
resp = retryResp2
break
}
@@ -4741,11 +4853,19 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
if applied && time.Since(retryStart) < maxRetryElapsed {
logger.LegacyPrintf("service.gateway", "Account %d: detected budget_tokens constraint error, retrying with rectified budget (budget_tokens=%d, max_tokens=%d)", account.ID, BudgetRectifyBudgetTokens, BudgetRectifyMaxTokens)
budgetRetryCtx, releaseBudgetRetryCtx := detachStreamUpstreamContext(ctx, reqStream)
- budgetRetryReq, buildErr := s.buildUpstreamRequest(budgetRetryCtx, c, account, rectifiedBody, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
+ budgetRetryReq, budgetWireBody, buildErr := s.buildUpstreamRequest(budgetRetryCtx, c, account, rectifiedBody, token, tokenType, reqModel, reqStream, shouldMimicClaudeCode)
releaseBudgetRetryCtx()
if buildErr == nil {
budgetRetryResp, retryErr := s.httpUpstream.DoWithTLS(budgetRetryReq, proxyURL, account.ID, account.Concurrency, tlsProfile)
if retryErr == nil {
+ if budgetRetryResp.StatusCode < 400 {
+ // budget 修正请求成功后,ParsedRequest 也要描述被接受的修正版。
+ lastWireBody = budgetWireBody
+ if err := replaceBody(budgetWireBody); err != nil {
+ _ = budgetRetryResp.Body.Close()
+ return nil, err
+ }
+ }
resp = budgetRetryResp
break
}
@@ -4780,7 +4900,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
break
}
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
@@ -4827,7 +4947,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
// 处理重试耗尽的情况
if resp.StatusCode >= 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) {
if s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -4854,7 +4974,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
return s.handleRetryExhaustedError(ctx, resp, c, account)
@@ -4862,7 +4982,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
// 处理可切换账号的错误
if resp.StatusCode >= 400 && s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -4870,7 +4990,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
logger.LegacyPrintf("service.gateway", "[Forward] Upstream error (failover): Account=%d(%s) Status=%d RequestID=%s Body=%s",
account.ID, account.Name, resp.StatusCode, resp.Header.Get("x-request-id"), truncateString(string(respBody), 1000))
- s.handleFailoverSideEffects(ctx, resp, account)
+ s.handleFailoverSideEffects(ctx, resp, account, reqModel)
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
@@ -4888,16 +5008,16 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
if resp.StatusCode >= 400 {
// 可选:对部分 400 触发 failover(默认关闭以保持语义)
if resp.StatusCode == 400 && s.cfg != nil && s.cfg.Gateway.FailoverOn400 {
- respBody, readErr := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, readErr := s.readUpstreamErrorBody(resp)
if readErr != nil {
// ReadAll failed, fall back to normal error handling without consuming the stream
- return s.handleErrorResponse(ctx, resp, c, account)
+ return s.handleErrorResponse(ctx, resp, c, account, reqModel)
}
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -4933,15 +5053,22 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
} else {
logger.LegacyPrintf("service.gateway", "Account %d: 400 error, attempting failover", account.ID)
}
- s.handleFailoverSideEffects(ctx, resp, account)
+ s.handleFailoverSideEffects(ctx, resp, account, reqModel)
return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBody}
}
}
- return s.handleErrorResponse(ctx, resp, c, account)
+ return s.handleErrorResponse(ctx, resp, c, account, reqModel)
}
// 处理正常响应
+ if !bytes.Equal(lastWireBody, body) {
+ // 成功后再同步最终 wire body,避免失败重试从已签名 CCH 的 body 继续派生。
+ if err := replaceBody(lastWireBody); err != nil {
+ return nil, err
+ }
+ }
+
// 触发上游接受回调(提前释放串行锁,不等流完成)
if parsed.OnUpstreamAccepted != nil {
parsed.OnUpstreamAccepted()
@@ -4984,6 +5111,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
type anthropicPassthroughForwardInput struct {
Body []byte
+ Parsed *ParsedRequest
RequestModel string
OriginalModel string
RequestStream bool
@@ -5036,16 +5164,29 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
}
// Pre-filter: strip empty text blocks (including nested in tool_result) to prevent upstream 400.
input.Body = StripEmptyTextBlocks(input.Body)
+ if input.Parsed != nil {
+ // 透传分支也会改写实际 wire body,成功 usage hash 依赖这里同步当前 body。
+ if err := input.Parsed.ReplaceBody(input.Body); err != nil {
+ return nil, err
+ }
+ }
var resp *http.Response
retryStart := time.Now()
for attempt := 1; attempt <= maxRetryAttempts; attempt++ {
upstreamCtx, releaseUpstreamCtx := detachStreamUpstreamContext(ctx, input.RequestStream)
- upstreamReq, err := s.buildUpstreamRequestAnthropicAPIKeyPassthrough(upstreamCtx, c, account, input.Body, token)
+ upstreamReq, wireBody, err := s.buildUpstreamRequestAnthropicAPIKeyPassthrough(upstreamCtx, c, account, input.Body, token)
releaseUpstreamCtx()
if err != nil {
return nil, err
}
+ if input.Parsed != nil && !bytes.Equal(wireBody, input.Body) {
+ // build 阶段会按 beta 能力清理 body,发送前同步到 ParsedRequest 当前视图。
+ if err := input.Parsed.ReplaceBody(wireBody); err != nil {
+ return nil, err
+ }
+ input.Body = input.Parsed.Body.Bytes()
+ }
resp, err = s.httpUpstream.DoWithTLS(upstreamReq, proxyURL, account.ID, account.Concurrency, s.tlsFPProfileService.ResolveTLSProfile(account))
if err != nil {
@@ -5091,7 +5232,7 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
break
}
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
@@ -5129,7 +5270,7 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
if resp.StatusCode >= 400 && s.shouldRetryUpstreamError(account, resp.StatusCode) {
if s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -5156,21 +5297,21 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
return s.handleRetryExhaustedError(ctx, resp, c, account)
}
if resp.StatusCode >= 400 && s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
logger.LegacyPrintf("service.gateway", "[Anthropic Passthrough] Upstream error (failover): Account=%d(%s) Status=%d RequestID=%s Body=%s",
account.ID, account.Name, resp.StatusCode, resp.Header.Get("x-request-id"), truncateString(string(respBody), 1000))
- s.handleFailoverSideEffects(ctx, resp, account)
+ s.handleFailoverSideEffects(ctx, resp, account, input.RequestModel)
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
@@ -5190,12 +5331,12 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput(
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
if resp.StatusCode >= 400 {
- return s.handleErrorResponse(ctx, resp, c, account)
+ return s.handleErrorResponse(ctx, resp, c, account, input.RequestModel)
}
var usage *ClaudeUsage
@@ -5237,20 +5378,31 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough(
account *Account,
body []byte,
token string,
-) (*http.Request, error) {
+) (*http.Request, []byte, error) {
targetURL := claudeAPIURL
baseURL := account.GetBaseURL()
if baseURL != "" {
validatedURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = validatedURL + "/v1/messages?beta=true"
}
+ // 能力维度 body sanitize:透传路径上 anthropic-beta header 原样透传客户端值,
+ // 依此决定是否保留 body 中的 context_management。避免“客户端 body 带字段但
+ // header 忘记带 beta token”的客户端 bug 在透传场景下让上游 400。
+ clientBeta := ""
+ if c != nil && c.Request != nil {
+ clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta")
+ }
+ if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed {
+ body = sanitized
+ }
+
req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body))
if err != nil {
- return nil, err
+ return nil, nil, err
}
if c != nil && c.Request != nil {
@@ -5280,7 +5432,7 @@ func (s *GatewayService) buildUpstreamRequestAnthropicAPIKeyPassthrough(
setHeaderRaw(req.Header, "anthropic-version", "2023-06-01")
}
- return req, nil
+ return req, body, nil
}
func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough(
@@ -5694,7 +5846,7 @@ func (s *GatewayService) forwardBedrock(
) (*ForwardResult, error) {
reqModel := parsed.Model
reqStream := parsed.Stream
- body := parsed.Body
+ body := parsed.Body.Bytes()
region := bedrockRuntimeRegion(account)
mappedModel, ok := ResolveBedrockModelID(account, reqModel)
@@ -5762,6 +5914,11 @@ func (s *GatewayService) forwardBedrock(
return s.handleBedrockUpstreamErrors(ctx, resp, c, account)
}
+ // Bedrock 分支绕过通用 Forward 成功路径,这里保持上游接受回调语义一致。
+ if parsed.OnUpstreamAccepted != nil {
+ parsed.OnUpstreamAccepted()
+ }
+
// 响应处理
var usage *ClaudeUsage
var firstTokenMs *int
@@ -5865,7 +6022,7 @@ func (s *GatewayService) executeBedrockUpstream(
break
}
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
@@ -5910,7 +6067,7 @@ func (s *GatewayService) handleBedrockUpstreamErrors(
// retry exhausted + failover
if s.shouldRetryUpstreamError(account, resp.StatusCode) {
if s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -5929,7 +6086,7 @@ func (s *GatewayService) handleBedrockUpstreamErrors(
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
return s.handleRetryExhaustedError(ctx, resp, c, account)
@@ -5937,7 +6094,7 @@ func (s *GatewayService) handleBedrockUpstreamErrors(
// non-retryable failover
if s.shouldFailoverUpstreamError(resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -5953,7 +6110,7 @@ func (s *GatewayService) handleBedrockUpstreamErrors(
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
@@ -6038,9 +6195,10 @@ func (s *GatewayService) handleBedrockNonStreamingResponse(
return usage, nil
}
-func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, reqStream bool, mimicClaudeCode bool) (*http.Request, error) {
+func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, reqStream bool, mimicClaudeCode bool) (*http.Request, []byte, error) {
if account.Platform == PlatformAnthropic && account.Type == AccountTypeServiceAccount {
- return s.buildUpstreamRequestAnthropicVertex(ctx, c, account, body, token, modelID, reqStream)
+ req, err := s.buildUpstreamRequestAnthropicVertex(ctx, c, account, body, token, modelID, reqStream)
+ return req, body, err
}
// 确定目标URL
@@ -6050,18 +6208,18 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
if baseURL != "" {
validatedURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = validatedURL + "/v1/messages?beta=true"
}
} else if account.IsCustomBaseURLEnabled() {
customURL := account.GetCustomBaseURL()
if customURL == "" {
- return nil, fmt.Errorf("custom_base_url is enabled but not configured for account %d", account.ID)
+ return nil, nil, fmt.Errorf("custom_base_url is enabled but not configured for account %d", account.ID)
}
validatedURL, err := s.validateUpstreamBaseURL(customURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = s.buildCustomRelayURL(validatedURL, "/v1/messages", account)
}
@@ -6106,6 +6264,29 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
if fingerprint != nil {
body = syncBillingHeaderVersion(body, fingerprint.UserAgent)
}
+
+ // === 计算最终 anthropic-beta header(先于 body sanitize 与 CCH 签名)===
+ //
+ // 顺序约束:
+ // 1) 算 finalBeta(纯函数,不依赖 req.Header;mimicry 路径会忽略客户端 beta,
+ // 与原“OAuth + mimicClaudeCode 跳过白名单透传”行为对齐)
+ // 2) 按 finalBeta 做能力维度 body sanitize(如 context-management beta 缺失 →
+ // strip body.context_management,与 Bedrock 路径对称)
+ // 3) CCH 签名(必须使用 strip 后的 body,否则 hash 与最终 body 不一致 →
+ // 被 Anthropic 判 third-party)
+ // 4) NewRequest(body 至此最终敲定)
+ // 5) 透传白名单 / fingerprint / mimic header / 写入 finalBeta
+ policyFilterSet := s.getBetaPolicyFilterSet(ctx, c, account, modelID)
+ effectiveDropSet := mergeDropSets(policyFilterSet)
+ finalBetaHeader, finalBetaShouldSet := s.computeFinalAnthropicBeta(
+ tokenType, mimicClaudeCode, modelID, clientHeaders, body, effectiveDropSet,
+ )
+
+ // 能力维度 body sanitize:与最终 anthropic-beta header 对称
+ if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed {
+ body = sanitized
+ }
+
// CCH 签名:将 cch=00000 占位符替换为 xxHash64 签名(需在所有 body 修改之后)
if enableCCH {
body = signBillingHeaderCCH(body)
@@ -6113,7 +6294,7 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
req, err := http.NewRequestWithContext(ctx, "POST", targetURL, bytes.NewReader(body))
if err != nil {
- return nil, err
+ return nil, nil, err
}
// 设置认证头(保持原始大小写)
@@ -6156,46 +6337,18 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
applyClaudeOAuthHeaderDefaults(req)
}
- // Build effective drop set: merge static defaults with dynamic beta policy filter rules
- policyFilterSet := s.getBetaPolicyFilterSet(ctx, c, account, modelID)
- effectiveDropSet := mergeDropSets(policyFilterSet)
+ // OAuth + mimic Claude Code:强制注入 CLI 指纹相关 header
+ // (user-agent/x-stainless-*/x-app/Accept/x-stainless-helper-method/x-client-request-id)
+ if tokenType == "oauth" && mimicClaudeCode {
+ applyClaudeCodeMimicHeaders(req, reqStream)
+ }
- // 处理 anthropic-beta header(OAuth 账号需要包含 oauth beta)
- if tokenType == "oauth" {
- if mimicClaudeCode {
- // 非 Claude Code 客户端:按 opencode 的策略处理:
- // - 强制 Claude Code 指纹相关请求头(尤其是 user-agent/x-stainless/x-app)
- // - 保留 incoming beta 的同时,确保 OAuth 所需 beta 存在
- applyClaudeCodeMimicHeaders(req, reqStream)
-
- incomingBeta := getHeaderRaw(req.Header, "anthropic-beta")
- // Claude Code OAuth credentials are scoped to Claude Code.
- // Non-haiku models MUST include claude-code beta for Anthropic to recognize
- // this as a legitimate Claude Code request; without it, the request is
- // rejected as third-party ("out of extra usage").
- // Haiku models are exempt from third-party detection and don't need it.
- requiredBetas := []string{claude.BetaOAuth, claude.BetaInterleavedThinking}
- if !strings.Contains(strings.ToLower(modelID), "haiku") {
- requiredBetas = claude.FullClaudeCodeMimicryBetas()
- }
- setHeaderRaw(req.Header, "anthropic-beta", mergeAnthropicBetaDropping(requiredBetas, incomingBeta, effectiveDropSet))
- } else {
- // Claude Code 客户端:尽量透传原始 header,仅补齐 oauth beta
- clientBetaHeader := getHeaderRaw(req.Header, "anthropic-beta")
- setHeaderRaw(req.Header, "anthropic-beta", stripBetaTokensWithSet(s.getBetaHeader(modelID, clientBetaHeader), effectiveDropSet))
- }
- } else {
- // API-key accounts: apply beta policy filter to strip controlled tokens
- if existingBeta := getHeaderRaw(req.Header, "anthropic-beta"); existingBeta != "" {
- setHeaderRaw(req.Header, "anthropic-beta", stripBetaTokensWithSet(existingBeta, effectiveDropSet))
- } else if s.cfg != nil && s.cfg.Gateway.InjectBetaForAPIKey {
- // API-key:仅在请求显式使用 beta 特性且客户端未提供时,按需补齐(默认关闭)
- if requestNeedsBetaFeatures(body) {
- if beta := defaultAPIKeyBetaHeader(body); beta != "" {
- setHeaderRaw(req.Header, "anthropic-beta", beta)
- }
- }
- }
+ // 写入最终 anthropic-beta header
+ // 注:透传分支白名单可能写入了客户端 anthropic-beta,无条件 Del 一次再按 finalBeta
+ // 决定是否 set,确保 dropSet 过滤后的结果一定覆盖客户端原始值。
+ deleteHeaderAllForms(req.Header, "anthropic-beta")
+ if finalBetaShouldSet {
+ setHeaderRaw(req.Header, "anthropic-beta", finalBetaHeader)
}
// 同步 X-Claude-Code-Session-Id 头:取 body 中已处理的 metadata.user_id 的 session_id 覆盖
@@ -6226,7 +6379,7 @@ func (s *GatewayService) buildUpstreamRequest(ctx context.Context, c *gin.Contex
logClaudeMimicDebug(req, body, account, tokenType, mimicClaudeCode)
}
- return req, nil
+ return req, body, nil
}
func (s *GatewayService) buildUpstreamRequestAnthropicVertex(
@@ -6242,6 +6395,16 @@ func (s *GatewayService) buildUpstreamRequestAnthropicVertex(
if err != nil {
return nil, err
}
+
+ // 能力维度 sanitize:Vertex 路径上 anthropic-beta header 原样透传客户端值
+ // (下面白名单跳过 anthropic-version 但保留 anthropic-beta),依此决定是否
+ // 保留 body 中的 context_management,与 Anthropic 直连 / Bedrock 路径对称。
+ if c != nil && c.Request != nil {
+ clientBeta := getHeaderRaw(c.Request.Header, "anthropic-beta")
+ if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(vertexBody, clientBeta); changed {
+ vertexBody = sanitized
+ }
+ }
fullURL, err := buildVertexAnthropicURL(account.VertexProjectID(), account.VertexLocation(modelID), modelID, reqStream)
if err != nil {
return nil, err
@@ -6410,6 +6573,121 @@ func mergeAnthropicBetaDropping(required []string, incoming string, drop map[str
return strings.Join(out, ",")
}
+// computeFinalAnthropicBeta 计算发往上游的最终 anthropic-beta header 值。
+//
+// 设计动机:将原本在 buildUpstreamRequest 内联在一起、依赖 req.Header 的
+// anthropic-beta 计算逻辑抽成纯函数。这样调用方可以在 NewRequest 之前
+// 就提前拿到最终 beta header,进而能按它对 body 做能力维度 sanitize 后再做
+// CCH 签名——一举修复了以下之前由顺序依赖导致的能力维度 sanitize
+// 无法部署的问题(签名与最终 body 不一致可以被判 third-party)。
+//
+// 返回 (value, shouldSet):
+// - shouldSet=false 意为“不主动设置 anthropic-beta header”,与原代码“
+// API-key 账号 + 客户端未传 anthropic-beta + InjectBetaForAPIKey 未开启或
+// requestNeedsBetaFeatures=false”的行为对齐。
+// - shouldSet=true 时 value 可能为空字符串(例如客户端透传的 beta 被 dropSet
+// 全部过滤掉),这与原代码中 setHeaderRaw 的结果一致。
+//
+// clientHeaders 是客户端原始 HTTP header(通常为 c.Request.Header);nil 时按“客户端
+// 未传”处理。body 是已经 metadata 重写 / billing version sync 之后但未 sanitize 上游
+// 不兼容字段之前的版本。
+func (s *GatewayService) computeFinalAnthropicBeta(
+ tokenType string,
+ mimicClaudeCode bool,
+ modelID string,
+ clientHeaders http.Header,
+ body []byte,
+ effectiveDropSet map[string]struct{},
+) (string, bool) {
+ clientBeta := ""
+ if clientHeaders != nil {
+ clientBeta = getHeaderRaw(clientHeaders, "anthropic-beta")
+ }
+
+ if tokenType == "oauth" {
+ if mimicClaudeCode {
+ // mimic 路径:原代码跳过白名单透传,incomingBeta 总是空字符串。
+ // 这里传空 string 以严格对齐原行为。
+ requiredBetas := []string{claude.BetaOAuth, claude.BetaInterleavedThinking}
+ if !strings.Contains(strings.ToLower(modelID), "haiku") {
+ requiredBetas = claude.FullClaudeCodeMimicryBetas()
+ }
+ return mergeAnthropicBetaDropping(requiredBetas, "", effectiveDropSet), true
+ }
+ // 真 Claude Code 客户端透传路径
+ return stripBetaTokensWithSet(s.getBetaHeader(modelID, clientBeta), effectiveDropSet), true
+ }
+
+ // API-key accounts
+ if clientBeta != "" {
+ return stripBetaTokensWithSet(clientBeta, effectiveDropSet), true
+ }
+ if s.cfg != nil && s.cfg.Gateway.InjectBetaForAPIKey {
+ if requestNeedsBetaFeatures(body) {
+ if beta := defaultAPIKeyBetaHeader(body); beta != "" {
+ return beta, true
+ }
+ }
+ }
+ return "", false
+}
+
+// computeFinalCountTokensAnthropicBeta 是 count_tokens 路径上 anthropic-beta header 的
+// 计算纯函数。语义与 computeFinalAnthropicBeta 对齐,但备份了 count_tokens 独有的
+// 两条特殊规则:
+//
+// - OAuth mimic:requiredBetas 为 FullClaudeCodeMimicryBetas + BetaTokenCounting
+// (与 messages 不同的是:不按 haiku 排除;count_tokens 始终携带 token-counting beta)
+// - OAuth 透传 + 客户端未传 anthropic-beta:补齐 CountTokensBetaHeader
+// - OAuth 透传 + 客户端传了:补齐 BetaTokenCounting(如果未含)
+//
+// 返回语义同 computeFinalAnthropicBeta。
+func (s *GatewayService) computeFinalCountTokensAnthropicBeta(
+ tokenType string,
+ mimicClaudeCode bool,
+ modelID string,
+ clientHeaders http.Header,
+ body []byte,
+ effectiveDropSet map[string]struct{},
+) (string, bool) {
+ clientBeta := ""
+ if clientHeaders != nil {
+ clientBeta = getHeaderRaw(clientHeaders, "anthropic-beta")
+ }
+
+ if tokenType == "oauth" {
+ if mimicClaudeCode {
+ // 与原代码严格等价:original buildCountTokensRequest 在 count_tokens mimic
+ // 分支上**不**会跳过白名单透传(与 messages mimic 路径不同),所以
+ // incomingBeta = req.Header[anthropic-beta] = 客户端透传过来的 client beta。
+ // 重构后直接从 clientHeaders 拿同一个值,保持行为一致。
+ requiredBetas := append(claude.FullClaudeCodeMimicryBetas(), claude.BetaTokenCounting)
+ return mergeAnthropicBetaDropping(requiredBetas, clientBeta, effectiveDropSet), true
+ }
+ if clientBeta == "" {
+ return claude.CountTokensBetaHeader, true
+ }
+ beta := s.getBetaHeader(modelID, clientBeta)
+ if !strings.Contains(beta, claude.BetaTokenCounting) {
+ beta = beta + "," + claude.BetaTokenCounting
+ }
+ return stripBetaTokensWithSet(beta, effectiveDropSet), true
+ }
+
+ // API-key accounts
+ if clientBeta != "" {
+ return stripBetaTokensWithSet(clientBeta, effectiveDropSet), true
+ }
+ if s.cfg != nil && s.cfg.Gateway.InjectBetaForAPIKey {
+ if requestNeedsBetaFeatures(body) {
+ if beta := defaultAPIKeyBetaHeader(body); beta != "" {
+ return beta, true
+ }
+ }
+ }
+ return "", false
+}
+
// stripBetaTokens removes the given beta tokens from a comma-separated header value.
func stripBetaTokens(header string, tokens []string) string {
if header == "" || len(tokens) == 0 {
@@ -6959,8 +7237,19 @@ func isCountTokensUnsupported404(statusCode int, body []byte) bool {
return strings.Contains(msg, "count_tokens") && strings.Contains(msg, "not found")
}
-func (s *GatewayService) handleErrorResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account) (*ForwardResult, error) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+func (s *GatewayService) readUpstreamErrorBody(resp *http.Response) ([]byte, error) {
+ if resp == nil || resp.Body == nil {
+ return nil, nil
+ }
+ limit := gatewayUpstreamErrorBodyReadLimit
+ if s != nil && s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody && s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes > int(limit) {
+ limit = int64(s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes)
+ }
+ return io.ReadAll(io.LimitReader(resp.Body, limit))
+}
+
+func (s *GatewayService) handleErrorResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, requestedModel ...string) (*ForwardResult, error) {
+ body, _ := s.readUpstreamErrorBody(resp)
// 调试日志:打印上游错误响应
logger.LegacyPrintf("service.gateway", "[Forward] Upstream error (non-retryable): Account=%d(%s) Status=%d RequestID=%s Body=%s",
@@ -7006,7 +7295,11 @@ func (s *GatewayService) handleErrorResponse(ctx context.Context, resp *http.Res
// 处理上游错误,标记账号状态
shouldDisable := false
if s.rateLimitService != nil {
- shouldDisable = s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, body)
+ if len(requestedModel) > 0 {
+ shouldDisable = s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, requestedModel[0])
+ } else {
+ shouldDisable = s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, body)
+ }
}
if shouldDisable {
return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: body}
@@ -7109,7 +7402,7 @@ func (s *GatewayService) handleErrorResponse(ctx context.Context, resp *http.Res
}
func (s *GatewayService) handleRetryExhaustedSideEffects(ctx context.Context, resp *http.Response, account *Account) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body, _ := s.readUpstreamErrorBody(resp)
statusCode := resp.StatusCode
// OAuth/Setup Token 账号的 403:标记账号异常
@@ -7122,8 +7415,12 @@ func (s *GatewayService) handleRetryExhaustedSideEffects(ctx context.Context, re
}
}
-func (s *GatewayService) handleFailoverSideEffects(ctx context.Context, resp *http.Response, account *Account) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+func (s *GatewayService) handleFailoverSideEffects(ctx context.Context, resp *http.Response, account *Account, requestedModel ...string) {
+ body, _ := s.readUpstreamErrorBody(resp)
+ if len(requestedModel) > 0 {
+ s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, requestedModel[0])
+ return
+ }
s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, body)
}
@@ -7132,7 +7429,7 @@ func (s *GatewayService) handleFailoverSideEffects(ctx context.Context, resp *ht
// API Key 未配置错误码:仅返回错误,不标记账号
func (s *GatewayService) handleRetryExhaustedError(ctx context.Context, resp *http.Response, c *gin.Context, account *Account) (*ForwardResult, error) {
// Capture upstream error body before side-effects consume the stream.
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody, _ := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -7956,10 +8253,10 @@ func (s *GatewayService) getUserGroupRateMultiplier(ctx context.Context, userID,
return resolver.Resolve(ctx, userID, groupID, groupDefaultMultiplier)
}
-// RecordUsageInput 记录使用量的输入参数
+// RecordUsageInput 记录使用量的输入参数。
+// 异步 worker 只接收计费所需快照,不能持有 ParsedRequest/RequestBodyRef 这类大请求体引用。
type RecordUsageInput struct {
Result *ForwardResult
- ParsedRequest *ParsedRequest
APIKey *APIKey
User *User
Account *Account
@@ -8084,18 +8381,23 @@ func postUsageBilling(ctx context.Context, p *postUsageBillingParams, deps *bill
}
}
- // Platform quota DB-only 累加(与 finalizePostUsageBilling 行为对齐的兜底):
- // - 仅对 standard(余额)模式生效;订阅模式豁免
- // - 直接走 DB,不经 Redis Incr 队列:legacy 路径在 repo==nil(仓库未注入)
- // 时被触发,此时整套 billing repo 都不可用,没有"双队列"风险
- // - 失败仅记 ALERT log + counter,不阻断主扣费流程;与正常路径一致
- //
- // 历史背景:原 legacy path 完全跳过此累加,导致部署中如果 repo 偶然为 nil
- // 时用户消费可绕过 platform quota,存在静默资金风险。
+ // Platform quota 累加(legacy 兜底路径):仅对 standard(余额)模式生效;订阅模式豁免;仅对有 limit 的用户写
+ // - HasUserPlatformQuotaLimit 守卫:与正常路径对齐,无 limit 公司跳过
+ // - 新增 Redis 同步写:enforcement 走 Redis,legacy 路径也必须同步写,否则 preflight 看不到消费
+ // - flusher_enabled=false(降级):保留原有同步直写 DB
+ // - flusher_enabled=true:跳过直写 DB,由 flusher 异步批量刷(markDirty 在 IncrementUserPlatformQuotaUsage 内部完成)
+ // - 失败仅记 ALERT log + counter,不阻断主扣费流程
if !p.IsSubscriptionBill && p.Platform != "" && cost.ActualCost > 0 && p.User != nil && deps.userPlatformQuotaRepo != nil {
- if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(billingCtx, p.User.ID, p.Platform, cost.ActualCost, time.Now().UTC()); err != nil {
- userPlatformQuotaDBIncrLegacyErrorTotal.Add(1)
- logger.LegacyPrintf("service.gateway", "ALERT: legacy incr user platform quota DB failed user=%d platform=%s cost=%f: %v", p.User.ID, p.Platform, cost.ActualCost, err)
+ if deps.billingCacheService.HasUserPlatformQuotaLimit(billingCtx, p.User.ID, p.Platform) {
+ deps.billingCacheService.IncrementUserPlatformQuotaUsage(p.User.ID, p.Platform, cost.ActualCost)
+ if deps.cfg == nil || !deps.cfg.Database.UserPlatformQuotaFlusherEnabled {
+ // 降级路径:flusher 未启用时保留原有同步直写 DB
+ if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(billingCtx, p.User.ID, p.Platform, cost.ActualCost, time.Now().UTC()); err != nil {
+ userPlatformQuotaDBIncrLegacyErrorTotal.Add(1)
+ logger.LegacyPrintf("service.gateway", "ALERT: legacy incr user platform quota DB failed user=%d platform=%s cost=%f: %v", p.User.ID, p.Platform, cost.ActualCost, err)
+ }
+ }
+ // flusher_enabled=true:不直写 DB,flusher 异步批量刷
}
}
@@ -8245,30 +8547,38 @@ func finalizePostUsageBilling(ctx context.Context, p *postUsageBillingParams, de
deps.deferredService.ScheduleLastUsedUpdate(p.Account.ID)
- // Platform quota 累加:仅在 standard(余额)模式生效;订阅模式豁免
- // Redis 同步写 + DB 异步持久化:
+ // Platform quota 累加:仅在 standard(余额)模式生效;订阅模式豁免;仅对有 limit 的用户写
+ // Redis 同步写 + DB 异步持久化(flag=false 降级)或 flusher 异步刷(flag=true):
+ // - HasUserPlatformQuotaLimit 守卫:无 limit 的公司跳过,避免无效写入 + 浪费 Redis 容量
// - Redis 同步:确保下次 preflight 立即看到最新 usage,把 TOCTOU 超支窗口
// 限制在并发 in-flight 请求数量内(旧实现的异步入队会让超支无限累积直到 worker 处理)
- // - DB 异步:在独立 goroutine 中走 detached context,失败用 ALERT log 触发 oncall 对账
+ // - DB 异步(flusher_enabled=false):在独立 goroutine 中走 detached context,失败用 ALERT log 触发 oncall 对账
+ // - flusher_enabled=true:不直写 DB,由 flusher 异步批量刷(markDirty 已在 IncrementUserPlatformQuotaUsage 内部完成)
if !p.IsSubscriptionBill && p.Platform != "" && p.Cost.ActualCost > 0 && p.User != nil && deps.userPlatformQuotaRepo != nil {
- deps.billingCacheService.IncrementUserPlatformQuotaUsage(p.User.ID, p.Platform, p.Cost.ActualCost)
- dbCtx, dbCancel := detachUpstreamContext(ctx)
- userID, platform, cost := p.User.ID, p.Platform, p.Cost.ActualCost
- go func() {
- defer func() {
- if r := recover(); r != nil {
- logger.LegacyPrintf("service.gateway", "ALERT: panic in user platform quota incr goroutine user=%d platform=%s: %v", userID, platform, r)
- }
- }()
- defer dbCancel()
- if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(dbCtx, userID, platform, cost, time.Now().UTC()); err != nil {
- // 失败计数器:暴露给 GatewayUserPlatformQuotaIncrStats(),由 ops 面板做斜率告警。
- userPlatformQuotaDBIncrErrorTotal.Add(1)
- // ALERT 级别:DB 持久化失败意味着 Redis cache 失效后该笔 cost 永久丢失,
- // 用户配额视图与实际消费会偏差,oncall 需要据此对账或人工补录。
- logger.LegacyPrintf("service.gateway", "ALERT: incr user platform quota DB failed user=%d platform=%s cost=%f: %v", userID, platform, cost, err)
+ if deps.billingCacheService.HasUserPlatformQuotaLimit(ctx, p.User.ID, p.Platform) {
+ deps.billingCacheService.IncrementUserPlatformQuotaUsage(p.User.ID, p.Platform, p.Cost.ActualCost)
+ if deps.cfg == nil || !deps.cfg.Database.UserPlatformQuotaFlusherEnabled {
+ // 降级路径:flusher 未启用时保留原有异步直写 DB
+ dbCtx, dbCancel := detachUpstreamContext(ctx)
+ userID, platform, cost := p.User.ID, p.Platform, p.Cost.ActualCost
+ go func() {
+ defer func() {
+ if r := recover(); r != nil {
+ logger.LegacyPrintf("service.gateway", "ALERT: panic in user platform quota incr goroutine user=%d platform=%s: %v", userID, platform, r)
+ }
+ }()
+ defer dbCancel()
+ if err := deps.userPlatformQuotaRepo.IncrementUsageWithReset(dbCtx, userID, platform, cost, time.Now().UTC()); err != nil {
+ // 失败计数器:暴露给 GatewayUserPlatformQuotaIncrStats(),由 ops 面板做斜率告警。
+ userPlatformQuotaDBIncrErrorTotal.Add(1)
+ // ALERT 级别:DB 持久化失败意味着 Redis cache 失效后该笔 cost 永久丢失,
+ // 用户配额视图与实际消费会偏差,oncall 需要据此对账或人工补录。
+ logger.LegacyPrintf("service.gateway", "ALERT: incr user platform quota DB failed user=%d platform=%s cost=%f: %v", userID, platform, cost, err)
+ }
+ }()
}
- }()
+ // flusher_enabled=true:不直写 DB,flusher 异步批量刷
+ }
}
// Notification checks run async — all parameters are already captured,
@@ -8383,6 +8693,7 @@ type billingDeps struct {
deferredService *DeferredService
balanceNotifyService *BalanceNotifyService
userPlatformQuotaRepo UserPlatformQuotaRepository
+ cfg *config.Config
}
func (s *GatewayService) billingDeps() *billingDeps {
@@ -8394,6 +8705,7 @@ func (s *GatewayService) billingDeps() *billingDeps {
deferredService: s.deferredService,
balanceNotifyService: s.balanceNotifyService,
userPlatformQuotaRepo: s.userPlatformQuotaRepo,
+ cfg: s.cfg,
}
}
@@ -8422,15 +8734,8 @@ func writeUsageLogBestEffort(ctx context.Context, repo UsageLogRepository, usage
}
}
-// recordUsageOpts 内部选项,参数化 RecordUsage 与 RecordUsageWithLongContext 的差异点。
+// recordUsageOpts 内部选项,参数化普通计费与长上下文计费的差异点。
type recordUsageOpts struct {
- // Claude Max 策略所需的 ParsedRequest(可选,仅 Claude 路径传入)
- ParsedRequest *ParsedRequest
-
- // EnableClaudePath 启用 Claude 路径特有逻辑:
- // - Claude Max 缓存计费策略
- EnableClaudePath bool
-
// 长上下文计费(仅 Gemini 路径需要)
LongContextThreshold int
LongContextMultiplier float64
@@ -8453,9 +8758,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
APIKeyService: input.APIKeyService,
QuotaPlatform: input.QuotaPlatform,
ChannelUsageFields: input.ChannelUsageFields,
- }, &recordUsageOpts{
- EnableClaudePath: true,
- })
+ }, &recordUsageOpts{})
}
// RecordUsageLongContextInput 记录使用量的输入参数(支持长上下文双倍计费)
@@ -8521,9 +8824,7 @@ type recordUsageCoreInput struct {
}
// recordUsageCore 是 RecordUsage 和 RecordUsageWithLongContext 的统一实现。
-// opts 中的字段控制两者之间的差异行为:
-// - ParsedRequest != nil → 启用 Claude Max 缓存计费策略
-// - LongContextThreshold > 0 → Token 计费回退走 CalculateCostWithLongContext
+// LongContextThreshold > 0 时 Token 计费回退走 CalculateCostWithLongContext。
func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsageCoreInput, opts *recordUsageOpts) error {
result := input.Result
apiKey := input.APIKey
@@ -8988,7 +9289,7 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
}
if account != nil && account.IsAnthropicAPIKeyPassthroughEnabled() {
- passthroughBody := parsed.Body
+ passthroughBody := parsed.Body.Bytes()
if reqModel := parsed.Model; reqModel != "" {
if mappedModel := account.GetMappedModel(reqModel); mappedModel != reqModel {
passthroughBody = s.replaceModelInBody(passthroughBody, mappedModel)
@@ -9004,24 +9305,43 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
return nil
}
- body := parsed.Body
+ body := parsed.Body.Bytes()
+ replaceBody := func(next []byte) error {
+ if err := parsed.ReplaceBody(next); err != nil {
+ return fmt.Errorf("rewrite count_tokens body: %w", err)
+ }
+ body = parsed.Body.Bytes()
+ return nil
+ }
reqModel := parsed.Model
// Pre-filter: strip empty text blocks to prevent upstream 400.
- body = StripEmptyTextBlocks(body)
+ if err := replaceBody(StripEmptyTextBlocks(body)); err != nil {
+ return err
+ }
isClaudeCodeCT := IsClaudeCodeClient(ctx) || isClaudeCodeClient(c.GetHeader("User-Agent"), parsed.MetadataUserID)
shouldMimicClaudeCode := account.IsOAuth() && !isClaudeCodeCT
if shouldMimicClaudeCode {
normalizeOpts := claudeOAuthNormalizeOptions{stripSystemCacheControl: true}
- body, reqModel = normalizeClaudeOAuthRequestBody(body, reqModel, normalizeOpts)
+ var normalizedBody []byte
+ normalizedBody, reqModel = normalizeClaudeOAuthRequestBody(body, reqModel, normalizeOpts)
+ if err := replaceBody(normalizedBody); err != nil {
+ return err
+ }
- body = s.rewriteMessageCacheControlIfEnabled(ctx, body)
+ if err := replaceBody(s.rewriteMessageCacheControlIfEnabled(ctx, body)); err != nil {
+ return err
+ }
if rw := buildToolNameRewriteFromBody(body); rw != nil {
- body = applyToolNameRewriteToBody(body, rw)
+ if err := replaceBody(applyToolNameRewriteToBody(body, rw)); err != nil {
+ return err
+ }
} else {
- body = applyToolsLastCacheBreakpoint(body)
+ if err := replaceBody(applyToolsLastCacheBreakpoint(body)); err != nil {
+ return err
+ }
}
}
@@ -9052,9 +9372,13 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
}
}
if mappedModel != reqModel {
- body = s.replaceModelInBody(body, mappedModel)
+ originalReqModel := reqModel
+ if err := replaceBody(s.replaceModelInBody(body, mappedModel)); err != nil {
+ return err
+ }
reqModel = mappedModel
- logger.LegacyPrintf("service.gateway", "CountTokens model mapping applied: %s -> %s (account: %s, source=%s)", parsed.Model, mappedModel, account.Name, mappingSource)
+ parsed.Model = mappedModel
+ logger.LegacyPrintf("service.gateway", "CountTokens model mapping applied: %s -> %s (account: %s, source=%s)", originalReqModel, mappedModel, account.Name, mappingSource)
}
}
@@ -9066,11 +9390,13 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
}
// 构建上游请求
- upstreamReq, err := s.buildCountTokensRequest(ctx, c, account, body, token, tokenType, reqModel, shouldMimicClaudeCode)
+ upstreamReq, wireBody, err := s.buildCountTokensRequest(ctx, c, account, body, token, tokenType, reqModel, shouldMimicClaudeCode)
if err != nil {
s.countTokensError(c, http.StatusInternalServerError, "api_error", "Failed to build request")
return err
}
+ // 先记录首发 wire body;如果后面进入 400 retry,retry 会基于未签名的逻辑 body 重新构建。
+ acceptedWireBody := wireBody
// 获取代理URL(自定义 base URL 模式下,proxy 通过 buildCustomRelayURL 作为查询参数传递)
proxyURL := ""
@@ -9106,10 +9432,14 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
logger.LegacyPrintf("service.gateway", "Account %d: detected thinking block signature error on count_tokens, retrying with filtered thinking blocks", account.ID)
filteredBody := FilterThinkingBlocksForRetry(body)
- retryReq, buildErr := s.buildCountTokensRequest(ctx, c, account, filteredBody, token, tokenType, reqModel, shouldMimicClaudeCode)
+ retryReq, retryWireBody, buildErr := s.buildCountTokensRequest(ctx, c, account, filteredBody, token, tokenType, reqModel, shouldMimicClaudeCode)
if buildErr == nil {
retryResp, retryErr := s.httpUpstream.DoWithTLS(retryReq, proxyURL, account.ID, account.Concurrency, s.tlsFPProfileService.ResolveTLSProfile(account))
if retryErr == nil {
+ if retryResp.StatusCode < 400 {
+ // count_tokens 签名重试成功后记录最终 wire body,错误响应仍保留原 body 便于后续处理。
+ acceptedWireBody = retryWireBody
+ }
resp = retryResp
respBody, err = ReadUpstreamResponseBody(resp.Body, s.cfg, c, countTokensTooLarge)
_ = resp.Body.Close()
@@ -9123,6 +9453,13 @@ func (s *GatewayService) ForwardCountTokens(ctx context.Context, c *gin.Context,
}
}
+ if resp.StatusCode < 400 && !bytes.Equal(acceptedWireBody, body) {
+ // count_tokens 成功后再同步最终 wire body,避免 retry 从已签名 body 派生。
+ if err := replaceBody(acceptedWireBody); err != nil {
+ return err
+ }
+ }
+
// 处理错误响应
if resp.StatusCode >= 400 {
// 标记账号状态(429/529等)
@@ -9303,6 +9640,16 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough(
}
targetURL = validatedURL + "/v1/messages/count_tokens?beta=true"
}
+ body = sanitizeCountTokensRequestBody(body)
+
+ // 同 buildUpstreamRequestAnthropicAPIKeyPassthrough:能力维度 sanitize。
+ clientBeta := ""
+ if c != nil && c.Request != nil {
+ clientBeta = getHeaderRaw(c.Request.Header, "anthropic-beta")
+ }
+ if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, clientBeta); changed {
+ body = sanitized
+ }
req, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body))
if err != nil {
@@ -9339,7 +9686,7 @@ func (s *GatewayService) buildCountTokensRequestAnthropicAPIKeyPassthrough(
}
// buildCountTokensRequest 构建 count_tokens 上游请求
-func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, mimicClaudeCode bool) (*http.Request, error) {
+func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, tokenType, modelID string, mimicClaudeCode bool) (*http.Request, []byte, error) {
// 确定目标 URL
targetURL := claudeAPICountTokensURL
if account.Type == AccountTypeAPIKey {
@@ -9347,18 +9694,18 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con
if baseURL != "" {
validatedURL, err := s.validateUpstreamBaseURL(baseURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = validatedURL + "/v1/messages/count_tokens?beta=true"
}
} else if account.IsCustomBaseURLEnabled() {
customURL := account.GetCustomBaseURL()
if customURL == "" {
- return nil, fmt.Errorf("custom_base_url is enabled but not configured for account %d", account.ID)
+ return nil, nil, fmt.Errorf("custom_base_url is enabled but not configured for account %d", account.ID)
}
validatedURL, err := s.validateUpstreamBaseURL(customURL)
if err != nil {
- return nil, err
+ return nil, nil, err
}
targetURL = s.buildCustomRelayURL(validatedURL, "/v1/messages/count_tokens", account)
}
@@ -9394,13 +9741,27 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con
if ctFingerprint != nil && ctEnableFP {
body = syncBillingHeaderVersion(body, ctFingerprint.UserAgent)
}
+
+ // === 计算最终 anthropic-beta header(先于 body sanitize 与 CCH 签名)===
+ // 顺序约束同 buildUpstreamRequest。
+ ctEffectiveDropSet := mergeDropSets(s.getBetaPolicyFilterSet(ctx, c, account, modelID))
+ finalBetaHeader, finalBetaShouldSet := s.computeFinalCountTokensAnthropicBeta(
+ tokenType, mimicClaudeCode, modelID, clientHeaders, body, ctEffectiveDropSet,
+ )
+
+ // 能力维度 body sanitize:与最终 anthropic-beta header 对称
+ if sanitized, changed := sanitizeAnthropicBodyForBetaTokens(body, finalBetaHeader); changed {
+ body = sanitized
+ }
+
if ctEnableCCH {
body = signBillingHeaderCCH(body)
}
+ body = sanitizeCountTokensRequestBody(body)
req, err := http.NewRequestWithContext(ctx, "POST", targetURL, bytes.NewReader(body))
if err != nil {
- return nil, err
+ return nil, nil, err
}
// 设置认证头(保持原始大小写)
@@ -9437,41 +9798,15 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con
applyClaudeOAuthHeaderDefaults(req)
}
- // Build effective drop set for count_tokens: merge static defaults with dynamic beta policy filter rules
- ctEffectiveDropSet := mergeDropSets(s.getBetaPolicyFilterSet(ctx, c, account, modelID))
+ // OAuth + mimic Claude Code:强制注入 CLI 指纹 header
+ if tokenType == "oauth" && mimicClaudeCode {
+ applyClaudeCodeMimicHeaders(req, false)
+ }
- // OAuth 账号:处理 anthropic-beta header
- if tokenType == "oauth" {
- if mimicClaudeCode {
- applyClaudeCodeMimicHeaders(req, false)
-
- incomingBeta := getHeaderRaw(req.Header, "anthropic-beta")
- requiredBetas := append(claude.FullClaudeCodeMimicryBetas(), claude.BetaTokenCounting)
- setHeaderRaw(req.Header, "anthropic-beta", mergeAnthropicBetaDropping(requiredBetas, incomingBeta, ctEffectiveDropSet))
- } else {
- clientBetaHeader := getHeaderRaw(req.Header, "anthropic-beta")
- if clientBetaHeader == "" {
- setHeaderRaw(req.Header, "anthropic-beta", claude.CountTokensBetaHeader)
- } else {
- beta := s.getBetaHeader(modelID, clientBetaHeader)
- if !strings.Contains(beta, claude.BetaTokenCounting) {
- beta = beta + "," + claude.BetaTokenCounting
- }
- setHeaderRaw(req.Header, "anthropic-beta", stripBetaTokensWithSet(beta, ctEffectiveDropSet))
- }
- }
- } else {
- // API-key accounts: apply beta policy filter to strip controlled tokens
- if existingBeta := getHeaderRaw(req.Header, "anthropic-beta"); existingBeta != "" {
- setHeaderRaw(req.Header, "anthropic-beta", stripBetaTokensWithSet(existingBeta, ctEffectiveDropSet))
- } else if s.cfg != nil && s.cfg.Gateway.InjectBetaForAPIKey {
- // API-key:与 messages 同步的按需 beta 注入(默认关闭)
- if requestNeedsBetaFeatures(body) {
- if beta := defaultAPIKeyBetaHeader(body); beta != "" {
- setHeaderRaw(req.Header, "anthropic-beta", beta)
- }
- }
- }
+ // 写入最终 anthropic-beta header(Del 一次避免白名单透传值残留)
+ deleteHeaderAllForms(req.Header, "anthropic-beta")
+ if finalBetaShouldSet {
+ setHeaderRaw(req.Header, "anthropic-beta", finalBetaHeader)
}
// 同步 X-Claude-Code-Session-Id 头:取 body 中已处理的 metadata.user_id 的 session_id 覆盖
@@ -9490,7 +9825,26 @@ func (s *GatewayService) buildCountTokensRequest(ctx context.Context, c *gin.Con
logClaudeMimicDebug(req, body, account, tokenType, mimicClaudeCode)
}
- return req, nil
+ return req, body, nil
+}
+
+func sanitizeCountTokensRequestBody(body []byte) []byte {
+ out := body
+ for _, path := range []string{
+ "temperature",
+ "top_p",
+ "top_k",
+ "stream",
+ "stop_sequences",
+ "stop",
+ } {
+ if gjson.GetBytes(out, path).Exists() {
+ if next, ok := deleteJSONPathBytes(out, path); ok {
+ out = next
+ }
+ }
+ }
+ return out
}
// countTokensError 返回 count_tokens 错误响应
diff --git a/backend/internal/service/gateway_service_benchmark_test.go b/backend/internal/service/gateway_service_benchmark_test.go
index c9c4d3dd..42a711db 100644
--- a/backend/internal/service/gateway_service_benchmark_test.go
+++ b/backend/internal/service/gateway_service_benchmark_test.go
@@ -2,10 +2,16 @@ package service
import (
"strconv"
+ "strings"
"testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/domain"
)
-var benchmarkStringSink string
+var (
+ benchmarkStringSink string
+ benchmarkIntSink int
+)
// BenchmarkGenerateSessionHash_Metadata 关注 JSON 解析与正则匹配开销。
func BenchmarkGenerateSessionHash_Metadata(b *testing.B) {
@@ -14,7 +20,7 @@ func BenchmarkGenerateSessionHash_Metadata(b *testing.B) {
b.ReportAllocs()
for i := 0; i < b.N; i++ {
- parsed, err := ParseGatewayRequest(body, "")
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
if err != nil {
b.Fatalf("解析请求失败: %v", err)
}
@@ -22,6 +28,179 @@ func BenchmarkGenerateSessionHash_Metadata(b *testing.B) {
}
}
+func BenchmarkParseGatewayRequest_LargeAnthropicMessages(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeAnthropicMessagesBody(size.bytes, false)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic)
+ if err != nil {
+ b.Fatalf("解析 Anthropic 请求失败: %v", err)
+ }
+ benchmarkIntSink = len(parsed.MessagesRaw())
+ }
+ })
+ }
+}
+
+func BenchmarkParseGatewayRequest_LargeGeminiContents(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeGeminiContentsBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformGemini)
+ if err != nil {
+ b.Fatalf("解析 Gemini 请求失败: %v", err)
+ }
+ benchmarkIntSink = len(parsed.MessagesRaw())
+ }
+ })
+ }
+}
+
+func BenchmarkGenerateSessionHash_LargeAnthropicMessages(b *testing.B) {
+ svc := &GatewayService{}
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeAnthropicMessagesBody(size.bytes, true)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), domain.PlatformAnthropic)
+ if err != nil {
+ b.Fatalf("解析请求失败: %v", err)
+ }
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ benchmarkStringSink = svc.GenerateSessionHash(parsed)
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputMeta(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ model, stream, promptCacheKey := extractOpenAIRequestMetaFromBody(body)
+ benchmarkStringSink = model + promptCacheKey
+ if stream {
+ benchmarkIntSink++
+ }
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputDecodeMap(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ reqBody, err := getOpenAIRequestBodyMap(nil, body)
+ if err != nil {
+ b.Fatalf("解析 OpenAI 请求失败: %v", err)
+ }
+ benchmarkIntSink = len(reqBody)
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputRawPatch(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ view := newOpenAIRequestView(body)
+ view.MarkPatchSet("instructions", "You are a helpful coding assistant.")
+ view.MarkPatchSet("reasoning.effort", "none")
+ patched, err := view.ApplyPatches()
+ if err != nil {
+ b.Fatalf("应用 OpenAI raw patch 失败: %v", err)
+ }
+ benchmarkIntSink = len(patched)
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputImageBillingRaw(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesImageToolBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ cfg, err := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, "gpt-5.4")
+ if err != nil {
+ b.Fatalf("解析 OpenAI 图片计费配置失败: %v", err)
+ }
+ benchmarkStringSink = cfg.Model + cfg.SizeTier + cfg.InputSize
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputEmptyBase64Guard(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ if openAIRequestBodyMayContainEmptyBase64InputImage(body) {
+ benchmarkIntSink++
+ }
+ }
+ })
+ }
+}
+
+func BenchmarkOpenAIResponses_LargeInputFunctionCallValidation(b *testing.B) {
+ for _, size := range benchmarkBodySizes() {
+ b.Run(size.name, func(b *testing.B) {
+ body := buildLargeOpenAIResponsesToolContinuationBody(size.bytes)
+
+ b.SetBytes(int64(len(body)))
+ b.ReportAllocs()
+ b.ResetTimer()
+ for i := 0; i < b.N; i++ {
+ validation := ValidateFunctionCallOutputContextBytes(body)
+ if !validation.HasFunctionCallOutput || !validation.HasItemReferenceForAllCallIDs {
+ b.Fatalf("工具续链校验结果异常: %+v", validation)
+ }
+ benchmarkIntSink++
+ }
+ })
+ }
+}
+
// BenchmarkExtractCacheableContent_System 关注字符串拼接路径的性能。
func BenchmarkExtractCacheableContent_System(b *testing.B) {
svc := &GatewayService{}
@@ -33,18 +212,130 @@ func BenchmarkExtractCacheableContent_System(b *testing.B) {
}
}
-func buildSystemCacheableRequest(parts int) *ParsedRequest {
- systemParts := make([]any, 0, parts)
- for i := 0; i < parts; i++ {
- systemParts = append(systemParts, map[string]any{
- "text": "system_part_" + strconv.Itoa(i),
- "cache_control": map[string]any{
- "type": "ephemeral",
- },
- })
- }
- return &ParsedRequest{
- System: systemParts,
- HasSystem: true,
+func benchmarkBodySizes() []struct {
+ name string
+ bytes int
+} {
+ return []struct {
+ name string
+ bytes int
+ }{
+ {name: "4MB", bytes: 4 << 20},
+ {name: "8MB", bytes: 8 << 20},
+ {name: "16MB", bytes: 16 << 20},
+ {name: "32MB", bytes: 32 << 20},
}
}
+
+func buildSystemCacheableRequest(parts int) *ParsedRequest {
+ var builder strings.Builder
+ _, _ = builder.WriteString(`{"system":[`)
+ for i := 0; i < parts; i++ {
+ if i > 0 {
+ _ = builder.WriteByte(',')
+ }
+ _, _ = builder.WriteString(`{"text":"system_part_`)
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`","cache_control":{"type":"ephemeral"}}`)
+ }
+ _, _ = builder.WriteString(`]}`)
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(builder.String())), "")
+ if err != nil {
+ panic(err)
+ }
+ return parsed
+}
+
+func buildLargeAnthropicMessagesBody(targetBytes int, includeCacheControl bool) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ _, _ = builder.WriteString(`{"model":"claude-sonnet-4-5","stream":true,"system":[{"type":"text","text":"system seed"}],"messages":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ _ = builder.WriteByte(',')
+ }
+ _, _ = builder.WriteString(`{"role":"user","content":[{"type":"text","text":"`)
+ _, _ = builder.WriteString(strings.Repeat("anthropic payload ", 64))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _ = builder.WriteByte('"')
+ if includeCacheControl && i%32 == 0 {
+ _, _ = builder.WriteString(`,"cache_control":{"type":"ephemeral"}`)
+ }
+ _, _ = builder.WriteString(`}]}`)
+ }
+ _, _ = builder.WriteString(`]}`)
+ return []byte(builder.String())
+}
+
+func buildLargeGeminiContentsBody(targetBytes int) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ _, _ = builder.WriteString(`{"model":"gemini-2.5-pro","systemInstruction":{"parts":[{"text":"system seed"}]},"contents":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ _ = builder.WriteByte(',')
+ }
+ _, _ = builder.WriteString(`{"role":"user","parts":[{"text":"`)
+ _, _ = builder.WriteString(strings.Repeat("gemini payload ", 64))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`"}]}`)
+ }
+ _, _ = builder.WriteString(`]}`)
+ return []byte(builder.String())
+}
+
+func buildLargeOpenAIResponsesBody(targetBytes int) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ _, _ = builder.WriteString(`{"model":"gpt-5.4","stream":true,"prompt_cache_key":"session-benchmark","input":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ _ = builder.WriteByte(',')
+ }
+ _, _ = builder.WriteString(`{"type":"message","role":"user","content":[{"type":"input_text","text":"`)
+ _, _ = builder.WriteString(strings.Repeat("openai responses payload ", 48))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`"}]}`)
+ }
+ _, _ = builder.WriteString(`],"tools":[{"type":"function","name":"lookup","parameters":{"type":"object","properties":{"query":{"type":"string"}}}}]}`)
+ return []byte(builder.String())
+}
+
+func buildLargeOpenAIResponsesToolContinuationBody(targetBytes int) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ _, _ = builder.WriteString(`{"model":"gpt-5.4","stream":true,"previous_response_id":"resp_benchmark","input":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ _ = builder.WriteByte(',')
+ }
+ callID := "call_" + strconv.Itoa(i)
+ _, _ = builder.WriteString(`{"type":"item_reference","id":"`)
+ _, _ = builder.WriteString(callID)
+ _, _ = builder.WriteString(`"},{"type":"function_call_output","call_id":"`)
+ _, _ = builder.WriteString(callID)
+ _, _ = builder.WriteString(`","output":"`)
+ _, _ = builder.WriteString(strings.Repeat("tool output payload ", 48))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`"}`)
+ }
+ _, _ = builder.WriteString(`]}`)
+ return []byte(builder.String())
+}
+
+func buildLargeOpenAIResponsesImageToolBody(targetBytes int) []byte {
+ var builder strings.Builder
+ builder.Grow(targetBytes + 1024)
+ _, _ = builder.WriteString(`{"model":"gpt-5.4","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[`)
+ for i := 0; builder.Len() < targetBytes; i++ {
+ if i > 0 {
+ _ = builder.WriteByte(',')
+ }
+ _, _ = builder.WriteString(`{"type":"message","role":"user","content":[{"type":"input_text","text":"`)
+ _, _ = builder.WriteString(strings.Repeat("openai image billing payload ", 48))
+ _, _ = builder.WriteString(strconv.Itoa(i))
+ _, _ = builder.WriteString(`"}]}`)
+ }
+ _, _ = builder.WriteString(`]}`)
+ return []byte(builder.String())
+}
diff --git a/backend/internal/service/gateway_websearch_emulation.go b/backend/internal/service/gateway_websearch_emulation.go
index a42b5585..2f9c8e0c 100644
--- a/backend/internal/service/gateway_websearch_emulation.go
+++ b/backend/internal/service/gateway_websearch_emulation.go
@@ -150,7 +150,7 @@ func (s *GatewayService) handleWebSearchEmulation(
parsed.OnUpstreamAccepted()
}
- query := extractSearchQueryFromBody(parsed.Body)
+ query := extractSearchQueryFromBody(parsed.Body.Bytes())
if query == "" {
return nil, fmt.Errorf("web search emulation: no query found in messages")
}
diff --git a/backend/internal/service/gemini_chat_completions_compat_service.go b/backend/internal/service/gemini_chat_completions_compat_service.go
index dcc3213b..ffea1595 100644
--- a/backend/internal/service/gemini_chat_completions_compat_service.go
+++ b/backend/internal/service/gemini_chat_completions_compat_service.go
@@ -151,7 +151,7 @@ func (s *GeminiMessagesCompatService) forwardClaudeBodyAsChatCompletions(
}
if resp.StatusCode >= 400 && s.shouldRetryGeminiUpstreamError(account, resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
if resp.StatusCode == http.StatusForbidden && isGeminiInsufficientScope(resp.Header, respBody) {
resp = &http.Response{
@@ -207,7 +207,7 @@ func (s *GeminiMessagesCompatService) forwardClaudeBodyAsChatCompletions(
reasoningEffort := extractCCReasoningEffortFromBody(originalChatBody)
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
s.handleGeminiUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
evBody := unwrapIfNeeded(account.Type == AccountTypeOAuth, respBody)
diff --git a/backend/internal/service/gemini_messages_compat_service.go b/backend/internal/service/gemini_messages_compat_service.go
index 516556ca..86073d9c 100644
--- a/backend/internal/service/gemini_messages_compat_service.go
+++ b/backend/internal/service/gemini_messages_compat_service.go
@@ -56,6 +56,18 @@ type GeminiMessagesCompatService struct {
responseHeaderFilter *responseheaders.CompiledHeaderFilter
}
+func (s *GeminiMessagesCompatService) readUpstreamErrorBody(resp *http.Response) []byte {
+ if resp == nil || resp.Body == nil {
+ return nil
+ }
+ limit := gatewayUpstreamErrorBodyReadLimit
+ if s != nil && s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody && s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes > int(limit) {
+ limit = int64(s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes)
+ }
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, limit))
+ return body
+}
+
func NewGeminiMessagesCompatService(
accountRepo AccountRepository,
groupRepo GroupRepository,
@@ -789,7 +801,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
// Special-case: signature/thought_signature validation errors are not transient, but may be fixed by
// downgrading Claude thinking/tool history to plain text (conservative two-stage retry).
if resp.StatusCode == http.StatusBadRequest && signatureRetryStage < 2 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
if isGeminiSignatureRelatedError(respBody) {
@@ -860,7 +872,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
}
if resp.StatusCode >= 400 && s.shouldRetryGeminiUpstreamError(account, resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
// Don't treat insufficient-scope as transient.
if resp.StatusCode == 403 && isGeminiInsufficientScope(resp.Header, respBody) {
@@ -919,7 +931,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
// 统一错误策略:自定义错误码 + 临时不可调度
if s.rateLimitService != nil {
switch s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, respBody) {
@@ -1329,7 +1341,7 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
}
if resp.StatusCode >= 400 && s.shouldRetryGeminiUpstreamError(account, resp.StatusCode) {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
// Don't treat insufficient-scope as transient.
if resp.StatusCode == 403 && isGeminiInsufficientScope(resp.Header, respBody) {
@@ -1410,7 +1422,7 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
isOAuth := account.Type == AccountTypeOAuth
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
// Best-effort fallback for OAuth tokens missing AI Studio scopes when calling countTokens.
// This avoids Gemini SDKs failing hard during preflight token counting.
// Checked before error policy so it always works regardless of custom error codes.
@@ -1619,7 +1631,7 @@ func (s *GeminiMessagesCompatService) checkErrorPolicyInLoop(
if resp.StatusCode < 400 || s.rateLimitService == nil {
return false, resp
}
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
rebuilt = &http.Response{
StatusCode: resp.StatusCode,
@@ -2031,6 +2043,22 @@ func (s *GeminiMessagesCompatService) handleStreamingResponse(c *gin.Context, re
parts := extractGeminiParts(geminiResp)
for _, part := range parts {
if text, ok := part["text"].(string); ok && text != "" {
+ // Close an open tool_use block before starting text, mirroring
+ // the functionCall branch (which closes open text blocks) and
+ // the chat-completions sibling's closeOpenTool(). Otherwise a
+ // tool→text sequence keeps the tool_use block open while the
+ // text block starts, emitting overlapping Anthropic content
+ // blocks that violate the SSE contract.
+ if openToolIndex >= 0 {
+ writeSSE(c.Writer, "content_block_stop", map[string]any{
+ "type": "content_block_stop",
+ "index": openToolIndex,
+ })
+ openToolIndex = -1
+ openToolName = ""
+ seenToolJSON = ""
+ }
+
delta, newSeen := computeGeminiTextDelta(seenText, text)
seenText = newSeen
if delta == "" {
diff --git a/backend/internal/service/gemini_messages_compat_service_test.go b/backend/internal/service/gemini_messages_compat_service_test.go
index d0560344..79db633a 100644
--- a/backend/internal/service/gemini_messages_compat_service_test.go
+++ b/backend/internal/service/gemini_messages_compat_service_test.go
@@ -832,3 +832,108 @@ func TestParseGeminiRateLimitResetTime(t *testing.T) {
})
}
}
+
+// TestGeminiMessagesHandleStreamingResponse_ClosesToolBlockBeforeText guards the
+// tool→text ordering in the Gemini→Anthropic (messages) streaming bridge. When
+// Gemini emits a functionCall part followed by a text part, the tool_use content
+// block must be closed before the text block opens; otherwise the Anthropic SSE
+// stream contains overlapping content blocks. The chat-completions sibling
+// already enforces this via closeOpenTool().
+func TestGeminiMessagesHandleStreamingResponse_ClosesToolBlockBeforeText(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ upstreamBody := `data: {"candidates":[{"content":{"parts":[{"functionCall":{"name":"get_weather","args":{"city":"SF"}}}]}}]}` + "\n\n" +
+ `data: {"candidates":[{"content":{"parts":[{"text":"All done."}]},"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":5,"candidatesTokenCount":3}}` + "\n\n" +
+ "data: [DONE]\n\n"
+
+ resp := &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"text/event-stream"}},
+ Body: io.NopCloser(strings.NewReader(upstreamBody)),
+ }
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+
+ svc := &GeminiMessagesCompatService{}
+ result, err := svc.handleStreamingResponse(c, resp, time.Now(), "claude-3-5-sonnet")
+ require.NoError(t, err)
+ require.NotNil(t, result)
+
+ events := parseAnthropicContentBlockEvents(t, rec.Body.String())
+
+ // Anthropic allows at most one content block open at a time: every
+ // content_block_start must be matched by a content_block_stop before the
+ // next start. Replay the lifecycle and assert there is no overlap.
+ open := -1
+ blockTypes := map[int]string{}
+ textStarted := false
+ toolClosed := false
+ toolClosedBeforeText := false
+ for _, ev := range events {
+ switch ev.event {
+ case "content_block_start":
+ require.Equalf(t, -1, open,
+ "content block %d opened while block %d was still open (overlapping blocks)", ev.index, open)
+ open = ev.index
+ blockTypes[ev.index] = ev.blockType
+ if ev.blockType == "text" {
+ textStarted = true
+ if toolClosed {
+ toolClosedBeforeText = true
+ }
+ }
+ case "content_block_stop":
+ require.Equalf(t, open, ev.index,
+ "content_block_stop index %d does not match the open block %d", ev.index, open)
+ if blockTypes[ev.index] == "tool_use" {
+ toolClosed = true
+ }
+ open = -1
+ }
+ }
+
+ require.True(t, textStarted, "expected a text content block to be emitted after the tool call")
+ require.True(t, toolClosedBeforeText, "tool_use block must be closed before the text block starts")
+ require.Equal(t, -1, open, "stream ended with a content block still open")
+}
+
+type anthropicContentBlockEvent struct {
+ event string
+ index int
+ blockType string
+}
+
+// parseAnthropicContentBlockEvents extracts content_block_start/stop events (with
+// their index and, for starts, the content block type) from an Anthropic SSE body.
+func parseAnthropicContentBlockEvents(t *testing.T, raw string) []anthropicContentBlockEvent {
+ t.Helper()
+ var events []anthropicContentBlockEvent
+ for _, chunk := range strings.Split(raw, "\n\n") {
+ var eventName, dataLine string
+ for _, line := range strings.Split(chunk, "\n") {
+ switch {
+ case strings.HasPrefix(line, "event:"):
+ eventName = strings.TrimSpace(strings.TrimPrefix(line, "event:"))
+ case strings.HasPrefix(line, "data:"):
+ dataLine = strings.TrimSpace(strings.TrimPrefix(line, "data:"))
+ }
+ }
+ if eventName != "content_block_start" && eventName != "content_block_stop" {
+ continue
+ }
+ var payload struct {
+ Index int `json:"index"`
+ ContentBlock struct {
+ Type string `json:"type"`
+ } `json:"content_block"`
+ }
+ require.NoError(t, json.Unmarshal([]byte(dataLine), &payload))
+ events = append(events, anthropicContentBlockEvent{
+ event: eventName,
+ index: payload.Index,
+ blockType: payload.ContentBlock.Type,
+ })
+ }
+ return events
+}
diff --git a/backend/internal/service/gemini_multiplatform_test.go b/backend/internal/service/gemini_multiplatform_test.go
index 5e09b95a..8f879b02 100644
--- a/backend/internal/service/gemini_multiplatform_test.go
+++ b/backend/internal/service/gemini_multiplatform_test.go
@@ -147,7 +147,7 @@ func (m *mockAccountRepoForGemini) ListSchedulableUngroupedByPlatforms(ctx conte
func (m *mockAccountRepoForGemini) SetRateLimited(ctx context.Context, id int64, resetAt time.Time) error {
return nil
}
-func (m *mockAccountRepoForGemini) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time) error {
+func (m *mockAccountRepoForGemini) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error {
return nil
}
func (m *mockAccountRepoForGemini) SetOverloaded(ctx context.Context, id int64, until time.Time) error {
diff --git a/backend/internal/service/generate_session_hash_test.go b/backend/internal/service/generate_session_hash_test.go
index 39679c3d..5ed3f0ae 100644
--- a/backend/internal/service/generate_session_hash_test.go
+++ b/backend/internal/service/generate_session_hash_test.go
@@ -3,12 +3,67 @@
package service
import (
+ "encoding/json"
"testing"
+ "github.com/Wei-Shaw/sub2api/internal/domain"
"github.com/stretchr/testify/require"
)
-// ============ 基础优先级测试 ============
+func mustParseSessionHashRequest(t *testing.T, body string, ctx *SessionContext) *ParsedRequest {
+ t.Helper()
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(body)), domain.PlatformAnthropic)
+ require.NoError(t, err)
+ parsed.SessionContext = ctx
+ return parsed
+}
+
+func mustParseGeminiSessionHashRequest(t *testing.T, body string, ctx *SessionContext) *ParsedRequest {
+ t.Helper()
+ parsed, err := ParseGatewayRequest(NewRequestBodyRef([]byte(body)), domain.PlatformGemini)
+ require.NoError(t, err)
+ parsed.SessionContext = ctx
+ return parsed
+}
+
+func anthropicSessionBody(system any, messages []any, metadataUserID string) string {
+ body := map[string]any{}
+ if system != nil {
+ body["system"] = system
+ }
+ if messages != nil {
+ body["messages"] = messages
+ }
+ if metadataUserID != "" {
+ body["metadata"] = map[string]any{"user_id": metadataUserID}
+ }
+ data, _ := json.Marshal(body)
+ return string(data)
+}
+
+func geminiSessionBody(systemParts []any, contents []any) string {
+ body := map[string]any{}
+ if systemParts != nil {
+ body["systemInstruction"] = map[string]any{"parts": systemParts}
+ }
+ if contents != nil {
+ body["contents"] = contents
+ }
+ data, _ := json.Marshal(body)
+ return string(data)
+}
+
+func msg(role string, content any) map[string]any {
+ return map[string]any{"role": role, "content": content}
+}
+
+func geminiMsg(role string, texts ...string) map[string]any {
+ parts := make([]any, 0, len(texts))
+ for _, text := range texts {
+ parts = append(parts, map[string]any{"text": text})
+ }
+ return map[string]any{"role": role, "parts": parts}
+}
func TestGenerateSessionHash_NilParsedRequest(t *testing.T) {
svc := &GatewayService{}
@@ -22,37 +77,17 @@ func TestGenerateSessionHash_EmptyRequest(t *testing.T) {
func TestGenerateSessionHash_MetadataHasHighestPriority(t *testing.T) {
svc := &GatewayService{}
-
- parsed := &ParsedRequest{
- MetadataUserID: "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000",
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ metadata := "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000"
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, metadata), nil)
hash := svc.GenerateSessionHash(parsed)
require.Equal(t, "123e4567-e89b-12d3-a456-426614174000", hash, "metadata session_id should have highest priority")
}
-// ============ System + Messages 基础测试 ============
-
func TestGenerateSessionHash_SystemPlusMessages(t *testing.T) {
svc := &GatewayService{}
-
- withSystem := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
- withoutSystem := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ withSystem := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, ""), nil)
+ withoutSystem := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "hello")}, ""), nil)
h1 := svc.GenerateSessionHash(withSystem)
h2 := svc.GenerateSessionHash(withoutSystem)
@@ -63,32 +98,16 @@ func TestGenerateSessionHash_SystemPlusMessages(t *testing.T) {
func TestGenerateSessionHash_SystemOnlyProducesHash(t *testing.T) {
svc := &GatewayService{}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", nil, ""), nil)
- parsed := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- }
hash := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, hash, "system prompt alone should produce a hash as part of full digest")
}
func TestGenerateSessionHash_DifferentSystemsSameMessages(t *testing.T) {
svc := &GatewayService{}
-
- parsed1 := &ParsedRequest{
- System: "You are assistant A.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
- parsed2 := &ParsedRequest{
- System: "You are assistant B.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ parsed1 := mustParseSessionHashRequest(t, anthropicSessionBody("You are assistant A.", []any{msg("user", "hello")}, ""), nil)
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody("You are assistant B.", []any{msg("user", "hello")}, ""), nil)
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
@@ -97,16 +116,8 @@ func TestGenerateSessionHash_DifferentSystemsSameMessages(t *testing.T) {
func TestGenerateSessionHash_SameSystemSameMessages(t *testing.T) {
svc := &GatewayService{}
-
mk := func() *ParsedRequest {
- return &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "hi"},
- },
- }
+ return mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello"), msg("assistant", "hi")}, ""), nil)
}
h1 := svc.GenerateSessionHash(mk())
@@ -116,53 +127,19 @@ func TestGenerateSessionHash_SameSystemSameMessages(t *testing.T) {
func TestGenerateSessionHash_DifferentMessagesProduceDifferentHash(t *testing.T) {
svc := &GatewayService{}
-
- parsed1 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "help me with Go"},
- },
- }
- parsed2 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "help me with Python"},
- },
- }
+ parsed1 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "help me with Go")}, ""), nil)
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "help me with Python")}, ""), nil)
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h1, h2, "same system but different messages should produce different hashes")
}
-// ============ SessionContext 核心测试 ============
-
func TestGenerateSessionHash_DifferentSessionContextProducesDifferentHash(t *testing.T) {
svc := &GatewayService{}
-
- // 相同消息 + 不同 SessionContext → 不同 hash(解决碰撞问题的核心场景)
- parsed1 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "192.168.1.1",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 100,
- },
- }
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "10.0.0.1",
- UserAgent: "curl/7.0",
- APIKeyID: 200,
- },
- }
+ body := anthropicSessionBody(nil, []any{msg("user", "hello")}, "")
+ parsed1 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "192.168.1.1", UserAgent: "Mozilla/5.0", APIKeyID: 100})
+ parsed2 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "10.0.0.1", UserAgent: "curl/7.0", APIKeyID: 200})
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
@@ -173,19 +150,9 @@ func TestGenerateSessionHash_DifferentSessionContextProducesDifferentHash(t *tes
func TestGenerateSessionHash_SameSessionContextProducesSameHash(t *testing.T) {
svc := &GatewayService{}
-
- mk := func() *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "192.168.1.1",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 100,
- },
- }
- }
+ ctx := &SessionContext{ClientIP: "192.168.1.1", UserAgent: "Mozilla/5.0", APIKeyID: 100}
+ body := anthropicSessionBody(nil, []any{msg("user", "hello")}, "")
+ mk := func() *ParsedRequest { return mustParseSessionHashRequest(t, body, ctx) }
h1 := svc.GenerateSessionHash(mk())
h2 := svc.GenerateSessionHash(mk())
@@ -194,35 +161,17 @@ func TestGenerateSessionHash_SameSessionContextProducesSameHash(t *testing.T) {
func TestGenerateSessionHash_MetadataOverridesSessionContext(t *testing.T) {
svc := &GatewayService{}
-
- parsed := &ParsedRequest{
- MetadataUserID: "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000",
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "192.168.1.1",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 100,
- },
- }
+ metadata := "user_a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2_account__session_123e4567-e89b-12d3-a456-426614174000"
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "hello")}, metadata), &SessionContext{ClientIP: "192.168.1.1", UserAgent: "Mozilla/5.0", APIKeyID: 100})
hash := svc.GenerateSessionHash(parsed)
- require.Equal(t, "123e4567-e89b-12d3-a456-426614174000", hash,
- "metadata session_id should take priority over SessionContext")
+ require.Equal(t, "123e4567-e89b-12d3-a456-426614174000", hash, "metadata session_id should take priority over SessionContext")
}
func TestGenerateSessionHash_MetadataJSON_HasHighestPriority(t *testing.T) {
svc := &GatewayService{}
-
- parsed := &ParsedRequest{
- MetadataUserID: `{"device_id":"a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2","account_uuid":"","session_id":"c72554f2-1234-5678-abcd-123456789abc"}`,
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ metadata := `{"device_id":"a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2c3d4e5f6a1b2","account_uuid":"","session_id":"c72554f2-1234-5678-abcd-123456789abc"}`
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, metadata), nil)
hash := svc.GenerateSessionHash(parsed)
require.Equal(t, "c72554f2-1234-5678-abcd-123456789abc", hash, "JSON format metadata session_id should have highest priority")
@@ -230,69 +179,25 @@ func TestGenerateSessionHash_MetadataJSON_HasHighestPriority(t *testing.T) {
func TestGenerateSessionHash_NilSessionContextBackwardCompatible(t *testing.T) {
svc := &GatewayService{}
-
- withCtx := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: nil,
- }
- withoutCtx := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- }
+ body := anthropicSessionBody(nil, []any{msg("user", "hello")}, "")
+ withCtx := mustParseSessionHashRequest(t, body, nil)
+ withoutCtx := mustParseSessionHashRequest(t, body, nil)
h1 := svc.GenerateSessionHash(withCtx)
h2 := svc.GenerateSessionHash(withoutCtx)
require.Equal(t, h1, h2, "nil SessionContext should produce same hash as no SessionContext")
}
-// ============ 多轮连续会话测试 ============
-
func TestGenerateSessionHash_ContinuousConversation_HashChangesWithMessages(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 模拟连续会话:每增加一轮对话,hash 应该不同(内容累积变化)
- round1 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: ctx,
- }
-
- round2 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "Hi there!"},
- map[string]any{"role": "user", "content": "How are you?"},
- },
- SessionContext: ctx,
- }
-
- round3 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "Hi there!"},
- map[string]any{"role": "user", "content": "How are you?"},
- map[string]any{"role": "assistant", "content": "I'm doing well!"},
- map[string]any{"role": "user", "content": "Tell me a joke"},
- },
- SessionContext: ctx,
- }
+ round1 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, ""), ctx)
+ round2 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello"), msg("assistant", "Hi there!"), msg("user", "How are you?")}, ""), ctx)
+ round3 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello"), msg("assistant", "Hi there!"), msg("user", "How are you?"), msg("assistant", "I'm doing well!"), msg("user", "Tell me a joke")}, ""), ctx)
h1 := svc.GenerateSessionHash(round1)
h2 := svc.GenerateSessionHash(round2)
h3 := svc.GenerateSessionHash(round3)
-
require.NotEmpty(t, h1)
require.NotEmpty(t, h2)
require.NotEmpty(t, h3)
@@ -303,62 +208,20 @@ func TestGenerateSessionHash_ContinuousConversation_HashChangesWithMessages(t *t
func TestGenerateSessionHash_ContinuousConversation_SameRoundSameHash(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 同一轮对话重复请求(如重试)应产生相同 hash
- mk := func() *ParsedRequest {
- return &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- map[string]any{"role": "assistant", "content": "Hi there!"},
- map[string]any{"role": "user", "content": "How are you?"},
- },
- SessionContext: ctx,
- }
- }
+ body := anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello"), msg("assistant", "Hi there!"), msg("user", "How are you?")}, "")
+ mk := func() *ParsedRequest { return mustParseSessionHashRequest(t, body, ctx) }
h1 := svc.GenerateSessionHash(mk())
h2 := svc.GenerateSessionHash(mk())
require.Equal(t, h1, h2, "same conversation state should produce identical hash on retry")
}
-// ============ 消息回退测试 ============
-
func TestGenerateSessionHash_MessageRollback(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 模拟消息回退:用户删掉最后一轮再重发
- original := &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "msg1"},
- map[string]any{"role": "assistant", "content": "reply1"},
- map[string]any{"role": "user", "content": "msg2"},
- map[string]any{"role": "assistant", "content": "reply2"},
- map[string]any{"role": "user", "content": "msg3"},
- },
- SessionContext: ctx,
- }
-
- // 回退到 msg2 后,用新的 msg3 替代
- rollback := &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "msg1"},
- map[string]any{"role": "assistant", "content": "reply1"},
- map[string]any{"role": "user", "content": "msg2"},
- map[string]any{"role": "assistant", "content": "reply2"},
- map[string]any{"role": "user", "content": "different msg3"},
- },
- SessionContext: ctx,
- }
+ original := mustParseSessionHashRequest(t, anthropicSessionBody("System prompt", []any{msg("user", "msg1"), msg("assistant", "reply1"), msg("user", "msg2"), msg("assistant", "reply2"), msg("user", "msg3")}, ""), ctx)
+ rollback := mustParseSessionHashRequest(t, anthropicSessionBody("System prompt", []any{msg("user", "msg1"), msg("assistant", "reply1"), msg("user", "msg2"), msg("assistant", "reply2"), msg("user", "different msg3")}, ""), ctx)
hOrig := svc.GenerateSessionHash(original)
hRollback := svc.GenerateSessionHash(rollback)
@@ -367,58 +230,19 @@ func TestGenerateSessionHash_MessageRollback(t *testing.T) {
func TestGenerateSessionHash_MessageRollbackSameContent(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 回退后重新发送相同内容 → 相同 hash(合理的粘性恢复)
- mk := func() *ParsedRequest {
- return &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "msg1"},
- map[string]any{"role": "assistant", "content": "reply1"},
- map[string]any{"role": "user", "content": "msg2"},
- },
- SessionContext: ctx,
- }
- }
+ body := anthropicSessionBody("System prompt", []any{msg("user", "msg1"), msg("assistant", "reply1"), msg("user", "msg2")}, "")
+ mk := func() *ParsedRequest { return mustParseSessionHashRequest(t, body, ctx) }
h1 := svc.GenerateSessionHash(mk())
h2 := svc.GenerateSessionHash(mk())
require.Equal(t, h1, h2, "rollback and resend same content should produce same hash")
}
-// ============ 相同 System、不同用户消息 ============
-
func TestGenerateSessionHash_SameSystemDifferentUsers(t *testing.T) {
svc := &GatewayService{}
-
- // 两个不同用户使用相同 system prompt 但发送不同消息
- user1 := &ParsedRequest{
- System: "You are a code reviewer.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "Review this Go code"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "vscode",
- APIKeyID: 1,
- },
- }
- user2 := &ParsedRequest{
- System: "You are a code reviewer.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "Review this Python code"},
- },
- SessionContext: &SessionContext{
- ClientIP: "2.2.2.2",
- UserAgent: "vscode",
- APIKeyID: 2,
- },
- }
+ user1 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a code reviewer.", []any{msg("user", "Review this Go code")}, ""), &SessionContext{ClientIP: "1.1.1.1", UserAgent: "vscode", APIKeyID: 1})
+ user2 := mustParseSessionHashRequest(t, anthropicSessionBody("You are a code reviewer.", []any{msg("user", "Review this Python code")}, ""), &SessionContext{ClientIP: "2.2.2.2", UserAgent: "vscode", APIKeyID: 2})
h1 := svc.GenerateSessionHash(user1)
h2 := svc.GenerateSessionHash(user2)
@@ -427,55 +251,20 @@ func TestGenerateSessionHash_SameSystemDifferentUsers(t *testing.T) {
func TestGenerateSessionHash_SameSystemSameMessageDifferentContext(t *testing.T) {
svc := &GatewayService{}
-
- // 这是修复的核心场景:两个不同用户发送完全相同的 system + messages(如 "hello")
- // 有了 SessionContext 后应该产生不同 hash
- user1 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 10,
- },
- }
- user2 := &ParsedRequest{
- System: "You are a helpful assistant.",
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "2.2.2.2",
- UserAgent: "Mozilla/5.0",
- APIKeyID: 20,
- },
- }
+ body := anthropicSessionBody("You are a helpful assistant.", []any{msg("user", "hello")}, "")
+ user1 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: "Mozilla/5.0", APIKeyID: 10})
+ user2 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "2.2.2.2", UserAgent: "Mozilla/5.0", APIKeyID: 20})
h1 := svc.GenerateSessionHash(user1)
h2 := svc.GenerateSessionHash(user2)
require.NotEqual(t, h1, h2, "CRITICAL: same system+messages but different users should get different hashes")
}
-// ============ SessionContext 各字段独立影响测试 ============
-
func TestGenerateSessionHash_SessionContext_IPDifference(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(ip string) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: ip,
- UserAgent: "same-ua",
- APIKeyID: 1,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: ip, UserAgent: "same-ua", APIKeyID: 1})
}
h1 := svc.GenerateSessionHash(base("1.1.1.1"))
@@ -485,18 +274,9 @@ func TestGenerateSessionHash_SessionContext_IPDifference(t *testing.T) {
func TestGenerateSessionHash_SessionContext_UADifference(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(ua string) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: ua,
- APIKeyID: 1,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: ua, APIKeyID: 1})
}
h1 := svc.GenerateSessionHash(base("Mozilla/5.0"))
@@ -506,18 +286,9 @@ func TestGenerateSessionHash_SessionContext_UADifference(t *testing.T) {
func TestGenerateSessionHash_SessionContext_UAVersionNoiseIgnored(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(ua string) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: ua,
- APIKeyID: 1,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: ua, APIKeyID: 1})
}
h1 := svc.GenerateSessionHash(base("Mozilla/5.0 codex_cli_rs/0.1.0"))
@@ -527,18 +298,9 @@ func TestGenerateSessionHash_SessionContext_UAVersionNoiseIgnored(t *testing.T)
func TestGenerateSessionHash_SessionContext_FreeformUAVersionNoiseIgnored(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(ua string) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: ua,
- APIKeyID: 1,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: ua, APIKeyID: 1})
}
h1 := svc.GenerateSessionHash(base("Codex CLI 0.1.0"))
@@ -548,18 +310,9 @@ func TestGenerateSessionHash_SessionContext_FreeformUAVersionNoiseIgnored(t *tes
func TestGenerateSessionHash_SessionContext_APIKeyIDDifference(t *testing.T) {
svc := &GatewayService{}
-
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
base := func(keyID int64) *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "same-ua",
- APIKeyID: keyID,
- },
- }
+ return mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: "same-ua", APIKeyID: keyID})
}
h1 := svc.GenerateSessionHash(base(1))
@@ -567,24 +320,12 @@ func TestGenerateSessionHash_SessionContext_APIKeyIDDifference(t *testing.T) {
require.NotEqual(t, h1, h2, "different APIKeyID should produce different hash")
}
-// ============ 多用户并发相同消息场景 ============
-
func TestGenerateSessionHash_MultipleUsersSameFirstMessage(t *testing.T) {
svc := &GatewayService{}
-
- // 模拟 5 个不同用户同时发送 "hello" → 应该产生 5 个不同的 hash
hashes := make(map[string]bool)
+ body := anthropicSessionBody(nil, []any{msg("user", "hello")}, "")
for i := 0; i < 5; i++ {
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "192.168.1." + string(rune('1'+i)),
- UserAgent: "client-" + string(rune('A'+i)),
- APIKeyID: int64(i + 1),
- },
- }
+ parsed := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "192.168.1." + string(rune('1'+i)), UserAgent: "client-" + string(rune('A'+i)), APIKeyID: int64(i + 1)})
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h)
require.False(t, hashes[h], "hash collision detected for user %d", i)
@@ -593,134 +334,56 @@ func TestGenerateSessionHash_MultipleUsersSameFirstMessage(t *testing.T) {
require.Len(t, hashes, 5, "5 different users should produce 5 unique hashes")
}
-// ============ 连续会话粘性:多轮对话同一用户 ============
-
func TestGenerateSessionHash_SameUserGrowingConversation(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "browser", APIKeyID: 42}
-
- // 模拟同一用户的连续会话,每轮 hash 不同但同用户重试保持一致
- messages := []map[string]any{
- {"role": "user", "content": "msg1"},
- {"role": "assistant", "content": "reply1"},
- {"role": "user", "content": "msg2"},
- {"role": "assistant", "content": "reply2"},
- {"role": "user", "content": "msg3"},
- {"role": "assistant", "content": "reply3"},
- {"role": "user", "content": "msg4"},
+ messages := []any{
+ msg("user", "msg1"), msg("assistant", "reply1"), msg("user", "msg2"), msg("assistant", "reply2"),
+ msg("user", "msg3"), msg("assistant", "reply3"), msg("user", "msg4"),
}
prevHash := ""
for round := 1; round <= len(messages); round += 2 {
- // 构建前 round 条消息
- msgs := make([]any, round)
- for j := 0; j < round; j++ {
- msgs[j] = messages[j]
- }
- parsed := &ParsedRequest{
- System: "System",
- HasSystem: true,
- Messages: msgs,
- SessionContext: ctx,
- }
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("System", messages[:round], ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "round %d hash should not be empty", round)
-
if prevHash != "" {
require.NotEqual(t, prevHash, h, "round %d hash should differ from previous round", round)
}
prevHash = h
-
- // 同一轮重试应该相同
h2 := svc.GenerateSessionHash(parsed)
require.Equal(t, h, h2, "retry of round %d should produce same hash", round)
}
}
-// ============ 多轮消息内容结构化测试 ============
-
func TestGenerateSessionHash_MultipleUserMessages(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 5 条用户消息(无 assistant 回复)
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "first"},
- map[string]any{"role": "user", "content": "second"},
- map[string]any{"role": "user", "content": "third"},
- map[string]any{"role": "user", "content": "fourth"},
- map[string]any{"role": "user", "content": "fifth"},
- },
- SessionContext: ctx,
- }
-
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "first"), msg("user", "second"), msg("user", "third"), msg("user", "fourth"), msg("user", "fifth")}, ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h)
- // 修改中间一条消息应该改变 hash
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "first"},
- map[string]any{"role": "user", "content": "CHANGED"},
- map[string]any{"role": "user", "content": "third"},
- map[string]any{"role": "user", "content": "fourth"},
- map[string]any{"role": "user", "content": "fifth"},
- },
- SessionContext: ctx,
- }
-
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "first"), msg("user", "CHANGED"), msg("user", "third"), msg("user", "fourth"), msg("user", "fifth")}, ""), ctx)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h, h2, "changing any message should change the hash")
}
func TestGenerateSessionHash_MessageOrderMatters(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- parsed1 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "alpha"},
- map[string]any{"role": "user", "content": "beta"},
- },
- SessionContext: ctx,
- }
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "beta"},
- map[string]any{"role": "user", "content": "alpha"},
- },
- SessionContext: ctx,
- }
+ parsed1 := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "alpha"), msg("user", "beta")}, ""), ctx)
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", "beta"), msg("user", "alpha")}, ""), ctx)
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h1, h2, "message order should affect the hash")
}
-// ============ 复杂内容格式测试 ============
-
func TestGenerateSessionHash_StructuredContent(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 结构化 content(数组形式)
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "content": []any{
- map[string]any{"type": "text", "text": "Look at this"},
- map[string]any{"type": "text", "text": "And this too"},
- },
- },
- },
- SessionContext: ctx,
- }
+ content := []any{map[string]any{"type": "text", "text": "Look at this"}, map[string]any{"type": "text", "text": "And this too"}}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{msg("user", content)}, ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "structured content should produce a hash")
@@ -728,100 +391,37 @@ func TestGenerateSessionHash_StructuredContent(t *testing.T) {
func TestGenerateSessionHash_ArraySystemPrompt(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 数组格式的 system prompt
- parsed := &ParsedRequest{
- System: []any{
- map[string]any{"type": "text", "text": "You are a helpful assistant."},
- map[string]any{"type": "text", "text": "Be concise."},
- },
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: ctx,
- }
+ system := []any{map[string]any{"type": "text", "text": "You are a helpful assistant."}, map[string]any{"type": "text", "text": "Be concise."}}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(system, []any{msg("user", "hello")}, ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "array system prompt should produce a hash")
}
-// ============ SessionContext 与 cache_control 优先级 ============
-
func TestGenerateSessionHash_CacheControlOverridesSessionContext(t *testing.T) {
svc := &GatewayService{}
-
- // 当有 cache_control: ephemeral 时,使用第 2 级优先级
- // SessionContext 不应影响结果
- parsed1 := &ParsedRequest{
- System: []any{
- map[string]any{
- "type": "text",
- "text": "You are a tool-specific assistant.",
- "cache_control": map[string]any{"type": "ephemeral"},
- },
- },
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "ua1",
- APIKeyID: 100,
- },
- }
- parsed2 := &ParsedRequest{
- System: []any{
- map[string]any{
- "type": "text",
- "text": "You are a tool-specific assistant.",
- "cache_control": map[string]any{"type": "ephemeral"},
- },
- },
- HasSystem: true,
- Messages: []any{
- map[string]any{"role": "user", "content": "hello"},
- },
- SessionContext: &SessionContext{
- ClientIP: "2.2.2.2",
- UserAgent: "ua2",
- APIKeyID: 200,
- },
- }
+ system := []any{map[string]any{"type": "text", "text": "You are a tool-specific assistant.", "cache_control": map[string]any{"type": "ephemeral"}}}
+ body := anthropicSessionBody(system, []any{msg("user", "hello")}, "")
+ parsed1 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: "ua1", APIKeyID: 100})
+ parsed2 := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "2.2.2.2", UserAgent: "ua2", APIKeyID: 200})
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
require.Equal(t, h1, h2, "cache_control ephemeral has higher priority, SessionContext should not affect result")
}
-// ============ 边界情况 ============
-
func TestGenerateSessionHash_EmptyMessages(t *testing.T) {
svc := &GatewayService{}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{}, ""), &SessionContext{ClientIP: "1.1.1.1", UserAgent: "test", APIKeyID: 1})
- parsed := &ParsedRequest{
- Messages: []any{},
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "test",
- APIKeyID: 1,
- },
- }
-
- // 空 messages + 只有 SessionContext 时,combined.Len() > 0 因为有 context 写入
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "empty messages with SessionContext should still produce a hash from context")
}
func TestGenerateSessionHash_EmptyMessagesNoContext(t *testing.T) {
svc := &GatewayService{}
-
- parsed := &ParsedRequest{
- Messages: []any{},
- }
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody(nil, []any{}, ""), nil)
h := svc.GenerateSessionHash(parsed)
require.Empty(t, h, "empty messages without SessionContext should produce empty hash")
@@ -829,98 +429,37 @@ func TestGenerateSessionHash_EmptyMessagesNoContext(t *testing.T) {
func TestGenerateSessionHash_SessionContextWithEmptyFields(t *testing.T) {
svc := &GatewayService{}
-
- // SessionContext 字段为空字符串和零值时仍应影响 hash
- withEmptyCtx := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- SessionContext: &SessionContext{
- ClientIP: "",
- UserAgent: "",
- APIKeyID: 0,
- },
- }
- withoutCtx := &ParsedRequest{
- Messages: []any{
- map[string]any{"role": "user", "content": "test"},
- },
- }
+ body := anthropicSessionBody(nil, []any{msg("user", "test")}, "")
+ withEmptyCtx := mustParseSessionHashRequest(t, body, &SessionContext{ClientIP: "", UserAgent: "", APIKeyID: 0})
+ withoutCtx := mustParseSessionHashRequest(t, body, nil)
h1 := svc.GenerateSessionHash(withEmptyCtx)
h2 := svc.GenerateSessionHash(withoutCtx)
- // 有 SessionContext(即使字段为空)仍然会写入分隔符 "::" 等
require.NotEqual(t, h1, h2, "empty-field SessionContext should still differ from nil SessionContext")
}
-// ============ 长对话历史测试 ============
-
func TestGenerateSessionHash_LongConversation(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "test", APIKeyID: 1}
-
- // 构建 20 轮对话
messages := make([]any, 0, 40)
for i := 0; i < 20; i++ {
- messages = append(messages, map[string]any{
- "role": "user",
- "content": "user message " + string(rune('A'+i)),
- })
- messages = append(messages, map[string]any{
- "role": "assistant",
- "content": "assistant reply " + string(rune('A'+i)),
- })
- }
-
- parsed := &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: messages,
- SessionContext: ctx,
+ messages = append(messages, msg("user", "user message "+string(rune('A'+i))))
+ messages = append(messages, msg("assistant", "assistant reply "+string(rune('A'+i))))
}
+ parsed := mustParseSessionHashRequest(t, anthropicSessionBody("System prompt", messages, ""), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h)
- // 再加一轮应该不同
- moreMessages := make([]any, len(messages)+2)
- copy(moreMessages, messages)
- moreMessages[len(messages)] = map[string]any{"role": "user", "content": "one more"}
- moreMessages[len(messages)+1] = map[string]any{"role": "assistant", "content": "ok"}
-
- parsed2 := &ParsedRequest{
- System: "System prompt",
- HasSystem: true,
- Messages: moreMessages,
- SessionContext: ctx,
- }
-
+ moreMessages := append(append([]any{}, messages...), msg("user", "one more"), msg("assistant", "ok"))
+ parsed2 := mustParseSessionHashRequest(t, anthropicSessionBody("System prompt", moreMessages, ""), ctx)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h, h2, "adding more messages to long conversation should change hash")
}
-// ============ Gemini 原生格式 session hash 测试 ============
-
func TestGenerateSessionHash_GeminiContentsProducesHash(t *testing.T) {
svc := &GatewayService{}
-
- // Gemini 格式: contents[].parts[].text
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Hello from Gemini"},
- },
- },
- },
- SessionContext: &SessionContext{
- ClientIP: "1.2.3.4",
- UserAgent: "gemini-cli",
- APIKeyID: 1,
- },
- }
+ parsed := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Hello from Gemini")}), &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1})
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "Gemini contents with parts should produce a non-empty hash")
@@ -928,31 +467,9 @@ func TestGenerateSessionHash_GeminiContentsProducesHash(t *testing.T) {
func TestGenerateSessionHash_GeminiDifferentContentsDifferentHash(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- parsed1 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Hello"},
- },
- },
- },
- SessionContext: ctx,
- }
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Goodbye"},
- },
- },
- },
- SessionContext: ctx,
- }
+ parsed1 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Hello")}), ctx)
+ parsed2 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Goodbye")}), ctx)
h1 := svc.GenerateSessionHash(parsed1)
h2 := svc.GenerateSessionHash(parsed2)
@@ -961,28 +478,9 @@ func TestGenerateSessionHash_GeminiDifferentContentsDifferentHash(t *testing.T)
func TestGenerateSessionHash_GeminiSameContentsSameHash(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- mk := func() *ParsedRequest {
- return &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Hello"},
- },
- },
- map[string]any{
- "role": "model",
- "parts": []any{
- map[string]any{"text": "Hi there!"},
- },
- },
- },
- SessionContext: ctx,
- }
- }
+ body := geminiSessionBody(nil, []any{geminiMsg("user", "Hello"), geminiMsg("model", "Hi there!")})
+ mk := func() *ParsedRequest { return mustParseGeminiSessionHashRequest(t, body, ctx) }
h1 := svc.GenerateSessionHash(mk())
h2 := svc.GenerateSessionHash(mk())
@@ -991,36 +489,9 @@ func TestGenerateSessionHash_GeminiSameContentsSameHash(t *testing.T) {
func TestGenerateSessionHash_GeminiMultiTurnHashChanges(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- round1 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: ctx,
- }
-
- round2 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- map[string]any{
- "role": "model",
- "parts": []any{map[string]any{"text": "Hi!"}},
- },
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "How are you?"}},
- },
- },
- SessionContext: ctx,
- }
+ round1 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "hello")}), ctx)
+ round2 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "hello"), geminiMsg("model", "Hi!"), geminiMsg("user", "How are you?")}), ctx)
h1 := svc.GenerateSessionHash(round1)
h2 := svc.GenerateSessionHash(round2)
@@ -1031,34 +502,9 @@ func TestGenerateSessionHash_GeminiMultiTurnHashChanges(t *testing.T) {
func TestGenerateSessionHash_GeminiDifferentUsersSameContentDifferentHash(t *testing.T) {
svc := &GatewayService{}
-
- // 核心场景:两个不同用户发送相同 Gemini 格式消息应得到不同 hash
- user1 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: &SessionContext{
- ClientIP: "1.1.1.1",
- UserAgent: "gemini-cli",
- APIKeyID: 10,
- },
- }
- user2 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: &SessionContext{
- ClientIP: "2.2.2.2",
- UserAgent: "gemini-cli",
- APIKeyID: 20,
- },
- }
+ body := geminiSessionBody(nil, []any{geminiMsg("user", "hello")})
+ user1 := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "1.1.1.1", UserAgent: "gemini-cli", APIKeyID: 10})
+ user2 := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "2.2.2.2", UserAgent: "gemini-cli", APIKeyID: 20})
h1 := svc.GenerateSessionHash(user1)
h2 := svc.GenerateSessionHash(user2)
@@ -1067,31 +513,9 @@ func TestGenerateSessionHash_GeminiDifferentUsersSameContentDifferentHash(t *tes
func TestGenerateSessionHash_GeminiSystemInstructionAffectsHash(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- // systemInstruction 经 ParseGatewayRequest 解析后存入 parsed.System
- withSys := &ParsedRequest{
- System: []any{
- map[string]any{"text": "You are a coding assistant."},
- },
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: ctx,
- }
- withoutSys := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{map[string]any{"text": "hello"}},
- },
- },
- SessionContext: ctx,
- }
+ withSys := mustParseGeminiSessionHashRequest(t, geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "hello")}), ctx)
+ withoutSys := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "hello")}), ctx)
h1 := svc.GenerateSessionHash(withSys)
h2 := svc.GenerateSessionHash(withoutSys)
@@ -1100,64 +524,21 @@ func TestGenerateSessionHash_GeminiSystemInstructionAffectsHash(t *testing.T) {
func TestGenerateSessionHash_GeminiMultiPartMessage(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- // 多 parts 的消息
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Part 1"},
- map[string]any{"text": "Part 2"},
- map[string]any{"text": "Part 3"},
- },
- },
- },
- SessionContext: ctx,
- }
-
+ parsed := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Part 1", "Part 2", "Part 3")}), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "multi-part Gemini message should produce a hash")
- // 不同内容的多 parts
- parsed2 := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Part 1"},
- map[string]any{"text": "CHANGED"},
- map[string]any{"text": "Part 3"},
- },
- },
- },
- SessionContext: ctx,
- }
-
+ parsed2 := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, []any{geminiMsg("user", "Part 1", "CHANGED", "Part 3")}), ctx)
h2 := svc.GenerateSessionHash(parsed2)
require.NotEqual(t, h, h2, "changing a part should change the hash")
}
func TestGenerateSessionHash_GeminiNonTextPartsIgnored(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "1.2.3.4", UserAgent: "gemini-cli", APIKeyID: 1}
-
- // 含非 text 类型 parts(如 inline_data),应被跳过但不报错
- parsed := &ParsedRequest{
- Messages: []any{
- map[string]any{
- "role": "user",
- "parts": []any{
- map[string]any{"text": "Describe this image"},
- map[string]any{"inline_data": map[string]any{"mime_type": "image/png", "data": "base64..."}},
- },
- },
- },
- SessionContext: ctx,
- }
+ content := []any{map[string]any{"role": "user", "parts": []any{map[string]any{"text": "Describe this image"}, map[string]any{"inline_data": map[string]any{"mime_type": "image/png", "data": "base64..."}}}}}
+ parsed := mustParseGeminiSessionHashRequest(t, geminiSessionBody(nil, content), ctx)
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "Gemini message with mixed parts should still produce a hash from text parts")
@@ -1165,107 +546,41 @@ func TestGenerateSessionHash_GeminiNonTextPartsIgnored(t *testing.T) {
func TestGenerateSessionHash_GeminiMultiTurnHashNotSticky(t *testing.T) {
svc := &GatewayService{}
-
ctx := &SessionContext{ClientIP: "10.0.0.1", UserAgent: "gemini-cli", APIKeyID: 42}
+ rounds := []string{
+ geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "Write a Go function")}),
+ geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "Write a Go function"), geminiMsg("model", "func hello() {}"), geminiMsg("user", "Add error handling")}),
+ geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "Write a Go function"), geminiMsg("model", "func hello() {}"), geminiMsg("user", "Add error handling"), geminiMsg("model", "func hello() error { return nil }"), geminiMsg("user", "Now add tests")}),
+ }
- // 模拟同一 Gemini 会话的三轮请求,每轮 contents 累积增长。
- // 验证预期行为:每轮 hash 都不同,即 GenerateSessionHash 不具备跨轮粘性。
- // 这是 by-design 的——Gemini 的跨轮粘性由 Digest Fallback(BuildGeminiDigestChain)负责。
- round1Body := []byte(`{
- "systemInstruction": {"parts": [{"text": "You are a coding assistant."}]},
- "contents": [
- {"role": "user", "parts": [{"text": "Write a Go function"}]}
- ]
- }`)
- round2Body := []byte(`{
- "systemInstruction": {"parts": [{"text": "You are a coding assistant."}]},
- "contents": [
- {"role": "user", "parts": [{"text": "Write a Go function"}]},
- {"role": "model", "parts": [{"text": "func hello() {}"}]},
- {"role": "user", "parts": [{"text": "Add error handling"}]}
- ]
- }`)
- round3Body := []byte(`{
- "systemInstruction": {"parts": [{"text": "You are a coding assistant."}]},
- "contents": [
- {"role": "user", "parts": [{"text": "Write a Go function"}]},
- {"role": "model", "parts": [{"text": "func hello() {}"}]},
- {"role": "user", "parts": [{"text": "Add error handling"}]},
- {"role": "model", "parts": [{"text": "func hello() error { return nil }"}]},
- {"role": "user", "parts": [{"text": "Now add tests"}]}
- ]
- }`)
-
- hashes := make([]string, 3)
- for i, body := range [][]byte{round1Body, round2Body, round3Body} {
- parsed, err := ParseGatewayRequest(body, "gemini")
- require.NoError(t, err)
- parsed.SessionContext = ctx
+ hashes := make([]string, len(rounds))
+ for i, body := range rounds {
+ parsed := mustParseGeminiSessionHashRequest(t, body, ctx)
hashes[i] = svc.GenerateSessionHash(parsed)
require.NotEmpty(t, hashes[i], "round %d hash should not be empty", i+1)
}
-
- // 每轮 hash 都不同——这是预期行为
require.NotEqual(t, hashes[0], hashes[1], "round 1 vs 2 hash should differ (contents grow)")
require.NotEqual(t, hashes[1], hashes[2], "round 2 vs 3 hash should differ (contents grow)")
require.NotEqual(t, hashes[0], hashes[2], "round 1 vs 3 hash should differ")
- // 同一轮重试应产生相同 hash
- parsed1Again, err := ParseGatewayRequest(round2Body, "gemini")
- require.NoError(t, err)
- parsed1Again.SessionContext = ctx
- h2Again := svc.GenerateSessionHash(parsed1Again)
+ parsedAgain := mustParseGeminiSessionHashRequest(t, rounds[1], ctx)
+ h2Again := svc.GenerateSessionHash(parsedAgain)
require.Equal(t, hashes[1], h2Again, "retry of same round should produce same hash")
}
func TestGenerateSessionHash_GeminiEndToEnd(t *testing.T) {
svc := &GatewayService{}
-
- // 端到端测试:模拟 ParseGatewayRequest + GenerateSessionHash 完整流程
- body := []byte(`{
- "model": "gemini-2.5-pro",
- "systemInstruction": {
- "parts": [{"text": "You are a coding assistant."}]
- },
- "contents": [
- {"role": "user", "parts": [{"text": "Write a Go function"}]},
- {"role": "model", "parts": [{"text": "Here is a function..."}]},
- {"role": "user", "parts": [{"text": "Now add error handling"}]}
- ]
- }`)
-
- parsed, err := ParseGatewayRequest(body, "gemini")
- require.NoError(t, err)
- parsed.SessionContext = &SessionContext{
- ClientIP: "10.0.0.1",
- UserAgent: "gemini-cli/1.0",
- APIKeyID: 42,
- }
+ body := geminiSessionBody([]any{map[string]any{"text": "You are a coding assistant."}}, []any{geminiMsg("user", "Write a Go function"), geminiMsg("model", "Here is a function..."), geminiMsg("user", "Now add error handling")})
+ parsed := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "10.0.0.1", UserAgent: "gemini-cli/1.0", APIKeyID: 42})
h := svc.GenerateSessionHash(parsed)
require.NotEmpty(t, h, "end-to-end Gemini flow should produce a hash")
- // 同一请求再次解析应产生相同 hash
- parsed2, err := ParseGatewayRequest(body, "gemini")
- require.NoError(t, err)
- parsed2.SessionContext = &SessionContext{
- ClientIP: "10.0.0.1",
- UserAgent: "gemini-cli/1.0",
- APIKeyID: 42,
- }
-
+ parsed2 := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "10.0.0.1", UserAgent: "gemini-cli/1.0", APIKeyID: 42})
h2 := svc.GenerateSessionHash(parsed2)
require.Equal(t, h, h2, "same request should produce same hash")
- // 不同用户发送相同请求应产生不同 hash
- parsed3, err := ParseGatewayRequest(body, "gemini")
- require.NoError(t, err)
- parsed3.SessionContext = &SessionContext{
- ClientIP: "10.0.0.2",
- UserAgent: "gemini-cli/1.0",
- APIKeyID: 99,
- }
-
+ parsed3 := mustParseGeminiSessionHashRequest(t, body, &SessionContext{ClientIP: "10.0.0.2", UserAgent: "gemini-cli/1.0", APIKeyID: 99})
h3 := svc.GenerateSessionHash(parsed3)
require.NotEqual(t, h, h3, "different user with same Gemini request should get different hash")
}
diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go
index f6155352..9aa2a52f 100644
--- a/backend/internal/service/group.go
+++ b/backend/internal/service/group.go
@@ -8,6 +8,7 @@ import (
)
type OpenAIMessagesDispatchModelConfig = domain.OpenAIMessagesDispatchModelConfig
+type GroupModelsListConfig = domain.GroupModelsListConfig
type Group struct {
ID int64
@@ -61,6 +62,7 @@ type Group struct {
RequirePrivacySet bool // 调度时仅允许 privacy 已成功设置的账号(OpenAI/Antigravity/Anthropic/Gemini)
DefaultMappedModel string
MessagesDispatchModelConfig OpenAIMessagesDispatchModelConfig
+ ModelsListConfig GroupModelsListConfig
// RPMLimit 分组级每分钟请求数上限(0 = 不限制)。
// 一旦设置即接管该分组用户的限流(覆盖用户级 rpm_limit),可被 user-group rpm_override 进一步覆盖。
diff --git a/backend/internal/service/group_models_list.go b/backend/internal/service/group_models_list.go
new file mode 100644
index 00000000..b10de724
--- /dev/null
+++ b/backend/internal/service/group_models_list.go
@@ -0,0 +1,32 @@
+package service
+
+import "strings"
+
+func normalizeGroupModelsListConfig(cfg GroupModelsListConfig) GroupModelsListConfig {
+ out := GroupModelsListConfig{Enabled: cfg.Enabled}
+ if len(cfg.Models) == 0 {
+ return out
+ }
+
+ seen := make(map[string]struct{}, len(cfg.Models))
+ out.Models = make([]string, 0, len(cfg.Models))
+ for _, model := range cfg.Models {
+ model = strings.TrimSpace(model)
+ if model == "" {
+ continue
+ }
+ if _, ok := seen[model]; ok {
+ continue
+ }
+ seen[model] = struct{}{}
+ out.Models = append(out.Models, model)
+ }
+ if len(out.Models) == 0 {
+ out.Models = nil
+ }
+ return out
+}
+
+func (g *Group) CustomModelsListEnabled() bool {
+ return g != nil && g.ModelsListConfig.Enabled && len(g.ModelsListConfig.Models) > 0
+}
diff --git a/backend/internal/service/header_util.go b/backend/internal/service/header_util.go
index 1091070d..f8da068d 100644
--- a/backend/internal/service/header_util.go
+++ b/backend/internal/service/header_util.go
@@ -109,6 +109,20 @@ func addHeaderRaw(h http.Header, key, value string) {
h[key] = append(h[key], value)
}
+// deleteHeaderAllForms removes a header in all common key forms (raw, wire casing,
+// canonical) so subsequent setHeaderRaw will not coexist with a passthrough value
+// written under a different casing.
+func deleteHeaderAllForms(h http.Header, key string) {
+ if h == nil || key == "" {
+ return
+ }
+ h.Del(key) // canonical
+ delete(h, key)
+ if wk := resolveWireCasing(key); wk != key {
+ delete(h, wk)
+ }
+}
+
// getHeaderRaw reads a header value, trying multiple key forms to handle the mismatch
// between Go canonical keys, wire casing keys, and raw keys:
// 1. exact key as provided
diff --git a/backend/internal/service/image_generation_intent.go b/backend/internal/service/image_generation_intent.go
index 4aca1239..80b6c66d 100644
--- a/backend/internal/service/image_generation_intent.go
+++ b/backend/internal/service/image_generation_intent.go
@@ -1,7 +1,6 @@
package service
import (
- "encoding/json"
"strings"
"github.com/tidwall/gjson"
@@ -91,7 +90,7 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool {
}
found := false
tools.ForEach(func(_, item gjson.Result) bool {
- if strings.TrimSpace(item.Get("type").String()) == "image_generation" {
+ if openAIJSONString(item.Get("type")) == "image_generation" {
found = true
return false
}
@@ -100,6 +99,36 @@ func openAIJSONToolsContainImageGeneration(tools gjson.Result) bool {
return found
}
+func openAIRequestBodyHasImageGenerationTool(body []byte) bool {
+ if len(body) == 0 || !gjson.ValidBytes(body) {
+ return false
+ }
+ return openAIJSONToolsContainImageGeneration(gjson.GetBytes(body, "tools"))
+}
+
+func openAIRequestBodyImageGenerationToolNeedsNormalization(body []byte) bool {
+ if len(body) == 0 || !gjson.ValidBytes(body) {
+ return false
+ }
+ tools := gjson.GetBytes(body, "tools")
+ if !tools.IsArray() {
+ return false
+ }
+ needsNormalization := false
+ tools.ForEach(func(_, item gjson.Result) bool {
+ if openAIJSONString(item.Get("type")) != "image_generation" {
+ return true
+ }
+ // 只有旧字段需要迁移时才进入 map 修改,纯计费读取保持 raw 路径。
+ if item.Get("format").Exists() || item.Get("compression").Exists() {
+ needsNormalization = true
+ return false
+ }
+ return true
+ })
+ return needsNormalization
+}
+
func openAIJSONToolChoiceSelectsImageGeneration(choice gjson.Result) bool {
if !choice.Exists() {
return false
@@ -159,17 +188,6 @@ func apiKeyGroup(apiKey *APIKey) *Group {
return apiKey.Group
}
-func cloneRequestMapForImageIntent(body []byte) map[string]any {
- if len(body) == 0 {
- return nil
- }
- var out map[string]any
- if err := json.Unmarshal(body, &out); err != nil {
- return nil
- }
- return out
-}
-
type OpenAIResponsesImageBillingConfig struct {
Model string
SizeTier string
@@ -225,8 +243,43 @@ func resolveOpenAIResponsesImageBillingConfigFromBody(body []byte, fallbackModel
}
func resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body []byte, fallbackModel string) (OpenAIResponsesImageBillingConfig, error) {
- reqBody := cloneRequestMapForImageIntent(body)
- return resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, fallbackModel)
+ imageModel := ""
+ imageSize := ""
+ hasImageTool := false
+ if len(body) > 0 && gjson.ValidBytes(body) {
+ tools := gjson.GetBytes(body, "tools")
+ if tools.IsArray() {
+ tools.ForEach(func(_, item gjson.Result) bool {
+ if openAIJSONString(item.Get("type")) != "image_generation" {
+ return true
+ }
+ hasImageTool = true
+ imageModel = openAIJSONString(item.Get("model"))
+ imageSize = openAIJSONString(item.Get("size"))
+ return false
+ })
+ }
+ if imageSize == "" {
+ imageSize = openAIJSONString(gjson.GetBytes(body, "size"))
+ }
+ if imageModel == "" {
+ bodyModel := openAIJSONString(gjson.GetBytes(body, "model"))
+ if isOpenAIImageBillingModelAlias(bodyModel) || !hasImageTool {
+ imageModel = bodyModel
+ }
+ }
+ }
+ if imageModel == "" && hasImageTool {
+ imageModel = "gpt-image-2"
+ }
+ if imageModel == "" {
+ imageModel = strings.TrimSpace(fallbackModel)
+ }
+ return OpenAIResponsesImageBillingConfig{
+ Model: imageModel,
+ SizeTier: normalizeOpenAIImageSizeTier(imageSize),
+ InputSize: imageSize,
+ }, nil
}
func isOpenAIImageBillingModelAlias(model string) bool {
@@ -236,3 +289,10 @@ func isOpenAIImageBillingModelAlias(model string) bool {
}
return isOpenAIImageGenerationModel(normalized) || strings.Contains(normalized, "image")
}
+
+func openAIJSONString(value gjson.Result) string {
+ if value.Type != gjson.String {
+ return ""
+ }
+ return strings.TrimSpace(value.String())
+}
diff --git a/backend/internal/service/image_generation_intent_test.go b/backend/internal/service/image_generation_intent_test.go
index 4621e9d9..59aab39c 100644
--- a/backend/internal/service/image_generation_intent_test.go
+++ b/backend/internal/service/image_generation_intent_test.go
@@ -84,6 +84,17 @@ func TestResolveOpenAIResponsesImageBillingConfigToolModelWins(t *testing.T) {
require.Equal(t, "2K", imageSize)
}
+func TestResolveOpenAIResponsesImageBillingConfigFromBodyIgnoresUnrelatedLargeInput(t *testing.T) {
+ cfg, err := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(
+ []byte(`{"model":"mapped-text-model","tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[{"type":"message","content":[{"type":"input_text","text":"hi","nonce":1e1000000}]}]}`),
+ "requested-model",
+ )
+ require.NoError(t, err)
+ require.Equal(t, "gpt-image-2", cfg.Model)
+ require.Equal(t, "2K", cfg.SizeTier)
+ require.Equal(t, "2048x1152", cfg.InputSize)
+}
+
func TestResolveOpenAIResponsesImageBillingConfigSupportsOfficialAndCustomSizes(t *testing.T) {
tests := []struct {
name string
diff --git a/backend/internal/service/model_not_found_error.go b/backend/internal/service/model_not_found_error.go
new file mode 100644
index 00000000..910a97d8
--- /dev/null
+++ b/backend/internal/service/model_not_found_error.go
@@ -0,0 +1,44 @@
+package service
+
+import (
+ "net/http"
+ "strings"
+)
+
+var upstreamModelNotFoundKeywords = []string{"model not found", "unknown model", "not found"}
+
+func isUpstreamModelNotFoundError(statusCode int, body []byte) bool {
+ if statusCode != http.StatusNotFound {
+ return false
+ }
+ normalized := normalizeModelNotFoundBody(body)
+ if normalized == "" || !strings.Contains(normalized, "model") {
+ return false
+ }
+ return containsModelNotFoundKeyword(normalized)
+}
+
+func isModelNotFoundError(statusCode int, body []byte) bool {
+ return isUpstreamModelNotFoundError(statusCode, body) || statusCode == http.StatusNotFound
+}
+
+func containsModelNotFoundKeyword(normalizedBody string) bool {
+ if normalizedBody == "" {
+ return false
+ }
+ for _, keyword := range upstreamModelNotFoundKeywords {
+ if strings.Contains(normalizedBody, keyword) {
+ return true
+ }
+ }
+ return false
+}
+
+func normalizeModelNotFoundBody(body []byte) string {
+ if len(body) == 0 {
+ return ""
+ }
+ normalized := strings.ToLower(string(body))
+ normalized = strings.NewReplacer("_", " ", "-", " ", "\n", " ", "\r", " ", "\t", " ").Replace(normalized)
+ return strings.Join(strings.Fields(normalized), " ")
+}
diff --git a/backend/internal/service/model_not_found_error_test.go b/backend/internal/service/model_not_found_error_test.go
new file mode 100644
index 00000000..a87340eb
--- /dev/null
+++ b/backend/internal/service/model_not_found_error_test.go
@@ -0,0 +1,66 @@
+package service
+
+import (
+ "net/http"
+ "testing"
+)
+
+func TestIsUpstreamModelNotFoundError(t *testing.T) {
+ tests := []struct {
+ name string
+ statusCode int
+ body []byte
+ want bool
+ }{
+ {
+ name: "404 model not found message",
+ statusCode: http.StatusNotFound,
+ body: []byte(`{"error":{"message":"model not found"}}`),
+ want: true,
+ },
+ {
+ name: "404 model_not_found code",
+ statusCode: http.StatusNotFound,
+ body: []byte(`{"error":{"code":"model_not_found","message":"The requested model was not found"}}`),
+ want: true,
+ },
+ {
+ name: "404 unknown model message",
+ statusCode: http.StatusNotFound,
+ body: []byte(`{"error":{"message":"unknown model gpt-5.4"}}`),
+ want: true,
+ },
+ {
+ name: "404 endpoint not found is not model specific",
+ statusCode: http.StatusNotFound,
+ body: []byte(`{"error":{"message":"endpoint not found"}}`),
+ want: false,
+ },
+ {
+ name: "404 arbitrary body is not model specific",
+ statusCode: http.StatusNotFound,
+ body: []byte(`404 page not found`),
+ want: false,
+ },
+ {
+ name: "non 404 does not match",
+ statusCode: http.StatusBadRequest,
+ body: []byte(`{"error":{"message":"model not found"}}`),
+ want: false,
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ if got := isUpstreamModelNotFoundError(tt.statusCode, tt.body); got != tt.want {
+ t.Fatalf("isUpstreamModelNotFoundError() = %v, want %v", got, tt.want)
+ }
+ })
+ }
+}
+
+func TestAntigravityModelNotFoundKeepsBare404Fallback(t *testing.T) {
+ if !isModelNotFoundError(http.StatusNotFound, []byte(`endpoint not found`)) {
+ t.Fatal("antigravity model-not-found helper should keep bare 404 fallback")
+ }
+}
diff --git a/backend/internal/service/model_rate_limit.go b/backend/internal/service/model_rate_limit.go
index c45615cc..420f1446 100644
--- a/backend/internal/service/model_rate_limit.go
+++ b/backend/internal/service/model_rate_limit.go
@@ -6,7 +6,10 @@ import (
"time"
)
-const modelRateLimitsKey = "model_rate_limits"
+const (
+ modelRateLimitsKey = "model_rate_limits"
+ antigravityGeminiModelRateLimitKey = "antigravity:gemini"
+)
// isRateLimitActiveForKey 检查指定 key 的限流是否生效
func (a *Account) isRateLimitActiveForKey(key string) bool {
@@ -35,6 +38,9 @@ func (a *Account) isModelRateLimitedWithContext(ctx context.Context, requestedMo
modelKey := a.GetMappedModel(requestedModel)
if a.Platform == PlatformAntigravity {
modelKey = resolveFinalAntigravityModelKey(ctx, a, requestedModel)
+ if isAntigravityGeminiModel(modelKey) && a.isRateLimitActiveForKey(antigravityGeminiModelRateLimitKey) {
+ return true
+ }
}
modelKey = strings.TrimSpace(modelKey)
if modelKey == "" {
@@ -62,7 +68,13 @@ func (a *Account) GetModelRateLimitRemainingTimeWithContext(ctx context.Context,
if modelKey == "" {
return 0
}
- return a.getRateLimitRemainingForKey(modelKey)
+ remaining := a.getRateLimitRemainingForKey(modelKey)
+ if a.Platform == PlatformAntigravity && isAntigravityGeminiModel(modelKey) {
+ if familyRemaining := a.getRateLimitRemainingForKey(antigravityGeminiModelRateLimitKey); familyRemaining > remaining {
+ return familyRemaining
+ }
+ }
+ return remaining
}
func resolveFinalAntigravityModelKey(ctx context.Context, account *Account, requestedModel string) string {
@@ -77,6 +89,22 @@ func resolveFinalAntigravityModelKey(ctx context.Context, account *Account, requ
return modelKey
}
+func isAntigravityGeminiModel(model string) bool {
+ return strings.HasPrefix(normalizeAntigravityModelName(model), "gemini-")
+}
+
+func antigravityModelRateLimitKeys(model string) []string {
+ model = strings.TrimSpace(model)
+ if model == "" {
+ return nil
+ }
+ keys := []string{model}
+ if isAntigravityGeminiModel(model) && model != antigravityGeminiModelRateLimitKey {
+ keys = append(keys, antigravityGeminiModelRateLimitKey)
+ }
+ return keys
+}
+
func (a *Account) modelRateLimitResetAt(scope string) *time.Time {
if a == nil || a.Extra == nil || scope == "" {
return nil
diff --git a/backend/internal/service/model_rate_limit_test.go b/backend/internal/service/model_rate_limit_test.go
index b79b9688..3cce6459 100644
--- a/backend/internal/service/model_rate_limit_test.go
+++ b/backend/internal/service/model_rate_limit_test.go
@@ -121,6 +121,36 @@ func TestIsModelRateLimited(t *testing.T) {
requestedModel: "gemini-3-pro-preview",
expected: true,
},
+ {
+ name: "antigravity platform - gemini family rate limit blocks mapped preview",
+ account: &Account{
+ Platform: PlatformAntigravity,
+ Extra: map[string]any{
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": future,
+ },
+ },
+ },
+ },
+ requestedModel: "gemini-3-pro-preview",
+ expected: true,
+ },
+ {
+ name: "antigravity platform - gemini family rate limit does not block claude",
+ account: &Account{
+ Platform: PlatformAntigravity,
+ Extra: map[string]any{
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": future,
+ },
+ },
+ },
+ },
+ requestedModel: "claude-sonnet-4-5",
+ expected: false,
+ },
{
name: "non-antigravity platform - gemini-3-pro-preview NOT mapped",
account: &Account{
@@ -306,6 +336,38 @@ func TestGetModelRateLimitRemainingTime(t *testing.T) {
minExpected: 4 * time.Minute,
maxExpected: 6 * time.Minute,
},
+ {
+ name: "antigravity platform - gemini family rate limit remaining",
+ account: &Account{
+ Platform: PlatformAntigravity,
+ Extra: map[string]any{
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": future10m,
+ },
+ },
+ },
+ },
+ requestedModel: "gemini-3-pro-preview",
+ minExpected: 9 * time.Minute,
+ maxExpected: 11 * time.Minute,
+ },
+ {
+ name: "antigravity platform - gemini family remaining ignored for claude",
+ account: &Account{
+ Platform: PlatformAntigravity,
+ Extra: map[string]any{
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": future10m,
+ },
+ },
+ },
+ },
+ requestedModel: "claude-sonnet-4-5",
+ minExpected: 0,
+ maxExpected: 0,
+ },
}
for _, tt := range tests {
diff --git a/backend/internal/service/openai_account_runtime_block_fastpath.go b/backend/internal/service/openai_account_runtime_block_fastpath.go
index 41a309fd..a4e905e8 100644
--- a/backend/internal/service/openai_account_runtime_block_fastpath.go
+++ b/backend/internal/service/openai_account_runtime_block_fastpath.go
@@ -31,7 +31,7 @@ func isOpenAIAccount(account *Account) bool {
return account != nil && account.Platform == PlatformOpenAI
}
-func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte) bool {
+func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte, requestedModel ...string) bool {
stateCtx, cancel := openAIAccountStateContext(ctx)
defer cancel()
@@ -41,6 +41,9 @@ func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Cont
if s == nil || account == nil || s.rateLimitService == nil {
return false
}
+ if len(requestedModel) > 0 && s.rateLimitService.HandleUpstreamModelNotFound(stateCtx, account, requestedModel[0], statusCode, responseBody) {
+ return true
+ }
shouldDisable := s.rateLimitService.HandleUpstreamError(stateCtx, account, statusCode, headers, responseBody)
if shouldDisable {
s.BlockAccountScheduling(account, time.Time{}, "upstream_disable")
diff --git a/backend/internal/service/openai_account_runtime_block_fastpath_test.go b/backend/internal/service/openai_account_runtime_block_fastpath_test.go
index 95336e81..3784dd33 100644
--- a/backend/internal/service/openai_account_runtime_block_fastpath_test.go
+++ b/backend/internal/service/openai_account_runtime_block_fastpath_test.go
@@ -57,6 +57,28 @@ func TestOpenAIRuntimeBlocker_IgnoresNonOpenAIFromRateLimitService(t *testing.T)
require.False(t, gateway.isOpenAIAccountRuntimeBlocked(account))
}
+func TestOpenAIModelNotFound_DoesNotRuntimeBlockWholeAccount(t *testing.T) {
+ repo := &modelNotFoundAccountRepoStub{}
+ svc := &OpenAIGatewayService{
+ rateLimitService: &RateLimitService{accountRepo: repo},
+ }
+ account := openAIModelNotFoundTempAccount()
+
+ shouldDisable := svc.handleOpenAIAccountUpstreamError(
+ context.Background(),
+ account,
+ http.StatusNotFound,
+ http.Header{},
+ []byte(`{"error":{"code":"model_not_found","message":"model not found"}}`),
+ "gpt-5.4",
+ )
+
+ require.True(t, shouldDisable)
+ require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
+ require.Zero(t, repo.tempCalls)
+ require.Len(t, repo.modelRateLimitCalls, 1)
+}
+
func TestOpenAIRuntimeBlock_DoesNotShortenExistingBlock(t *testing.T) {
svc := &OpenAIGatewayService{}
account := &Account{ID: 46, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go
index a8ac391a..6e1ab3fd 100644
--- a/backend/internal/service/openai_account_scheduler.go
+++ b/backend/internal/service/openai_account_scheduler.go
@@ -41,9 +41,11 @@ type OpenAIAccountScheduleRequest struct {
GroupID *int64
SessionHash string
StickyAccountID int64
+ PreserveStickyBinding bool
PreviousResponseID string
RequestedModel string
RequiredTransport OpenAIUpstreamTransport
+ RequiredCapability OpenAIEndpointCapability
RequiredImageCapability OpenAIImagesCapability
RequireCompact bool
ExcludedIDs map[int64]struct{}
@@ -240,6 +242,12 @@ type defaultOpenAIAccountScheduler struct {
stats *openAIAccountRuntimeStats
}
+type openAIStickyEscapeConfig struct {
+ enabled bool
+ ttftMs float64
+ errorRate float64
+}
+
func newDefaultOpenAIAccountScheduler(service *OpenAIGatewayService, stats *openAIAccountRuntimeStats) OpenAIAccountScheduler {
if stats == nil {
stats = newOpenAIAccountRuntimeStats()
@@ -263,12 +271,13 @@ func (s *defaultOpenAIAccountScheduler) Select(
previousResponseID := strings.TrimSpace(req.PreviousResponseID)
if previousResponseID != "" {
- selection, err := s.service.SelectAccountByPreviousResponseID(
+ selection, err := s.service.selectAccountByPreviousResponseIDForCapability(
ctx,
req.GroupID,
previousResponseID,
req.RequestedModel,
req.ExcludedIDs,
+ req.RequiredCapability,
req.RequireCompact,
)
if err != nil {
@@ -294,7 +303,7 @@ func (s *defaultOpenAIAccountScheduler) Select(
}
}
- selection, err := s.selectBySessionHash(ctx, req)
+ selection, escapedSticky, err := s.selectBySessionHash(ctx, req)
if err != nil {
return nil, decision, err
}
@@ -305,6 +314,9 @@ func (s *defaultOpenAIAccountScheduler) Select(
decision.SelectedAccountType = selection.Account.Type
return selection, decision, nil
}
+ if escapedSticky {
+ req.PreserveStickyBinding = true
+ }
selection, candidateCount, topK, loadSkew, err := s.selectByLoadBalance(ctx, req)
decision.Layer = openAIAccountScheduleLayerLoadBalance
@@ -324,10 +336,10 @@ func (s *defaultOpenAIAccountScheduler) Select(
func (s *defaultOpenAIAccountScheduler) selectBySessionHash(
ctx context.Context,
req OpenAIAccountScheduleRequest,
-) (*AccountSelectionResult, error) {
+) (*AccountSelectionResult, bool, error) {
sessionHash := strings.TrimSpace(req.SessionHash)
if sessionHash == "" || s == nil || s.service == nil || s.service.cache == nil {
- return nil, nil
+ return nil, false, nil
}
accountID := req.StickyAccountID
@@ -335,40 +347,49 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash(
var err error
accountID, err = s.service.getStickySessionAccountID(ctx, req.GroupID, sessionHash)
if err != nil || accountID <= 0 {
- return nil, nil
+ return nil, false, nil
}
}
if accountID <= 0 {
- return nil, nil
+ return nil, false, nil
}
if req.ExcludedIDs != nil {
if _, excluded := req.ExcludedIDs[accountID]; excluded {
- return nil, nil
+ return nil, false, nil
}
}
account, err := s.service.getSchedulableAccount(ctx, accountID)
if err != nil || account == nil {
_ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash)
- return nil, nil
+ return nil, false, nil
}
if shouldClearStickySession(account, req.RequestedModel) || !account.IsOpenAI() || !account.IsSchedulable() {
_ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash)
- return nil, nil
+ return nil, false, nil
}
if !s.isAccountRequestCompatible(ctx, account, req) {
- return nil, nil
+ return nil, false, nil
}
if !s.isAccountTransportCompatible(account, req.RequiredTransport) {
_ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash)
- return nil, nil
+ return nil, false, nil
}
- account = s.service.recheckSelectedOpenAIAccountFromDB(ctx, account, req.RequestedModel, req.RequireCompact)
+ account = s.service.recheckSelectedOpenAIAccountFromDB(ctx, account, req.RequestedModel, req.RequireCompact, req.RequiredCapability)
if account == nil || !s.isAccountTransportCompatible(account, req.RequiredTransport) {
_ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash)
- return nil, nil
+ return nil, false, nil
+ }
+ escapeCfg := s.service.openAIStickyEscapeConfig()
+ if reason, errorRate, ttft, shouldEscape := s.shouldEscapeStickyAccount(accountID, escapeCfg); shouldEscape {
+ slog.Info("sticky_escape_triggered",
+ "account_id", accountID,
+ "reason", reason,
+ "error_rate", errorRate,
+ "ttft", ttft,
+ )
+ return nil, true, nil
}
-
result, acquireErr := s.service.tryAcquireAccountSlot(ctx, accountID, account.Concurrency)
if acquireErr == nil && result != nil && result.Acquired {
_ = s.service.refreshStickySessionTTL(ctx, req.GroupID, sessionHash, s.service.openAIWSSessionStickyTTL())
@@ -376,12 +397,22 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash(
Account: account,
Acquired: true,
ReleaseFunc: result.ReleaseFunc,
- }, nil
+ }, false, nil
}
cfg := s.service.schedulingConfig()
// WaitPlan.MaxConcurrency 使用 Concurrency(非 EffectiveLoadFactor),因为 WaitPlan 控制的是 Redis 实际并发槽位等待。
if s.service.concurrencyService != nil {
+ if escapeCfg.enabled && acquireErr == nil && result != nil && !result.Acquired {
+ errorRate, ttft, _ := s.stats.snapshot(accountID)
+ slog.Info("sticky_escape_triggered",
+ "account_id", accountID,
+ "reason", "concurrency_full",
+ "error_rate", errorRate,
+ "ttft", ttft,
+ )
+ return nil, true, nil
+ }
return &AccountSelectionResult{
Account: account,
WaitPlan: &AccountWaitPlan{
@@ -390,9 +421,23 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash(
Timeout: cfg.StickySessionWaitTimeout,
MaxWaiting: cfg.StickySessionMaxWaiting,
},
- }, nil
+ }, false, nil
}
- return nil, nil
+ return nil, false, nil
+}
+
+func (s *defaultOpenAIAccountScheduler) shouldEscapeStickyAccount(accountID int64, cfg openAIStickyEscapeConfig) (reason string, errorRate float64, ttft float64, shouldEscape bool) {
+ if !cfg.enabled || s == nil || s.stats == nil || accountID <= 0 {
+ return "", 0, 0, false
+ }
+ errorRate, ttft, hasTTFT := s.stats.snapshot(accountID)
+ if hasTTFT && ttft > cfg.ttftMs {
+ return "ttft", errorRate, ttft, true
+ }
+ if errorRate > cfg.errorRate {
+ return "error_rate", errorRate, ttft, true
+ }
+ return "", errorRate, ttft, false
}
type openAIAccountCandidateScore struct {
@@ -791,11 +836,11 @@ func (s *defaultOpenAIAccountScheduler) tryAcquireOpenAISelectionOrder(
compactBlocked := false
for i := 0; i < len(selectionOrder); i++ {
candidate := selectionOrder[i]
- fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false)
+ fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false, req.RequiredCapability)
if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) {
continue
}
- fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false)
+ fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false, req.RequiredCapability)
if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) {
continue
}
@@ -808,7 +853,7 @@ func (s *defaultOpenAIAccountScheduler) tryAcquireOpenAISelectionOrder(
return nil, compactBlocked, acquireErr
}
if result != nil && result.Acquired {
- if req.SessionHash != "" {
+ if req.SessionHash != "" && !req.PreserveStickyBinding {
_ = s.service.BindStickySession(ctx, req.GroupID, req.SessionHash, fresh.ID)
}
return &AccountSelectionResult{
@@ -930,11 +975,11 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance(
cfg := s.service.schedulingConfig()
// WaitPlan.MaxConcurrency 使用 Concurrency(非 EffectiveLoadFactor),因为 WaitPlan 控制的是 Redis 实际并发槽位等待。
for _, candidate := range selectionOrder {
- fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false)
+ fresh := s.service.resolveFreshSchedulableOpenAIAccount(ctx, candidate.account, req.RequestedModel, false, req.RequiredCapability)
if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) {
continue
}
- fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false)
+ fresh = s.service.recheckSelectedOpenAIAccountFromDB(ctx, fresh, req.RequestedModel, false, req.RequiredCapability)
if fresh == nil || !s.isAccountTransportCompatible(fresh, req.RequiredTransport) || !s.isAccountRequestCompatible(ctx, fresh, req) {
continue
}
@@ -973,6 +1018,13 @@ func (s *defaultOpenAIAccountScheduler) isAccountRequestCompatible(ctx context.C
if s != nil && s.service != nil && s.service.isOpenAIAccountRuntimeBlocked(account) {
return false
}
+ // Quota auto-pause must be evaluated during the initial filter too. Without it the
+ // TopK candidate pool can be filled with paused accounts and the later fresh/DB
+ // rechecks won't reach healthy accounts that fell outside TopK — manifesting as
+ // "no available accounts" even though healthy ones exist.
+ if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused {
+ return false
+ }
if req.RequestedModel != "" && !account.IsModelSupported(req.RequestedModel) {
return false
}
@@ -981,7 +1033,7 @@ func (s *defaultOpenAIAccountScheduler) isAccountRequestCompatible(ctx context.C
s.service.isUpstreamModelRestrictedByChannel(ctx, *req.GroupID, account, req.RequestedModel, req.RequireCompact) {
return false
}
- return account.SupportsOpenAIImageCapability(req.RequiredImageCapability)
+ return accountSupportsOpenAICapabilities(account, req.RequiredCapability, req.RequiredImageCapability)
}
func (s *defaultOpenAIAccountScheduler) ReportResult(accountID int64, success bool, firstTokenMs *int) {
@@ -1104,7 +1156,21 @@ func (s *OpenAIGatewayService) SelectAccountWithScheduler(
requiredTransport OpenAIUpstreamTransport,
requireCompact bool,
) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) {
- return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", requireCompact)
+ return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, "", "", requireCompact)
+}
+
+func (s *OpenAIGatewayService) SelectAccountWithSchedulerForCapability(
+ ctx context.Context,
+ groupID *int64,
+ previousResponseID string,
+ sessionHash string,
+ requestedModel string,
+ excludedIDs map[int64]struct{},
+ requiredTransport OpenAIUpstreamTransport,
+ requiredCapability OpenAIEndpointCapability,
+ requireCompact bool,
+) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) {
+ return s.selectAccountWithScheduler(ctx, groupID, previousResponseID, sessionHash, requestedModel, excludedIDs, requiredTransport, requiredCapability, "", requireCompact)
}
func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages(
@@ -1115,13 +1181,13 @@ func (s *OpenAIGatewayService) SelectAccountWithSchedulerForImages(
excludedIDs map[int64]struct{},
requiredCapability OpenAIImagesCapability,
) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) {
- selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, requiredCapability, false)
+ selection, decision, err := s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", requiredCapability, false)
if err == nil && selection != nil && selection.Account != nil {
return selection, decision, nil
}
// 如果要求 native 能力(如指定了模型)但没有可用的 APIKey 账号,回退到 basic(OAuth 账号)
if requiredCapability == OpenAIImagesCapabilityNative {
- return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, OpenAIImagesCapabilityBasic, false)
+ return s.selectAccountWithScheduler(ctx, groupID, "", sessionHash, requestedModel, excludedIDs, OpenAIUpstreamTransportHTTPSSE, "", OpenAIImagesCapabilityBasic, false)
}
return selection, decision, err
}
@@ -1134,9 +1200,11 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler(
requestedModel string,
excludedIDs map[int64]struct{},
requiredTransport OpenAIUpstreamTransport,
+ requiredCapability OpenAIEndpointCapability,
requiredImageCapability OpenAIImagesCapability,
requireCompact bool,
) (*AccountSelectionResult, OpenAIAccountScheduleDecision, error) {
+ ctx = s.withOpenAIQuotaAutoPauseContext(ctx)
decision := OpenAIAccountScheduleDecision{}
scheduler := s.getOpenAIAccountScheduler(ctx)
if scheduler == nil {
@@ -1144,14 +1212,14 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler(
if requiredTransport == OpenAIUpstreamTransportAny || requiredTransport == OpenAIUpstreamTransportHTTPSSE {
effectiveExcludedIDs := cloneExcludedAccountIDs(excludedIDs)
for {
- selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact)
+ selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact, requiredCapability)
if err != nil {
return nil, decision, err
}
if selection == nil || selection.Account == nil {
return selection, decision, nil
}
- if selection.Account.SupportsOpenAIImageCapability(requiredImageCapability) {
+ if accountSupportsOpenAICapabilities(selection.Account, requiredCapability, requiredImageCapability) {
return selection, decision, nil
}
if selection.ReleaseFunc != nil {
@@ -1169,14 +1237,15 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler(
effectiveExcludedIDs := cloneExcludedAccountIDs(excludedIDs)
for {
- selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact)
+ selection, err := s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, effectiveExcludedIDs, requireCompact, requiredCapability)
if err != nil {
return nil, decision, err
}
if selection == nil || selection.Account == nil {
return selection, decision, nil
}
- if s.isOpenAIAccountTransportCompatible(selection.Account, requiredTransport) {
+ if s.isOpenAIAccountTransportCompatible(selection.Account, requiredTransport) &&
+ accountSupportsOpenAICapabilities(selection.Account, requiredCapability, requiredImageCapability) {
return selection, decision, nil
}
if selection.ReleaseFunc != nil {
@@ -1213,12 +1282,21 @@ func (s *OpenAIGatewayService) selectAccountWithScheduler(
PreviousResponseID: previousResponseID,
RequestedModel: requestedModel,
RequiredTransport: requiredTransport,
+ RequiredCapability: requiredCapability,
RequiredImageCapability: requiredImageCapability,
RequireCompact: requireCompact,
ExcludedIDs: excludedIDs,
})
}
+func accountSupportsOpenAICapabilities(account *Account, requiredCapability OpenAIEndpointCapability, requiredImageCapability OpenAIImagesCapability) bool {
+ if account == nil {
+ return false
+ }
+ return account.SupportsOpenAIEndpointCapability(requiredCapability) &&
+ account.SupportsOpenAIImageCapability(requiredImageCapability)
+}
+
func cloneExcludedAccountIDs(excludedIDs map[int64]struct{}) map[int64]struct{} {
if len(excludedIDs) == 0 {
return nil
@@ -1278,6 +1356,37 @@ func (s *OpenAIGatewayService) openAIWSLBTopK() int {
return 7
}
+func (s *OpenAIGatewayService) openAIStickyEscapeConfig() openAIStickyEscapeConfig {
+ if s != nil && s.cfg != nil {
+ cfg := s.cfg.Gateway.OpenAIScheduler
+ enabled := cfg.StickyEscapeEnabled
+ if !enabled && cfg.StickyEscapeTTFTMs == 0 && cfg.StickyEscapeErrorRate == 0 {
+ enabled = true
+ }
+ ttftMs := float64(cfg.StickyEscapeTTFTMs)
+ if ttftMs <= 0 {
+ ttftMs = 15000
+ }
+ errorRate := cfg.StickyEscapeErrorRate
+ if errorRate < 0 || errorRate > 1 {
+ errorRate = 0.5
+ }
+ if errorRate == 0 && cfg.StickyEscapeTTFTMs == 0 && cfg.StickyEscapeErrorRate == 0 {
+ errorRate = 0.5
+ }
+ return openAIStickyEscapeConfig{
+ enabled: enabled,
+ ttftMs: ttftMs,
+ errorRate: errorRate,
+ }
+ }
+ return openAIStickyEscapeConfig{
+ enabled: true,
+ ttftMs: 15000,
+ errorRate: 0.5,
+ }
+}
+
func (s *OpenAIGatewayService) openAIWSSchedulerWeights() GatewayOpenAIWSSchedulerScoreWeightsView {
if s != nil && s.cfg != nil {
return GatewayOpenAIWSSchedulerScoreWeightsView{
diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go
index 0950ee54..845d60c4 100644
--- a/backend/internal/service/openai_account_scheduler_test.go
+++ b/backend/internal/service/openai_account_scheduler_test.go
@@ -393,6 +393,64 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Require
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
}
+func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_EmbeddingsSkipsChatOnlyAccount(t *testing.T) {
+ resetOpenAIAdvancedSchedulerSettingCacheForTest()
+
+ ctx := context.Background()
+ groupID := int64(10110)
+ accounts := []Account{
+ {
+ ID: 36031,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Credentials: map[string]any{
+ "openai_capabilities": []any{"chat_completions"},
+ },
+ },
+ {
+ ID: 36032,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 5,
+ Credentials: map[string]any{
+ "openai_capabilities": []any{"chat_completions", "embeddings"},
+ },
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Gateway.Scheduling.LoadBatchEnabled = false
+ svc := &OpenAIGatewayService{
+ accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
+ cache: &schedulerTestGatewayCache{},
+ cfg: cfg,
+ concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}),
+ }
+
+ selection, decision, err := svc.SelectAccountWithSchedulerForCapability(
+ ctx,
+ &groupID,
+ "",
+ "",
+ "text-embedding-3-small",
+ nil,
+ OpenAIUpstreamTransportHTTPSSE,
+ OpenAIEndpointCapabilityEmbeddings,
+ false,
+ )
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(36032), selection.Account.ID)
+ require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
+}
+
func TestOpenAIGatewayService_SelectAccountWithScheduler_EnabledUsesAdvancedPreviousResponseRouting(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
@@ -458,6 +516,141 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_EnabledUsesAdvancedPrev
require.True(t, decision.StickyPreviousHit)
}
+func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkipsChatOnlyAccount(t *testing.T) {
+ resetOpenAIAdvancedSchedulerSettingCacheForTest()
+
+ ctx := context.Background()
+ groupID := int64(10111)
+ accounts := []Account{
+ {
+ ID: 37011,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Credentials: map[string]any{
+ "openai_capabilities": []any{"chat_completions"},
+ },
+ },
+ {
+ ID: 37012,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 5,
+ Credentials: map[string]any{
+ "openai_capabilities": []any{"chat_completions", "embeddings"},
+ },
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Gateway.Scheduling.LoadBatchEnabled = false
+ svc := &OpenAIGatewayService{
+ accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
+ cache: &schedulerTestGatewayCache{},
+ cfg: cfg,
+ rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
+ concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}),
+ }
+
+ selection, decision, err := svc.SelectAccountWithSchedulerForCapability(
+ ctx,
+ &groupID,
+ "",
+ "",
+ "text-embedding-3-small",
+ nil,
+ OpenAIUpstreamTransportHTTPSSE,
+ OpenAIEndpointCapabilityEmbeddings,
+ false,
+ )
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(37012), selection.Account.ID)
+ require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
+ require.Equal(t, 1, decision.CandidateCount)
+}
+
+func TestOpenAIGatewayService_SelectAccountWithScheduler_Enabled_EmbeddingsSkipsChatOnlyStickyBindings(t *testing.T) {
+ resetOpenAIAdvancedSchedulerSettingCacheForTest()
+
+ ctx := context.Background()
+ groupID := int64(10112)
+ accounts := []Account{
+ {
+ ID: 37021,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Credentials: map[string]any{
+ "openai_capabilities": []any{"chat_completions"},
+ },
+ Extra: map[string]any{
+ "openai_apikey_responses_websockets_v2_enabled": true,
+ },
+ },
+ {
+ ID: 37022,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 5,
+ Credentials: map[string]any{
+ "openai_capabilities": []any{"chat_completions", "embeddings"},
+ },
+ Extra: map[string]any{
+ "openai_apikey_responses_websockets_v2_enabled": true,
+ },
+ },
+ }
+ cfg := newSchedulerTestOpenAIWSV2Config()
+ cfg.Gateway.Scheduling.LoadBatchEnabled = false
+ cache := &schedulerTestGatewayCache{
+ sessionBindings: map[string]int64{
+ "openai:session_hash_embeddings": 37021,
+ },
+ }
+ svc := &OpenAIGatewayService{
+ accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
+ cache: cache,
+ cfg: cfg,
+ rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
+ concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{}),
+ }
+ store := svc.getOpenAIWSStateStore()
+ require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_embeddings_chat_only", 37021, time.Hour))
+
+ selection, decision, err := svc.SelectAccountWithSchedulerForCapability(
+ ctx,
+ &groupID,
+ "resp_embeddings_chat_only",
+ "session_hash_embeddings",
+ "text-embedding-3-small",
+ nil,
+ OpenAIUpstreamTransportHTTPSSE,
+ OpenAIEndpointCapabilityEmbeddings,
+ false,
+ )
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(37022), selection.Account.ID)
+ require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
+ require.False(t, decision.StickyPreviousHit)
+ require.False(t, decision.StickySessionHit)
+ require.Equal(t, int64(37022), cache.sessionBindings["openai:session_hash_embeddings"])
+}
+
func TestOpenAIGatewayService_OpenAIAccountSchedulerMetrics_DisabledNoOp(t *testing.T) {
resetOpenAIAdvancedSchedulerSettingCacheForTest()
@@ -498,6 +691,287 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyRateLimite
require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
}
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_AutoPauseBy5hThreshold(t *testing.T) {
+ ctx := context.Background()
+ primary := Account{
+ ID: 35001,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 95.0,
+ "auto_pause_5h_threshold": 0.95,
+ },
+ }
+ secondary := Account{ID: 35002, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35002), account.ID)
+}
+
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_AllowsBelow5hThreshold(t *testing.T) {
+ ctx := context.Background()
+ primary := Account{
+ ID: 35101,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 80.0,
+ "auto_pause_5h_threshold": 0.95,
+ },
+ }
+ secondary := Account{ID: 35102, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35101), account.ID)
+}
+
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_AutoPauseBy7dThreshold(t *testing.T) {
+ ctx := context.Background()
+ primary := Account{
+ ID: 35201,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_7d_used_percent": 95.0,
+ "auto_pause_7d_threshold": 0.95,
+ },
+ }
+ secondary := Account{ID: 35202, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35202), account.ID)
+}
+
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_UnconfiguredThresholdKeepsLegacyBehavior(t *testing.T) {
+ ctx := context.Background()
+ primary := Account{ID: 35301, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, Extra: map[string]any{"codex_5h_used_percent": 99.0, "codex_7d_used_percent": 99.0}}
+ secondary := Account{ID: 35302, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35301), account.ID)
+}
+
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_UsesGlobalDefaultThreshold(t *testing.T) {
+ ctx := withOpenAIQuotaAutoPauseSettings(context.Background(), OpsOpenAIAccountQuotaAutoPauseSettings{DefaultThreshold5h: 0.95})
+ primary := Account{
+ ID: 35401,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 95.0,
+ },
+ }
+ secondary := Account{ID: 35402, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35402), account.ID)
+}
+
+// Regression: a per-account explicit-disable flag exempts the account from auto-pause
+// even when a global default threshold is set. Without this, "leave threshold blank"
+// silently falls back to global default and admins have no way to whitelist a single
+// account.
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_PerAccountDisableOverridesGlobalDefault(t *testing.T) {
+ ctx := withOpenAIQuotaAutoPauseSettings(context.Background(), OpsOpenAIAccountQuotaAutoPauseSettings{DefaultThreshold5h: 0.95})
+ // Account has high usage AND no per-account threshold (would normally fall back to
+ // the global default and get paused), but the explicit disable flag is set.
+ primary := Account{
+ ID: 35701,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 99.0,
+ "auto_pause_5h_disabled": true,
+ },
+ }
+ secondary := Account{ID: 35702, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35701), account.ID)
+}
+
+// Disable is per-window: disabling only 5h must still allow 7d auto-pause to fire.
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_PerWindowDisableScoped(t *testing.T) {
+ ctx := context.Background()
+ primary := Account{
+ ID: 35801,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 99.0,
+ "codex_7d_used_percent": 99.0,
+ "auto_pause_5h_disabled": true,
+ "auto_pause_7d_threshold": 0.95,
+ },
+ }
+ secondary := Account{ID: 35802, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35802), account.ID, "7d auto-pause must still fire even though 5h is disabled")
+}
+
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_StaleUsageWindowResetSkipsPause(t *testing.T) {
+ ctx := context.Background()
+ // Usage is over threshold but the window's reset time has already passed, so the
+ // cached percentage is stale (the real window rolled over) and the account must NOT
+ // stay paused — otherwise it could be skipped forever with no traffic to refresh it.
+ primary := Account{
+ ID: 35501,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 99.0,
+ "auto_pause_5h_threshold": 0.95,
+ "codex_5h_reset_at": time.Now().Add(-time.Minute).Format(time.RFC3339),
+ },
+ }
+ secondary := Account{ID: 35502, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35501), account.ID)
+}
+
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_FreshUsageWindowStillPauses(t *testing.T) {
+ ctx := context.Background()
+ // Same as above but the window has not reset yet, so the account stays paused.
+ primary := Account{
+ ID: 35601,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 99.0,
+ "auto_pause_5h_threshold": 0.95,
+ "codex_5h_reset_at": time.Now().Add(time.Hour).Format(time.RFC3339),
+ },
+ }
+ secondary := Account{ID: 35602, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35602), account.ID)
+}
+
+// Issue #2994: an account poisoned with an inflated used% (e.g. from the reverted #2918
+// inversion) gets excluded from scheduling, and a paused account never receives traffic to
+// refresh its snapshot. When the snapshot is stale (codex_usage_updated_at older than the
+// staleness bound) the account must be allowed a request so it can self-heal from the real
+// response headers — independent of the window's reset time.
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_StaleUsageSnapshotSkipsPause_Issue2994(t *testing.T) {
+ ctx := context.Background()
+ primary := Account{
+ ID: 35701,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 99.0,
+ "auto_pause_5h_threshold": 0.95,
+ // Window has NOT reset yet, so the reset guard stays inactive.
+ "codex_5h_reset_at": time.Now().Add(time.Hour).Format(time.RFC3339),
+ // Snapshot is stale: older than openAICodexAutoPauseStaleAfter (2h).
+ "codex_usage_updated_at": time.Now().Add(-3 * time.Hour).Format(time.RFC3339),
+ },
+ }
+ secondary := Account{ID: 35702, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35701), account.ID)
+}
+
+// Issue #2994 guardrail: a genuinely-exhausted account whose snapshot was refreshed recently
+// (codex_usage_updated_at fresh) must STILL be auto-paused. The stale self-heal must not let a
+// real 99%-used account escape pause.
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_FreshExhaustedSnapshotStillPauses_Issue2994(t *testing.T) {
+ ctx := context.Background()
+ primary := Account{
+ ID: 35801,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 99.0,
+ "auto_pause_5h_threshold": 0.95,
+ "codex_5h_reset_at": time.Now().Add(time.Hour).Format(time.RFC3339),
+ // Snapshot refreshed 1 minute ago: not stale, so the account stays paused.
+ "codex_usage_updated_at": time.Now().Add(-time.Minute).Format(time.RFC3339),
+ },
+ }
+ secondary := Account{ID: 35802, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5}
+ svc := &OpenAIGatewayService{accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}}, cfg: &config.Config{}}
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.1", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(35802), account.ID)
+}
+
func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_SkipsFreshlyRateLimitedSnapshotCandidate(t *testing.T) {
ctx := context.Background()
groupID := int64(10102)
@@ -521,6 +995,50 @@ func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_SkipsFreshlyRa
require.Equal(t, int64(32002), account.ID)
}
+func TestOpenAIGatewayService_SelectAccountForModelWithExclusions_ModelRateLimitOnlySkipsThatModel(t *testing.T) {
+ ctx := context.Background()
+ resetAt := time.Now().Add(30 * time.Minute).Format(time.RFC3339)
+ primary := Account{
+ ID: 32101,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ modelRateLimitsKey: map[string]any{
+ "gpt-5.4": map[string]any{
+ "rate_limit_reset_at": resetAt,
+ },
+ },
+ },
+ }
+ secondary := Account{
+ ID: 32102,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 5,
+ }
+ svc := &OpenAIGatewayService{
+ accountRepo: schedulerTestOpenAIAccountRepo{accounts: []Account{primary, secondary}},
+ cfg: &config.Config{},
+ }
+
+ account, err := svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.4", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(32102), account.ID)
+
+ account, err = svc.SelectAccountForModelWithExclusions(ctx, nil, "", "gpt-5.3", nil)
+ require.NoError(t, err)
+ require.NotNil(t, account)
+ require.Equal(t, int64(32101), account.ID)
+}
+
func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyDBRuntimeRecheckSkipsStaleCachedAccount(t *testing.T) {
ctx := context.Background()
groupID := int64(10103)
@@ -711,6 +1229,9 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyBusyKeepsS
cfg := &config.Config{}
cfg.Gateway.Scheduling.StickySessionMaxWaiting = 2
cfg.Gateway.Scheduling.StickySessionWaitTimeout = 45 * time.Second
+ cfg.Gateway.OpenAIScheduler.StickyEscapeEnabled = false
+ cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs = 15000
+ cfg.Gateway.OpenAIScheduler.StickyEscapeErrorRate = 0.5
cfg.Gateway.OpenAIWS.Enabled = true
cfg.Gateway.OpenAIWS.APIKeyEnabled = true
cfg.Gateway.OpenAIWS.OAuthEnabled = true
@@ -759,6 +1280,253 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyBusyKeepsS
require.True(t, decision.StickySessionHit)
}
+func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyEscapeByTTFT(t *testing.T) {
+ ctx := context.Background()
+ groupID := int64(10101)
+ accounts := []Account{
+ {
+ ID: 21101,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ },
+ {
+ ID: 21102,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 1,
+ },
+ }
+ cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{"openai:session_hash_sticky_ttft": 21101}}
+ cfg := &config.Config{}
+ cfg.Gateway.OpenAIScheduler.StickyEscapeEnabled = true
+ cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs = 15000
+ cfg.Gateway.OpenAIScheduler.StickyEscapeErrorRate = 0.5
+ concurrencyCache := schedulerTestConcurrencyCache{acquireResults: map[int64]bool{21102: true}}
+ svc := &OpenAIGatewayService{
+ accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
+ cache: cache,
+ cfg: cfg,
+ rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
+ concurrencyService: NewConcurrencyService(concurrencyCache),
+ openaiAccountStats: newOpenAIAccountRuntimeStats(),
+ }
+ fastTTFT := 14999
+ svc.openaiAccountStats.report(21101, true, &fastTTFT)
+ stableTTFT := 14999
+ svc.openaiAccountStats.report(21101, true, &stableTTFT)
+
+ selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_hash_sticky_ttft", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(21101), selection.Account.ID)
+ require.Equal(t, openAIAccountScheduleLayerSessionSticky, decision.Layer)
+ require.True(t, decision.StickySessionHit)
+ if selection.ReleaseFunc != nil {
+ selection.ReleaseFunc()
+ }
+
+ slowTTFT := 20000
+ for i := 0; i < 3; i++ {
+ svc.openaiAccountStats.report(21101, true, &slowTTFT)
+ }
+
+ selection, decision, err = svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_hash_sticky_ttft", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(21102), selection.Account.ID)
+ require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
+ require.False(t, decision.StickySessionHit)
+ require.Equal(t, int64(21101), cache.sessionBindings["openai:session_hash_sticky_ttft"])
+ if selection.ReleaseFunc != nil {
+ selection.ReleaseFunc()
+ }
+}
+
+func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyEscapeByErrorRate(t *testing.T) {
+ ctx := context.Background()
+ groupID := int64(10102)
+ accounts := []Account{
+ {ID: 21201, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0},
+ {ID: 21202, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 1},
+ }
+ cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{"openai:session_hash_sticky_error_rate": 21201}}
+ cfg := &config.Config{}
+ cfg.Gateway.OpenAIScheduler.StickyEscapeEnabled = true
+ cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs = 15000
+ cfg.Gateway.OpenAIScheduler.StickyEscapeErrorRate = 0.5
+ svc := &OpenAIGatewayService{
+ accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
+ cache: cache,
+ cfg: cfg,
+ rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
+ concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{acquireResults: map[int64]bool{21202: true}}),
+ openaiAccountStats: newOpenAIAccountRuntimeStats(),
+ }
+ for i := 0; i < 3; i++ {
+ svc.openaiAccountStats.report(21201, false, nil)
+ }
+ selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_hash_sticky_error_rate", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(21201), selection.Account.ID)
+ require.Equal(t, openAIAccountScheduleLayerSessionSticky, decision.Layer)
+ require.True(t, decision.StickySessionHit)
+ if selection.ReleaseFunc != nil {
+ selection.ReleaseFunc()
+ }
+ for i := 0; i < 2; i++ {
+ svc.openaiAccountStats.report(21201, false, nil)
+ }
+
+ selection, decision, err = svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_hash_sticky_error_rate", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(21202), selection.Account.ID)
+ require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
+ require.False(t, decision.StickySessionHit)
+ require.Equal(t, int64(21201), cache.sessionBindings["openai:session_hash_sticky_error_rate"])
+ if selection.ReleaseFunc != nil {
+ selection.ReleaseFunc()
+ }
+}
+
+func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyBusyEscapes(t *testing.T) {
+ ctx := context.Background()
+ groupID := int64(10103)
+ accounts := []Account{
+ {ID: 21301, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0},
+ {ID: 21302, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 1},
+ }
+ cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{"openai:session_hash_sticky_busy_escape": 21301}}
+ cfg := &config.Config{}
+ cfg.Gateway.OpenAIScheduler.StickyEscapeEnabled = true
+ cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs = 15000
+ cfg.Gateway.OpenAIScheduler.StickyEscapeErrorRate = 0.5
+ cfg.Gateway.Scheduling.StickySessionMaxWaiting = 2
+ cfg.Gateway.Scheduling.StickySessionWaitTimeout = 45 * time.Second
+ concurrencyCache := schedulerTestConcurrencyCache{
+ acquireResults: map[int64]bool{21301: false, 21302: true},
+ waitCounts: map[int64]int{21301: 999},
+ loadMap: map[int64]*AccountLoadInfo{
+ 21301: {AccountID: 21301, LoadRate: 95, WaitingCount: 9},
+ 21302: {AccountID: 21302, LoadRate: 1, WaitingCount: 0},
+ },
+ }
+ svc := &OpenAIGatewayService{
+ accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
+ cache: cache,
+ cfg: cfg,
+ rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
+ concurrencyService: NewConcurrencyService(concurrencyCache),
+ }
+
+ selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_hash_sticky_busy_escape", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(21302), selection.Account.ID)
+ require.Nil(t, selection.WaitPlan)
+ require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
+ require.False(t, decision.StickySessionHit)
+ if selection.ReleaseFunc != nil {
+ selection.ReleaseFunc()
+ }
+}
+
+func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionStickyEscapeDisabledKeepsLegacyBehavior(t *testing.T) {
+ ctx := context.Background()
+ groupID := int64(10104)
+ accounts := []Account{
+ {ID: 21401, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0},
+ {ID: 21402, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 1},
+ }
+ cache := &schedulerTestGatewayCache{sessionBindings: map[string]int64{"openai:session_hash_sticky_disabled": 21401}}
+ cfg := &config.Config{}
+ cfg.Gateway.OpenAIScheduler.StickyEscapeEnabled = false
+ cfg.Gateway.OpenAIScheduler.StickyEscapeTTFTMs = 15000
+ cfg.Gateway.OpenAIScheduler.StickyEscapeErrorRate = 0.5
+ cfg.Gateway.Scheduling.StickySessionMaxWaiting = 2
+ cfg.Gateway.Scheduling.StickySessionWaitTimeout = 45 * time.Second
+ concurrencyCache := schedulerTestConcurrencyCache{
+ acquireResults: map[int64]bool{21401: false, 21402: true},
+ waitCounts: map[int64]int{21401: 999},
+ }
+ svc := &OpenAIGatewayService{
+ accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
+ cache: cache,
+ cfg: cfg,
+ rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
+ concurrencyService: NewConcurrencyService(concurrencyCache),
+ openaiAccountStats: newOpenAIAccountRuntimeStats(),
+ }
+ slowTTFT := 20000
+ svc.openaiAccountStats.report(21401, true, &slowTTFT)
+ for i := 0; i < 5; i++ {
+ svc.openaiAccountStats.report(21401, false, nil)
+ }
+
+ selection, decision, err := svc.SelectAccountWithScheduler(ctx, &groupID, "", "session_hash_sticky_disabled", "gpt-5.1", nil, OpenAIUpstreamTransportAny, false)
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(21401), selection.Account.ID)
+ require.NotNil(t, selection.WaitPlan)
+ require.Equal(t, int64(21401), selection.WaitPlan.AccountID)
+ require.Equal(t, openAIAccountScheduleLayerSessionSticky, decision.Layer)
+ require.True(t, decision.StickySessionHit)
+}
+
+func TestDefaultOpenAIAccountScheduler_ShouldEscapeStickyAccount_ThresholdBoundary(t *testing.T) {
+ stats := newOpenAIAccountRuntimeStats()
+ accountID := int64(21501)
+ ttft := 15000
+ stats.report(accountID, true, &ttft)
+ stats.report(accountID, false, nil)
+ stats.report(accountID, true, nil)
+ scheduler := &defaultOpenAIAccountScheduler{stats: stats}
+
+ reason, errorRate, observedTTFT, shouldEscape := scheduler.shouldEscapeStickyAccount(accountID, openAIStickyEscapeConfig{
+ enabled: true,
+ ttftMs: 15000,
+ errorRate: 0.5,
+ })
+ require.False(t, shouldEscape)
+ require.Empty(t, reason)
+ require.InDelta(t, 0.16, errorRate, 1e-9)
+ require.InDelta(t, 15000, observedTTFT, 1e-9)
+
+ for i := 0; i < 4; i++ {
+ stats.report(accountID, false, nil)
+ }
+ reason, errorRate, _, shouldEscape = scheduler.shouldEscapeStickyAccount(accountID, openAIStickyEscapeConfig{
+ enabled: true,
+ ttftMs: 15000,
+ errorRate: 1,
+ })
+ require.False(t, shouldEscape)
+ require.Empty(t, reason)
+ reason, errorRate, observedTTFT, shouldEscape = scheduler.shouldEscapeStickyAccount(accountID, openAIStickyEscapeConfig{
+ enabled: true,
+ ttftMs: 15000,
+ errorRate: errorRate,
+ })
+ require.False(t, shouldEscape)
+ require.Empty(t, reason)
+ require.InDelta(t, 0.655936, errorRate, 1e-9)
+ require.InDelta(t, 15000, observedTTFT, 1e-9)
+}
+
func TestOpenAIGatewayService_SelectAccountWithScheduler_SessionSticky_ForceHTTP(t *testing.T) {
ctx := context.Background()
groupID := int64(1010)
@@ -1001,6 +1769,85 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_LoadBalanceTopKFallback
}
}
+// Regression: TopK initial filter must drop quota-auto-paused accounts. Otherwise
+// the candidate pool is filled with paused accounts, healthy accounts fall outside
+// TopK, and the scheduler returns "no available accounts" even though healthy ones
+// exist.
+func TestOpenAIGatewayService_SelectAccountWithScheduler_LoadBalanceTopKExcludesQuotaPaused(t *testing.T) {
+ ctx := context.Background()
+ groupID := int64(110)
+ accounts := []Account{
+ {
+ ID: 37001,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 0,
+ Extra: map[string]any{
+ "codex_5h_used_percent": 96.0,
+ "auto_pause_5h_threshold": 0.95,
+ },
+ },
+ {
+ ID: 37002,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 5,
+ },
+ }
+
+ cfg := &config.Config{}
+ cfg.Gateway.OpenAIWS.LBTopK = 1 // TopK=1 makes the bug fatal: paused account would crowd out the healthy one entirely
+ cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Priority = 0.4
+ cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Load = 1.0
+ cfg.Gateway.OpenAIWS.SchedulerScoreWeights.Queue = 1.0
+
+ concurrencyCache := schedulerTestConcurrencyCache{
+ loadMap: map[int64]*AccountLoadInfo{
+ 37001: {AccountID: 37001, LoadRate: 5, WaitingCount: 0},
+ 37002: {AccountID: 37002, LoadRate: 5, WaitingCount: 0},
+ },
+ acquireResults: map[int64]bool{
+ 37002: true,
+ },
+ }
+
+ svc := &OpenAIGatewayService{
+ accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts},
+ cache: &schedulerTestGatewayCache{},
+ cfg: cfg,
+ rateLimitService: newOpenAIAdvancedSchedulerRateLimitService("true"),
+ concurrencyService: NewConcurrencyService(concurrencyCache),
+ }
+
+ selection, decision, err := svc.SelectAccountWithScheduler(
+ ctx,
+ &groupID,
+ "",
+ "",
+ "gpt-5.1",
+ nil,
+ OpenAIUpstreamTransportAny,
+ false,
+ )
+ require.NoError(t, err)
+ require.NotNil(t, selection)
+ require.NotNil(t, selection.Account)
+ require.Equal(t, int64(37002), selection.Account.ID)
+ require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer)
+ // Only the healthy account should ever enter the candidate pool; the paused one
+ // must be filtered out at the initial-filter stage.
+ require.Equal(t, 1, decision.CandidateCount)
+ if selection.ReleaseFunc != nil {
+ selection.ReleaseFunc()
+ }
+}
+
func TestOpenAIGatewayService_OpenAIAccountSchedulerMetrics(t *testing.T) {
ctx := context.Background()
groupID := int64(12)
diff --git a/backend/internal/service/openai_client_restriction_detector.go b/backend/internal/service/openai_client_restriction_detector.go
index d1784e11..8589737a 100644
--- a/backend/internal/service/openai_client_restriction_detector.go
+++ b/backend/internal/service/openai_client_restriction_detector.go
@@ -13,6 +13,10 @@ const (
CodexClientRestrictionReasonMatchedUA = "official_client_user_agent_matched"
// CodexClientRestrictionReasonMatchedOriginator 表示请求命中官方客户端 originator 白名单。
CodexClientRestrictionReasonMatchedOriginator = "official_client_originator_matched"
+ // CodexClientRestrictionReasonMatchedAllowedClient 表示请求命中账号级额外放行的命名客户端预设。
+ CodexClientRestrictionReasonMatchedAllowedClient = "allowed_client_matched"
+ // CodexClientRestrictionReasonMatchedGlobalAllowedClient 表示请求命中全局额外放行的命名客户端预设。
+ CodexClientRestrictionReasonMatchedGlobalAllowedClient = "global_allowed_client_matched"
// CodexClientRestrictionReasonNotMatchedUA 表示请求未命中官方客户端 UA 白名单。
CodexClientRestrictionReasonNotMatchedUA = "official_client_user_agent_not_matched"
// CodexClientRestrictionReasonForceCodexCLI 表示通过 ForceCodexCLI 配置兜底放行。
@@ -28,7 +32,7 @@ type CodexClientRestrictionDetectionResult struct {
// CodexClientRestrictionDetector 定义 codex_cli_only 统一检测入口。
type CodexClientRestrictionDetector interface {
- Detect(c *gin.Context, account *Account) CodexClientRestrictionDetectionResult
+ Detect(c *gin.Context, account *Account, globalAllowedClients []string) CodexClientRestrictionDetectionResult
}
// OpenAICodexClientRestrictionDetector 为 OpenAI OAuth codex_cli_only 的默认实现。
@@ -40,7 +44,7 @@ func NewOpenAICodexClientRestrictionDetector(cfg *config.Config) *OpenAICodexCli
return &OpenAICodexClientRestrictionDetector{cfg: cfg}
}
-func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *Account) CodexClientRestrictionDetectionResult {
+func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *Account, globalAllowedClients []string) CodexClientRestrictionDetectionResult {
if account == nil || !account.IsCodexCLIOnlyEnabled() {
return CodexClientRestrictionDetectionResult{
Enabled: false,
@@ -78,6 +82,26 @@ func (d *OpenAICodexClientRestrictionDetector) Detect(c *gin.Context, account *A
}
}
+ // 官方客户端白名单未命中时,先尝试账号级额外放行的命名客户端预设(如 Claude Code codex 插件)。
+ if allowed := account.GetCodexCLIOnlyAllowedClients(); len(allowed) > 0 &&
+ openai.MatchAllowedClients(userAgent, originator, allowed) {
+ return CodexClientRestrictionDetectionResult{
+ Enabled: true,
+ Matched: true,
+ Reason: CodexClientRestrictionReasonMatchedAllowedClient,
+ }
+ }
+
+ // 再尝试由更高作用域(全局设置)注入的额外放行客户端列表。
+ if len(globalAllowedClients) > 0 &&
+ openai.MatchAllowedClients(userAgent, originator, globalAllowedClients) {
+ return CodexClientRestrictionDetectionResult{
+ Enabled: true,
+ Matched: true,
+ Reason: CodexClientRestrictionReasonMatchedGlobalAllowedClient,
+ }
+ }
+
return CodexClientRestrictionDetectionResult{
Enabled: true,
Matched: false,
diff --git a/backend/internal/service/openai_client_restriction_detector_test.go b/backend/internal/service/openai_client_restriction_detector_test.go
index 984b4ff6..fc115128 100644
--- a/backend/internal/service/openai_client_restriction_detector_test.go
+++ b/backend/internal/service/openai_client_restriction_detector_test.go
@@ -30,7 +30,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) {
detector := NewOpenAICodexClientRestrictionDetector(nil)
account := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{}}
- result := detector.Detect(newCodexDetectorTestContext("curl/8.0", ""), account)
+ result := detector.Detect(newCodexDetectorTestContext("curl/8.0", ""), account, nil)
require.False(t, result.Enabled)
require.False(t, result.Matched)
require.Equal(t, CodexClientRestrictionReasonDisabled, result.Reason)
@@ -44,7 +44,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) {
Extra: map[string]any{"codex_cli_only": true},
}
- result := detector.Detect(newCodexDetectorTestContext("codex_cli_rs/0.99.0", ""), account)
+ result := detector.Detect(newCodexDetectorTestContext("codex_cli_rs/0.99.0", ""), account, nil)
require.True(t, result.Enabled)
require.True(t, result.Matched)
require.Equal(t, CodexClientRestrictionReasonMatchedUA, result.Reason)
@@ -58,7 +58,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) {
Extra: map[string]any{"codex_cli_only": true},
}
- result := detector.Detect(newCodexDetectorTestContext("codex_vscode/1.0.0", ""), account)
+ result := detector.Detect(newCodexDetectorTestContext("codex_vscode/1.0.0", ""), account, nil)
require.True(t, result.Enabled)
require.True(t, result.Matched)
require.Equal(t, CodexClientRestrictionReasonMatchedUA, result.Reason)
@@ -72,7 +72,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) {
Extra: map[string]any{"codex_cli_only": true},
}
- result := detector.Detect(newCodexDetectorTestContext("codex_app/2.1.0", ""), account)
+ result := detector.Detect(newCodexDetectorTestContext("codex_app/2.1.0", ""), account, nil)
require.True(t, result.Enabled)
require.True(t, result.Matched)
require.Equal(t, CodexClientRestrictionReasonMatchedUA, result.Reason)
@@ -86,7 +86,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) {
Extra: map[string]any{"codex_cli_only": true},
}
- result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "codex_chatgpt_desktop"), account)
+ result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "codex_chatgpt_desktop"), account, nil)
require.True(t, result.Enabled)
require.True(t, result.Matched)
require.Equal(t, CodexClientRestrictionReasonMatchedOriginator, result.Reason)
@@ -100,7 +100,7 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) {
Extra: map[string]any{"codex_cli_only": true},
}
- result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account)
+ result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account, nil)
require.True(t, result.Enabled)
require.False(t, result.Matched)
require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason)
@@ -116,9 +116,146 @@ func TestOpenAICodexClientRestrictionDetector_Detect(t *testing.T) {
Extra: map[string]any{"codex_cli_only": true},
}
- result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account)
+ result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account, nil)
require.True(t, result.Enabled)
require.True(t, result.Matched)
require.Equal(t, CodexClientRestrictionReasonForceCodexCLI, result.Reason)
})
}
+
+func TestOpenAICodexClientRestrictionDetector_Detect_AllowedClients(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ const (
+ claudeCodeUA = "Claude Code/0.5.0 (Macos 15.5; arm64) iTerm2.app (Claude Code; 1.0.4)"
+ claudeCodeOriginator = "Claude Code"
+ )
+
+ t.Run("配置 claude_code 白名单且命中真实签名时放行", func(t *testing.T) {
+ detector := NewOpenAICodexClientRestrictionDetector(nil)
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{
+ "codex_cli_only": true,
+ "codex_cli_only_allowed_clients": []any{"claude_code"},
+ },
+ }
+
+ result := detector.Detect(newCodexDetectorTestContext(claudeCodeUA, claudeCodeOriginator), account, nil)
+ require.True(t, result.Enabled)
+ require.True(t, result.Matched)
+ require.Equal(t, CodexClientRestrictionReasonMatchedAllowedClient, result.Reason)
+ })
+
+ t.Run("配置白名单但伪造 originator 仍拒绝", func(t *testing.T) {
+ detector := NewOpenAICodexClientRestrictionDetector(nil)
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{
+ "codex_cli_only": true,
+ "codex_cli_only_allowed_clients": []any{"claude_code"},
+ },
+ }
+
+ result := detector.Detect(newCodexDetectorTestContext(claudeCodeUA, "my_client"), account, nil)
+ require.True(t, result.Enabled)
+ require.False(t, result.Matched)
+ require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason)
+ })
+
+ t.Run("未配置白名单时 Claude Code 签名仍拒绝", func(t *testing.T) {
+ detector := NewOpenAICodexClientRestrictionDetector(nil)
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{"codex_cli_only": true},
+ }
+
+ result := detector.Detect(newCodexDetectorTestContext(claudeCodeUA, claudeCodeOriginator), account, nil)
+ require.True(t, result.Enabled)
+ require.False(t, result.Matched)
+ require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason)
+ })
+
+ t.Run("未开启 codex_cli_only 时白名单不参与,直接绕过", func(t *testing.T) {
+ detector := NewOpenAICodexClientRestrictionDetector(nil)
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{"codex_cli_only_allowed_clients": []any{"claude_code"}},
+ }
+
+ result := detector.Detect(newCodexDetectorTestContext(claudeCodeUA, claudeCodeOriginator), account, nil)
+ require.False(t, result.Enabled)
+ require.False(t, result.Matched)
+ require.Equal(t, CodexClientRestrictionReasonDisabled, result.Reason)
+ })
+
+ t.Run("全局列表含 claude_code + 命中签名 → 放行(global)", func(t *testing.T) {
+ detector := NewOpenAICodexClientRestrictionDetector(nil)
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{"codex_cli_only": true},
+ }
+ result := detector.Detect(
+ newCodexDetectorTestContext("Claude Code/0.5.0 (Macos 15.5; arm64) iTerm2.app (Claude Code; 1.0.4)", "Claude Code"),
+ account,
+ []string{"claude_code"},
+ )
+ require.True(t, result.Enabled)
+ require.True(t, result.Matched)
+ require.Equal(t, CodexClientRestrictionReasonMatchedGlobalAllowedClient, result.Reason)
+ })
+
+ t.Run("全局列表含 claude_code + 非签名 → 403", func(t *testing.T) {
+ detector := NewOpenAICodexClientRestrictionDetector(nil)
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{"codex_cli_only": true},
+ }
+ result := detector.Detect(newCodexDetectorTestContext("curl/8.0", "my_client"), account, []string{"claude_code"})
+ require.True(t, result.Enabled)
+ require.False(t, result.Matched)
+ require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason)
+ })
+
+ t.Run("全局列表为空 + 账号未配 → 403", func(t *testing.T) {
+ detector := NewOpenAICodexClientRestrictionDetector(nil)
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{"codex_cli_only": true},
+ }
+ result := detector.Detect(
+ newCodexDetectorTestContext("Claude Code/0.5.0 (Macos) (Claude Code; 1.0.4)", "Claude Code"),
+ account,
+ nil,
+ )
+ require.True(t, result.Enabled)
+ require.False(t, result.Matched)
+ require.Equal(t, CodexClientRestrictionReasonNotMatchedUA, result.Reason)
+ })
+
+ t.Run("账号白名单优先于全局列表(reason=account)", func(t *testing.T) {
+ detector := NewOpenAICodexClientRestrictionDetector(nil)
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Extra: map[string]any{
+ "codex_cli_only": true,
+ "codex_cli_only_allowed_clients": []any{"claude_code"},
+ },
+ }
+ result := detector.Detect(
+ newCodexDetectorTestContext("Claude Code/0.5.0 (Macos) (Claude Code; 1.0.4)", "Claude Code"),
+ account,
+ []string{"claude_code"},
+ )
+ require.True(t, result.Matched)
+ require.Equal(t, CodexClientRestrictionReasonMatchedAllowedClient, result.Reason)
+ })
+}
diff --git a/backend/internal/service/openai_embeddings.go b/backend/internal/service/openai_embeddings.go
new file mode 100644
index 00000000..7c710259
--- /dev/null
+++ b/backend/internal/service/openai_embeddings.go
@@ -0,0 +1,240 @@
+package service
+
+import (
+ "bytes"
+ "context"
+ "errors"
+ "fmt"
+ "io"
+ "net/http"
+ "strings"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/pkg/logger"
+ "github.com/Wei-Shaw/sub2api/internal/util/responseheaders"
+ "github.com/gin-gonic/gin"
+ "github.com/tidwall/gjson"
+ "go.uber.org/zap"
+)
+
+func (s *OpenAIGatewayService) ForwardEmbeddings(
+ ctx context.Context,
+ c *gin.Context,
+ account *Account,
+ body []byte,
+ defaultMappedModel string,
+) (*OpenAIForwardResult, error) {
+ startTime := time.Now()
+
+ originalModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
+ if originalModel == "" {
+ writeOpenAIEmbeddingsError(c, http.StatusBadRequest, "invalid_request_error", "model is required")
+ return nil, fmt.Errorf("missing model in request")
+ }
+
+ billingModel := resolveOpenAIForwardModel(account, originalModel, defaultMappedModel)
+ upstreamModel := normalizeOpenAIModelForUpstream(account, billingModel)
+ upstreamBody := body
+ if upstreamModel != originalModel {
+ upstreamBody = ReplaceModelInBody(body, upstreamModel)
+ }
+
+ logger.L().Debug("openai embeddings: forwarding",
+ zap.Int64("account_id", account.ID),
+ zap.String("original_model", originalModel),
+ zap.String("billing_model", billingModel),
+ zap.String("upstream_model", upstreamModel),
+ )
+
+ apiKey := account.GetOpenAIApiKey()
+ if apiKey == "" {
+ return nil, fmt.Errorf("account %d missing api_key", account.ID)
+ }
+ baseURL := account.GetOpenAIBaseURL()
+ if baseURL == "" {
+ baseURL = "https://api.openai.com"
+ }
+ validatedURL, err := s.validateUpstreamBaseURL(baseURL)
+ if err != nil {
+ return nil, fmt.Errorf("invalid base_url: %w", err)
+ }
+ targetURL := buildOpenAIEmbeddingsURL(validatedURL)
+
+ upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
+ upstreamReq, err := http.NewRequestWithContext(upstreamCtx, http.MethodPost, targetURL, bytes.NewReader(upstreamBody))
+ releaseUpstreamCtx()
+ if err != nil {
+ return nil, fmt.Errorf("build upstream request: %w", err)
+ }
+ upstreamReq = upstreamReq.WithContext(WithHTTPUpstreamProfile(upstreamReq.Context(), HTTPUpstreamProfileOpenAI))
+ upstreamReq.Header.Set("Content-Type", "application/json")
+ upstreamReq.Header.Set("Authorization", "Bearer "+apiKey)
+ upstreamReq.Header.Set("Accept", "application/json")
+ for key, values := range c.Request.Header {
+ lowerKey := strings.ToLower(key)
+ if openaiCCRawAllowedHeaders[lowerKey] {
+ for _, v := range values {
+ upstreamReq.Header.Add(key, v)
+ }
+ }
+ }
+ if customUA := account.GetOpenAIUserAgent(); customUA != "" {
+ upstreamReq.Header.Set("user-agent", customUA)
+ }
+
+ proxyURL := ""
+ if account.Proxy != nil {
+ proxyURL = account.Proxy.URL()
+ }
+ resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency)
+ if err != nil {
+ safeErr := sanitizeUpstreamErrorMessage(err.Error())
+ setOpsUpstreamError(c, 0, safeErr, "")
+ appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
+ Platform: account.Platform,
+ AccountID: account.ID,
+ AccountName: account.Name,
+ UpstreamStatusCode: 0,
+ Kind: "request_error",
+ Message: safeErr,
+ })
+ writeOpenAIEmbeddingsError(c, http.StatusBadGateway, "upstream_error", "Upstream request failed")
+ return nil, fmt.Errorf("upstream request failed: %s", safeErr)
+ }
+ defer func() { _ = resp.Body.Close() }()
+
+ if resp.StatusCode >= 400 {
+ respBody := s.readUpstreamErrorBody(resp)
+ _ = resp.Body.Close()
+ resp.Body = io.NopCloser(bytes.NewReader(respBody))
+
+ upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody))
+ upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
+ if s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody) {
+ upstreamDetail := ""
+ if s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody {
+ maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes
+ if maxBytes <= 0 {
+ maxBytes = 2048
+ }
+ upstreamDetail = truncateString(string(respBody), maxBytes)
+ }
+ appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
+ Platform: account.Platform,
+ AccountID: account.ID,
+ AccountName: account.Name,
+ UpstreamStatusCode: resp.StatusCode,
+ UpstreamRequestID: resp.Header.Get("x-request-id"),
+ Kind: "failover",
+ Message: upstreamMsg,
+ Detail: upstreamDetail,
+ })
+ s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel)
+ return nil, &UpstreamFailoverError{
+ StatusCode: resp.StatusCode,
+ ResponseBody: respBody,
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
+ }
+ }
+ writeOpenAIEmbeddingsUpstreamResponse(c, resp, respBody, s.responseHeaderFilter)
+ return nil, fmt.Errorf("upstream returned status %d", resp.StatusCode)
+ }
+
+ respBody, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError)
+ if err != nil {
+ if !errors.Is(err, ErrUpstreamResponseBodyTooLarge) {
+ writeOpenAIEmbeddingsError(c, http.StatusBadGateway, "api_error", "Failed to read upstream response")
+ }
+ return nil, fmt.Errorf("read upstream body: %w", err)
+ }
+
+ writeOpenAIEmbeddingsUpstreamResponse(c, resp, respBody, s.responseHeaderFilter)
+
+ return &OpenAIForwardResult{
+ RequestID: firstNonEmptyString(resp.Header.Get("x-request-id"), resp.Header.Get("request-id")),
+ Usage: extractOpenAIEmbeddingsUsage(respBody),
+ Model: originalModel,
+ BillingModel: billingModel,
+ UpstreamModel: upstreamModel,
+ Stream: false,
+ Duration: time.Since(startTime),
+ }, nil
+}
+
+func writeOpenAIEmbeddingsUpstreamResponse(c *gin.Context, resp *http.Response, body []byte, filter *responseheaders.CompiledHeaderFilter) {
+ if c == nil || resp == nil {
+ return
+ }
+ if c.Writer.Written() {
+ return
+ }
+ if resp.Header != nil {
+ responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, filter)
+ }
+ if ct := resp.Header.Get("Content-Type"); ct != "" {
+ c.Writer.Header().Set("Content-Type", ct)
+ } else {
+ c.Writer.Header().Set("Content-Type", "application/json")
+ }
+ c.Writer.WriteHeader(resp.StatusCode)
+ _, _ = c.Writer.Write(body)
+}
+
+func writeOpenAIEmbeddingsError(c *gin.Context, statusCode int, errType, message string) {
+ c.JSON(statusCode, gin.H{
+ "error": gin.H{
+ "type": errType,
+ "message": message,
+ },
+ })
+}
+
+func extractOpenAIEmbeddingsUsage(body []byte) OpenAIUsage {
+ usage := gjson.GetBytes(body, "usage")
+ if !usage.Exists() || !usage.IsObject() {
+ return OpenAIUsage{}
+ }
+ inputTokens := firstPositiveGJSONInt(
+ usage.Get("prompt_tokens"),
+ usage.Get("input_tokens"),
+ usage.Get("total_tokens"),
+ )
+ outputTokens := firstPositiveGJSONInt(
+ usage.Get("completion_tokens"),
+ usage.Get("output_tokens"),
+ )
+ cacheReadTokens := firstPositiveGJSONInt(
+ usage.Get("prompt_tokens_details.cached_tokens"),
+ usage.Get("input_tokens_details.cached_tokens"),
+ usage.Get("cache_read_tokens"),
+ usage.Get("cache_read_input_tokens"),
+ )
+ cacheCreationTokens := firstPositiveGJSONInt(
+ usage.Get("cache_creation_tokens"),
+ usage.Get("cache_creation_input_tokens"),
+ usage.Get("input_tokens_details.cache_creation_tokens"),
+ )
+ return OpenAIUsage{
+ InputTokens: inputTokens,
+ OutputTokens: outputTokens,
+ CacheReadInputTokens: cacheReadTokens,
+ CacheCreationInputTokens: cacheCreationTokens,
+ }
+}
+
+func firstPositiveGJSONInt(values ...gjson.Result) int {
+ for _, value := range values {
+ if !value.Exists() {
+ continue
+ }
+ n := int(value.Int())
+ if n > 0 {
+ return n
+ }
+ }
+ return 0
+}
+
+func buildOpenAIEmbeddingsURL(base string) string {
+ return buildOpenAIEndpointURL(base, "/v1/embeddings")
+}
diff --git a/backend/internal/service/openai_embeddings_test.go b/backend/internal/service/openai_embeddings_test.go
new file mode 100644
index 00000000..c7e89d64
--- /dev/null
+++ b/backend/internal/service/openai_embeddings_test.go
@@ -0,0 +1,106 @@
+package service
+
+import (
+ "bytes"
+ "context"
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
+)
+
+func TestBuildOpenAIEmbeddingsURL(t *testing.T) {
+ t.Parallel()
+
+ tests := []struct {
+ name string
+ base string
+ want string
+ }{
+ {"bare domain", "https://api.openai.com", "https://api.openai.com/v1/embeddings"},
+ {"bare /v1", "https://api.openai.com/v1", "https://api.openai.com/v1/embeddings"},
+ {"already embeddings", "https://api.openai.com/v1/embeddings", "https://api.openai.com/v1/embeddings"},
+ {"third-party versioned path", "https://open.bigmodel.cn/api/paas/v4", "https://open.bigmodel.cn/api/paas/v4/embeddings"},
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ t.Parallel()
+ require.Equal(t, tt.want, buildOpenAIEmbeddingsURL(tt.base))
+ })
+ }
+}
+
+func TestForwardEmbeddings_APIKeyPassthroughRecordsUsageAndBatchInput(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ reqBody := []byte(`{
+ "model":"nowledge-embedding",
+ "input":["hello","world"],
+ "encoding_format":"float",
+ "dimensions":256
+ }`)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/embeddings", bytes.NewReader(reqBody))
+ c.Request.Header.Set("Content-Type", "application/json")
+
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "Content-Type": []string{"application/json"},
+ "X-Request-Id": []string{"emb-rid"},
+ },
+ Body: io.NopCloser(strings.NewReader(`{
+ "object":"list",
+ "data":[
+ {"object":"embedding","index":0,"embedding":[0.1,0.2]},
+ {"object":"embedding","index":1,"embedding":[0.3,0.4]}
+ ],
+ "model":"jina-embeddings-v5-text-small",
+ "usage":{"prompt_tokens":13,"total_tokens":13}
+ }`)),
+ }}
+ svc := &OpenAIGatewayService{
+ cfg: &config.Config{},
+ httpUpstream: upstream,
+ }
+ account := &Account{
+ ID: 42,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://api.jina.ai",
+ "model_mapping": map[string]any{
+ "nowledge-embedding": "jina-embeddings-v5-text-small",
+ },
+ },
+ }
+
+ result, err := svc.ForwardEmbeddings(context.Background(), c, account, reqBody, "")
+
+ require.NoError(t, err)
+ require.Equal(t, http.StatusOK, rec.Code)
+ require.NotNil(t, result)
+ require.Equal(t, "emb-rid", result.RequestID)
+ require.Equal(t, "nowledge-embedding", result.Model)
+ require.Equal(t, "jina-embeddings-v5-text-small", result.BillingModel)
+ require.Equal(t, "jina-embeddings-v5-text-small", result.UpstreamModel)
+ require.Equal(t, 13, result.Usage.InputTokens)
+ require.Equal(t, 0, result.Usage.OutputTokens)
+ require.Equal(t, "https://api.jina.ai/v1/embeddings", upstream.lastReq.URL.String())
+ require.Equal(t, "Bearer sk-test", upstream.lastReq.Header.Get("Authorization"))
+ require.Equal(t, "jina-embeddings-v5-text-small", gjson.GetBytes(upstream.lastBody, "model").String())
+ require.Equal(t, int64(2), gjson.GetBytes(upstream.lastBody, "input.#").Int())
+ require.Equal(t, "hello", gjson.GetBytes(upstream.lastBody, "input.0").String())
+ require.Equal(t, "world", gjson.GetBytes(upstream.lastBody, "input.1").String())
+ require.Equal(t, "float", gjson.GetBytes(upstream.lastBody, "encoding_format").String())
+ require.Equal(t, int64(256), gjson.GetBytes(upstream.lastBody, "dimensions").Int())
+}
diff --git a/backend/internal/service/openai_failover_cached_body_test.go b/backend/internal/service/openai_failover_cached_body_test.go
new file mode 100644
index 00000000..776f1b1a
--- /dev/null
+++ b/backend/internal/service/openai_failover_cached_body_test.go
@@ -0,0 +1,134 @@
+package service
+
+import (
+ "bytes"
+ "context"
+ "errors"
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
+)
+
+func TestOpenAIGatewayService_Forward_FailoverReparsesCachedBodyForNextAccount(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ tests := []struct {
+ name string
+ requestModel string
+ firstMapping map[string]any
+ secondMapping map[string]any
+ wantFirst string
+ wantSecond string
+ }{
+ {
+ name: "both accounts have mapping",
+ firstMapping: map[string]any{"alias-model": "base-model-a"},
+ secondMapping: map[string]any{"alias-model": "base-model-b"},
+ wantFirst: "base-model-a",
+ wantSecond: "base-model-b",
+ },
+ {
+ name: "first account has mapping second account has none",
+ requestModel: "gpt-5.4-high",
+ firstMapping: map[string]any{"gpt-5.4-high": "gpt-5.4"},
+ wantFirst: "gpt-5.4",
+ wantSecond: "gpt-5.4",
+ },
+ {
+ name: "first account has no mapping second account has mapping",
+ secondMapping: map[string]any{"alias-model": "base-model-b"},
+ wantFirst: "alias-model",
+ wantSecond: "base-model-b",
+ },
+ {
+ name: "legacy context cache is ignored when mappings differ",
+ firstMapping: map[string]any{"alias-model": "base-model-a"},
+ secondMapping: map[string]any{"alias-model": "base-model-b"},
+ wantFirst: "base-model-a",
+ wantSecond: "base-model-b",
+ },
+ }
+
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ requestModel := tt.requestModel
+ if requestModel == "" {
+ requestModel = "alias-model"
+ }
+ body := []byte(`{"model":"` + requestModel + `","stream":false,"instructions":"cache-test","input":"hello"}`)
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
+ c.Request.Header.Set("Content-Type", "application/json")
+
+ upstream := &httpUpstreamRecorder{responses: []*http.Response{
+ {
+ StatusCode: http.StatusTooManyRequests,
+ Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-failover-a"}},
+ Body: io.NopCloser(strings.NewReader(`{"error":{"type":"rate_limit_error","message":"rate limited"}}`)),
+ },
+ {
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-ok-b"}},
+ Body: io.NopCloser(strings.NewReader(`{"id":"resp_123","status":"completed","model":"ok","output":[],"usage":{"input_tokens":1,"output_tokens":1}}`)),
+ },
+ }}
+ svc := &OpenAIGatewayService{httpUpstream: upstream}
+
+ firstAccount := openAIFailoverCachedBodyTestAccount(1, "account-a", tt.firstMapping)
+ secondAccount := openAIFailoverCachedBodyTestAccount(2, "account-b", tt.secondMapping)
+
+ _, err := svc.Forward(context.Background(), c, firstAccount, body)
+ require.Error(t, err)
+ var failoverErr *UpstreamFailoverError
+ require.True(t, errors.As(err, &failoverErr))
+ require.Len(t, upstream.bodies, 1)
+ require.Equal(t, tt.wantFirst, gjson.GetBytes(upstream.bodies[0], "model").String())
+
+ c.Set("openai_parsed_request_body", map[string]any{"model": tt.wantFirst, "stream": true})
+ result, err := svc.Forward(context.Background(), c, secondAccount, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Len(t, upstream.bodies, 2)
+ require.Equal(t, tt.wantSecond, gjson.GetBytes(upstream.bodies[1], "model").String())
+ })
+ }
+}
+
+func TestGetOpenAIRequestBodyMap_IgnoresLegacyContextCache(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Set("openai_parsed_request_body", map[string]any{"model": "base-model-a", "stream": true})
+
+ got, err := getOpenAIRequestBodyMap(c, []byte(`{"model":"alias-model","stream":false}`))
+ require.NoError(t, err)
+ require.Equal(t, "alias-model", got["model"])
+ require.Equal(t, false, got["stream"])
+}
+
+func openAIFailoverCachedBodyTestAccount(id int64, name string, mapping map[string]any) *Account {
+ credentials := map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-account"}
+ if mapping != nil {
+ credentials["model_mapping"] = mapping
+ }
+ return &Account{
+ ID: id,
+ Name: name,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Concurrency: 1,
+ Credentials: credentials,
+ Status: StatusActive,
+ Schedulable: true,
+ RateMultiplier: f64p(1),
+ }
+}
diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go
index f44d88cf..6e91d85c 100644
--- a/backend/internal/service/openai_gateway_chat_completions.go
+++ b/backend/internal/service/openai_gateway_chat_completions.go
@@ -241,7 +241,7 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions(
// 8. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -276,14 +276,14 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions(
Message: upstreamMsg,
Detail: upstreamDetail,
})
- s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
+ s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && (isPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
+ RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
}
}
- return s.handleChatCompletionsErrorResponse(resp, c, account)
+ return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel)
}
// 9. Handle normal response
@@ -358,8 +358,9 @@ func (s *OpenAIGatewayService) handleChatCompletionsErrorResponse(
resp *http.Response,
c *gin.Context,
account *Account,
+ requestedModel ...string,
) (*OpenAIForwardResult, error) {
- return s.handleCompatErrorResponse(resp, c, account, writeChatCompletionsError)
+ return s.handleCompatErrorResponse(resp, c, account, writeChatCompletionsError, requestedModel...)
}
// handleChatBufferedStreamingResponse reads all Responses SSE events from the
diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go
index efac4671..3ff6fac4 100644
--- a/backend/internal/service/openai_gateway_chat_completions_raw.go
+++ b/backend/internal/service/openai_gateway_chat_completions_raw.go
@@ -183,7 +183,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
// 7. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -208,14 +208,14 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions(
Message: upstreamMsg,
Detail: upstreamDetail,
})
- s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
+ s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && (isPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
+ RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
}
}
- return s.handleChatCompletionsErrorResponse(resp, c, account)
+ return s.handleChatCompletionsErrorResponse(resp, c, account, billingModel)
}
// 8. Forward response
diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go
index a624175b..4398bd27 100644
--- a/backend/internal/service/openai_gateway_messages.go
+++ b/backend/internal/service/openai_gateway_messages.go
@@ -300,7 +300,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
// 8. Handle error response with failover
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -338,15 +338,15 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic(
Message: upstreamMsg,
Detail: upstreamDetail,
})
- s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
+ s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && (isPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
+ RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
}
}
// Non-failover error: return Anthropic-formatted error to client
- return s.handleAnthropicErrorResponse(resp, c, account)
+ return s.handleAnthropicErrorResponse(resp, c, account, billingModel)
}
if account.Type == AccountTypeOAuth && promptCacheKey != "" {
@@ -413,8 +413,9 @@ func (s *OpenAIGatewayService) handleAnthropicErrorResponse(
resp *http.Response,
c *gin.Context,
account *Account,
+ requestedModel ...string,
) (*OpenAIForwardResult, error) {
- return s.handleCompatErrorResponse(resp, c, account, writeAnthropicError)
+ return s.handleCompatErrorResponse(resp, c, account, writeAnthropicError, requestedModel...)
}
// handleAnthropicBufferedStreamingResponse reads all Responses SSE events from
diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go
index 9769a82e..318c0861 100644
--- a/backend/internal/service/openai_gateway_record_usage_test.go
+++ b/backend/internal/service/openai_gateway_record_usage_test.go
@@ -721,6 +721,37 @@ func TestOpenAIGatewayServiceRecordUsage_PrefersClientRequestIDOverUpstreamReque
require.Equal(t, "client:openai-client-stable-123", usageRepo.lastLog.RequestID)
}
+func TestOpenAIGatewayServiceRecordUsage_WSModePrefersUpstreamRequestIDOverClientRequestID(t *testing.T) {
+ usageRepo := &openAIRecordUsageLogRepoStub{}
+ billingRepo := &openAIRecordUsageBillingRepoStub{result: &UsageBillingApplyResult{Applied: true}}
+ userRepo := &openAIRecordUsageUserRepoStub{}
+ subRepo := &openAIRecordUsageSubRepoStub{}
+ svc := newOpenAIRecordUsageServiceWithBillingRepoForTest(usageRepo, billingRepo, userRepo, subRepo, nil)
+
+ ctx := context.WithValue(context.Background(), ctxkey.ClientRequestID, "openai-ws-connection-123")
+ err := svc.RecordUsage(ctx, &OpenAIRecordUsageInput{
+ Result: &OpenAIForwardResult{
+ RequestID: "resp_openai_ws_turn_456",
+ OpenAIWSMode: true,
+ Usage: OpenAIUsage{
+ InputTokens: 8,
+ OutputTokens: 4,
+ },
+ Model: "gpt-5.1",
+ Duration: time.Second,
+ },
+ APIKey: &APIKey{ID: 10050},
+ User: &User{ID: 20050},
+ Account: &Account{ID: 30050},
+ })
+
+ require.NoError(t, err)
+ require.NotNil(t, billingRepo.lastCmd)
+ require.Equal(t, "resp_openai_ws_turn_456", billingRepo.lastCmd.RequestID)
+ require.NotNil(t, usageRepo.lastLog)
+ require.Equal(t, "resp_openai_ws_turn_456", usageRepo.lastLog.RequestID)
+}
+
func TestOpenAIGatewayServiceRecordUsage_GeneratesRequestIDWhenAllSourcesMissing(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{}
billingRepo := &openAIRecordUsageBillingRepoStub{result: &UsageBillingApplyResult{Applied: true}}
diff --git a/backend/internal/service/openai_gateway_responses_chat_fallback.go b/backend/internal/service/openai_gateway_responses_chat_fallback.go
index 91203bc1..205b27f7 100644
--- a/backend/internal/service/openai_gateway_responses_chat_fallback.go
+++ b/backend/internal/service/openai_gateway_responses_chat_fallback.go
@@ -163,7 +163,7 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -188,14 +188,14 @@ func (s *OpenAIGatewayService) forwardResponsesViaRawChatCompletions(
Message: upstreamMsg,
Detail: upstreamDetail,
})
- s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
+ s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, upstreamModel)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && (isPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
+ RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
}
}
- return s.handleErrorResponse(ctx, resp, c, account, chatBody)
+ return s.handleErrorResponse(ctx, resp, c, account, chatBody, billingModel)
}
if clientStream {
diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go
index d4921511..cfe92757 100644
--- a/backend/internal/service/openai_gateway_service.go
+++ b/backend/internal/service/openai_gateway_service.go
@@ -46,11 +46,11 @@ const (
// codex_cli_only 拒绝时单个请求头日志长度上限(字符)
codexCLIOnlyHeaderValueMaxBytes = 256
- // OpenAIParsedRequestBodyKey 缓存 handler 侧已解析的请求体,避免重复解析。
- OpenAIParsedRequestBodyKey = "openai_parsed_request_body"
// OpenAI WS Mode 失败后的重连次数上限(不含首次尝试)。
// 与 Codex 客户端保持一致:失败后最多重连 5 次。
openAIWSReconnectRetryLimit = 5
+ // 上游错误体只需要提取错误 JSON/日志摘要,默认 512KiB 避免错误风暴叠加大请求体。
+ openAIUpstreamErrorBodyReadLimit int64 = 512 << 10
// OpenAI WS Mode 重连退避默认值(可由配置覆盖)。
openAIWSRetryBackoffInitialDefault = 120 * time.Millisecond
openAIWSRetryBackoffMaxDefault = 2 * time.Second
@@ -59,6 +59,10 @@ const (
codexCLIVersion = "0.125.0"
// Codex 限额快照仅用于后台展示/诊断,不需要每个成功请求都立即落库。
openAICodexSnapshotPersistMinInterval = 30 * time.Second
+ // 配额自动暂停时,超过该时长仍未刷新的 used% 快照视为陈旧,不再据此暂停账号。
+ // 被暂停的账号收不到流量,其快照永远不会从上游响应头刷新;该兜底让账号在快照
+ // 陈旧时放行一次请求,从而通过正常响应头自愈,而无需等待整个窗口(5h/7d)重置。
+ openAICodexAutoPauseStaleAfter = 2 * time.Hour
)
// OpenAI allowed headers whitelist (for non-passthrough).
@@ -243,6 +247,9 @@ type OpenAIForwardResult struct {
ImageOutputSizes []string
ImageSizeSource string
ImageSizeBreakdown map[string]int
+
+ wsReplayInput []json.RawMessage
+ wsReplayInputExists bool
}
type OpenAIWSRetryMetricsSnapshot struct {
@@ -901,7 +908,17 @@ func SnapshotOpenAICompatibilityFallbackMetrics() OpenAICompatibilityFallbackMet
}
func (s *OpenAIGatewayService) detectCodexClientRestriction(c *gin.Context, account *Account) CodexClientRestrictionDetectionResult {
- return s.getCodexClientRestrictionDetector().Detect(c, account)
+ var globalAllowedClients []string
+ if account != nil && account.IsCodexCLIOnlyEnabled() && s != nil && s.settingService != nil {
+ ctx := context.Background()
+ if c != nil && c.Request != nil {
+ ctx = c.Request.Context()
+ }
+ if s.settingService.IsOpenAIAllowClaudeCodeCodexPluginEnabled(ctx) {
+ globalAllowedClients = []string{openai.AllowedClientClaudeCode}
+ }
+ }
+ return s.getCodexClientRestrictionDetector().Detect(c, account, globalAllowedClients)
}
func getAPIKeyIDFromContext(c *gin.Context) int64 {
@@ -959,6 +976,7 @@ func logCodexCLIOnlyDetection(ctx context.Context, c *gin.Context, account *Acco
}
log := logger.FromContext(ctx).With(fields...)
if result.Matched {
+ log.Info("OpenAI codex_cli_only 放行请求")
return
}
log.Warn("OpenAI codex_cli_only 拒绝非官方客户端请求")
@@ -1279,7 +1297,7 @@ func (s *OpenAIGatewayService) SelectAccountForModel(ctx context.Context, groupI
// SelectAccountForModelWithExclusions selects an account supporting the requested model while excluding specified accounts.
// SelectAccountForModelWithExclusions 选择支持指定模型的账号,同时排除指定的账号。
func (s *OpenAIGatewayService) SelectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*Account, error) {
- return s.selectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, excludedIDs, false, 0)
+ return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, sessionHash, requestedModel, excludedIDs, false, 0, "")
}
// noAvailableOpenAISelectionError builds the standard "no account available" error
@@ -1312,19 +1330,255 @@ func openAICompactSupportTier(account *Account) int {
// isOpenAIAccountEligibleForRequest centralises the schedulable / OpenAI / model /
// compact-support checks used during account selection.
-func isOpenAIAccountEligibleForRequest(account *Account, requestedModel string, requireCompact bool) bool {
- if account == nil || !account.IsSchedulable() || !account.IsOpenAI() {
+func isOpenAIAccountEligibleForRequest(ctx context.Context, account *Account, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) bool {
+ if account == nil || !account.IsOpenAI() || !account.IsSchedulableForModelWithContext(ctx, requestedModel) {
+ return false
+ }
+ if paused, reason := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused {
+ // Debug level: this fires per-candidate on the scheduling hot path, so Info
+ // would amplify into log spam once several accounts cross the threshold.
+ slog.Debug("account_auto_paused_by_quota",
+ "account_id", account.ID,
+ "window", reason.window,
+ "threshold", reason.threshold,
+ "utilization", reason.utilization,
+ )
return false
}
if requestedModel != "" && !account.IsModelSupported(requestedModel) {
return false
}
+ if !account.SupportsOpenAIEndpointCapability(requiredCapability) {
+ return false
+ }
if requireCompact && openAICompactSupportTier(account) == 0 {
return false
}
return true
}
+type openAIQuotaAutoPauseDecision struct {
+ window string
+ threshold float64
+ utilization float64
+}
+
+func shouldAutoPauseOpenAIAccountByQuota(ctx context.Context, account *Account) (bool, openAIQuotaAutoPauseDecision) {
+ if account == nil || !account.IsOpenAI() {
+ return false, openAIQuotaAutoPauseDecision{}
+ }
+ // Per-account explicit-disable flags must take precedence over the global default.
+ // Without these, leaving the account threshold blank means "use global default",
+ // so an admin has no way to exempt a single account from auto-pause once a global
+ // default exists. The disable flag is per-window so an account can opt out of
+ // only 5h or only 7d auto-pause.
+ disabled5h := resolveAccountExtraBool(account.Extra, "auto_pause_5h_disabled")
+ disabled7d := resolveAccountExtraBool(account.Extra, "auto_pause_7d_disabled")
+ threshold5h, threshold7d := resolveOpenAIQuotaAutoPauseThresholds(ctx, account)
+ now := time.Now()
+ if !disabled5h && threshold5h > 0 {
+ if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "5h", now); ok && utilization >= threshold5h {
+ return true, openAIQuotaAutoPauseDecision{window: "5h", threshold: threshold5h, utilization: utilization}
+ }
+ }
+ if !disabled7d && threshold7d > 0 {
+ if utilization, ok := resolveOpenAIQuotaUtilization(account.Extra, "7d", now); ok && utilization >= threshold7d {
+ return true, openAIQuotaAutoPauseDecision{window: "7d", threshold: threshold7d, utilization: utilization}
+ }
+ }
+ return false, openAIQuotaAutoPauseDecision{}
+}
+
+// resolveAccountExtraBool reads a bool-like value from account extra, tolerating
+// the few shapes JSON unmarshalling may produce (real bool, "true"/"false"
+// strings, 0/1 numbers).
+func resolveAccountExtraBool(extra map[string]any, key string) bool {
+ if len(extra) == 0 {
+ return false
+ }
+ value, ok := extra[key]
+ if !ok || value == nil {
+ return false
+ }
+ switch v := value.(type) {
+ case bool:
+ return v
+ case string:
+ parsed, err := strconv.ParseBool(strings.TrimSpace(v))
+ return err == nil && parsed
+ case float64:
+ return v != 0
+ case float32:
+ return v != 0
+ case int:
+ return v != 0
+ case int64:
+ return v != 0
+ case json.Number:
+ if i, err := v.Int64(); err == nil {
+ return i != 0
+ }
+ }
+ return false
+}
+
+func resolveOpenAIQuotaAutoPauseThresholds(ctx context.Context, account *Account) (float64, float64) {
+ threshold5h, _ := resolveAccountExtraNumber(account.Extra, "auto_pause_5h_threshold")
+ threshold7d, _ := resolveAccountExtraNumber(account.Extra, "auto_pause_7d_threshold")
+ threshold5h = clamp01(threshold5h)
+ threshold7d = clamp01(threshold7d)
+ if threshold5h > 0 && threshold7d > 0 {
+ return threshold5h, threshold7d
+ }
+ settings := openAIQuotaAutoPauseSettingsFromContext(ctx)
+ if threshold5h <= 0 {
+ threshold5h = clamp01(settings.DefaultThreshold5h)
+ }
+ if threshold7d <= 0 {
+ threshold7d = clamp01(settings.DefaultThreshold7d)
+ }
+ return threshold5h, threshold7d
+}
+
+func resolveAccountExtraNumber(extra map[string]any, keys ...string) (float64, bool) {
+ if len(extra) == 0 {
+ return 0, false
+ }
+ for _, key := range keys {
+ value, ok := extra[key]
+ if !ok || value == nil {
+ continue
+ }
+ switch v := value.(type) {
+ case float64:
+ return v, true
+ case float32:
+ return float64(v), true
+ case int:
+ return float64(v), true
+ case int64:
+ return float64(v), true
+ case json.Number:
+ parsed, err := v.Float64()
+ if err == nil {
+ return parsed, true
+ }
+ case string:
+ parsed, err := strconv.ParseFloat(strings.TrimSpace(v), 64)
+ if err == nil {
+ return parsed, true
+ }
+ }
+ }
+ return 0, false
+}
+
+// resolveOpenAIQuotaUtilization returns the current utilization ratio (0..1) for the
+// given Codex usage window. ok=false means there is no usable signal to pause on:
+// either no snapshot exists, or the window has already rolled over so the cached
+// percentage is stale. The stale guard matters because a paused account stops
+// receiving requests, so its snapshot is never refreshed from upstream headers —
+// without this check an old used_percent would keep the account paused forever even
+// after the real window reset.
+func resolveOpenAIQuotaUtilization(extra map[string]any, window string, now time.Time) (float64, bool) {
+ usedPercent := readOpenAIQuotaUsedPercent(extra, window)
+ if usedPercent <= 0 {
+ return 0, false
+ }
+ if openAIQuotaWindowReset(extra, window, now) {
+ return 0, false
+ }
+ // 快照过于陈旧(账号长期未收到流量刷新)时,不再据此暂停。放行后下一次响应头
+ // 会刷新快照实现自愈,避免账号在错误/过期的 used% 上被永久跳过(issue #2994)。
+ if openAICodexSnapshotStaleForPause(extra, now) {
+ return 0, false
+ }
+ return usedPercent / 100, true
+}
+
+// openAICodexSnapshotStaleForPause reports whether the Codex usage snapshot is stale
+// enough that it should no longer keep an account auto-paused. It anchors on
+// codex_usage_updated_at (always written by buildCodexUsageExtraUpdates). A missing or
+// unparseable timestamp returns false (treated as fresh, so the account stays paused) —
+// this is deliberate: it prevents any snapshot without a write time from silently escaping
+// auto-pause, and a genuinely-exhausted account that is actively served refreshes the
+// timestamp on every response so it never crosses the staleness bound.
+func openAICodexSnapshotStaleForPause(extra map[string]any, now time.Time) bool {
+ if len(extra) == 0 {
+ return false
+ }
+ updatedRaw, ok := extra["codex_usage_updated_at"]
+ if !ok {
+ return false
+ }
+ updatedAt, err := parseTime(fmt.Sprint(updatedRaw))
+ if err != nil {
+ return false
+ }
+ return now.Sub(updatedAt) >= openAICodexAutoPauseStaleAfter
+}
+
+// openAIQuotaWindowReset reports whether the Codex usage window's reset time has
+// already passed relative to now. It prefers the absolute codex__reset_at
+// timestamp and falls back to codex__reset_after_seconds anchored at
+// codex_usage_updated_at, mirroring AccountUsageService's window-progress logic.
+func openAIQuotaWindowReset(extra map[string]any, window string, now time.Time) bool {
+ if len(extra) == 0 {
+ return false
+ }
+ if resetAtRaw, ok := extra["codex_"+window+"_reset_at"]; ok {
+ if resetAt, err := parseTime(fmt.Sprint(resetAtRaw)); err == nil {
+ return !now.Before(resetAt)
+ }
+ }
+ resetAfter := parseExtraInt(extra["codex_"+window+"_reset_after_seconds"])
+ if resetAfter <= 0 {
+ return false
+ }
+ base := now
+ if updatedRaw, ok := extra["codex_usage_updated_at"]; ok {
+ if updatedAt, err := parseTime(fmt.Sprint(updatedRaw)); err == nil {
+ base = updatedAt
+ }
+ }
+ resetAt := base.Add(time.Duration(resetAfter) * time.Second)
+ return !now.Before(resetAt)
+}
+
+func readOpenAIQuotaUsedPercent(extra map[string]any, window string) float64 {
+ if len(extra) == 0 {
+ return 0
+ }
+ if value, ok := resolveAccountExtraNumber(extra, "codex_"+window+"_used_percent"); ok {
+ return value
+ }
+ return 0
+}
+
+type openAIQuotaAutoPauseCtxKey struct{}
+
+func withOpenAIQuotaAutoPauseSettings(ctx context.Context, settings OpsOpenAIAccountQuotaAutoPauseSettings) context.Context {
+ if ctx == nil {
+ ctx = context.Background()
+ }
+ return context.WithValue(ctx, openAIQuotaAutoPauseCtxKey{}, settings)
+}
+
+func openAIQuotaAutoPauseSettingsFromContext(ctx context.Context) OpsOpenAIAccountQuotaAutoPauseSettings {
+ if ctx == nil {
+ return OpsOpenAIAccountQuotaAutoPauseSettings{}
+ }
+ settings, _ := ctx.Value(openAIQuotaAutoPauseCtxKey{}).(OpsOpenAIAccountQuotaAutoPauseSettings)
+ return settings
+}
+
+func (s *OpenAIGatewayService) withOpenAIQuotaAutoPauseContext(ctx context.Context) context.Context {
+ if s == nil || s.settingService == nil {
+ return ctx
+ }
+ return withOpenAIQuotaAutoPauseSettings(ctx, s.settingService.GetOpenAIQuotaAutoPauseSettings(ctx))
+}
+
// prioritizeOpenAICompactAccounts re-orders a slice so that accounts with known
// compact support are tried first, followed by unknown, then explicitly unsupported.
// The relative order within each tier is preserved.
@@ -1366,7 +1620,7 @@ func resolveOpenAIAccountUpstreamModelForRequest(account *Account, requestedMode
return upstreamModel
}
-func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64) (*Account, error) {
+func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) (*Account, error) {
if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) {
slog.Warn("channel pricing restriction blocked request",
"group_id", derefGroupID(groupID),
@@ -1376,7 +1630,7 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C
// 1. 尝试粘性会话命中
// Try sticky session hit
- if account := s.tryStickySessionHit(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID); account != nil {
+ if account := s.tryStickySessionHit(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability); account != nil {
return account, nil
}
@@ -1389,7 +1643,7 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C
// 3. 按优先级 + LRU 选择最佳账号
// Select by priority + LRU
- selected, compactBlocked := s.selectBestAccount(ctx, groupID, accounts, requestedModel, excludedIDs, requireCompact)
+ selected, compactBlocked := s.selectBestAccount(ctx, groupID, accounts, requestedModel, excludedIDs, requireCompact, requiredCapability)
if selected == nil {
return nil, noAvailableOpenAISelectionError(requestedModel, compactBlocked)
@@ -1414,7 +1668,7 @@ func (s *OpenAIGatewayService) selectAccountForModelWithExclusions(ctx context.C
//
// tryStickySessionHit attempts to get account from sticky session.
// Returns account if hit and usable; clears session and returns nil if account is unavailable.
-func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID *int64, sessionHash, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64) *Account {
+func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID *int64, sessionHash, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, stickyAccountID int64, requiredCapability OpenAIEndpointCapability) *Account {
if sessionHash == "" {
return nil
}
@@ -1446,14 +1700,14 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID
// 验证账号是否可用于当前请求
// Verify account is usable for current request
- if !isOpenAIAccountEligibleForRequest(account, requestedModel, false) {
+ if !isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, false, requiredCapability) {
return nil
}
if s.isOpenAIAccountRuntimeBlocked(account) {
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
return nil
}
- account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact)
+ account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact, requiredCapability)
if account == nil {
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
return nil
@@ -1477,7 +1731,7 @@ func (s *OpenAIGatewayService) tryStickySessionHit(ctx context.Context, groupID
// Returns nil if no available account. The second return reports whether at
// least one candidate was filtered out solely because it lacks compact support
// (only meaningful when requireCompact=true).
-func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool) (*Account, bool) {
+func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *int64, accounts []Account, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*Account, bool) {
var selected *Account
selectedCompactTier := -1
compactBlocked := false
@@ -1492,11 +1746,11 @@ func (s *OpenAIGatewayService) selectBestAccount(ctx context.Context, groupID *i
continue
}
- fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false)
+ fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false, requiredCapability)
if fresh == nil {
continue
}
- fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, false)
+ fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, false, requiredCapability)
if fresh == nil {
continue
}
@@ -1573,10 +1827,10 @@ func (s *OpenAIGatewayService) isBetterAccount(candidate, current *Account) bool
// SelectAccountWithLoadAwareness selects an account with load-awareness and wait plan.
func (s *OpenAIGatewayService) SelectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}) (*AccountSelectionResult, error) {
- return s.selectAccountWithLoadAwareness(ctx, groupID, sessionHash, requestedModel, excludedIDs, false)
+ return s.selectAccountWithLoadAwareness(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, sessionHash, requestedModel, excludedIDs, false, "")
}
-func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool) (*AccountSelectionResult, error) {
+func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Context, groupID *int64, sessionHash string, requestedModel string, excludedIDs map[int64]struct{}, requireCompact bool, requiredCapability OpenAIEndpointCapability) (*AccountSelectionResult, error) {
if s.checkChannelPricingRestriction(ctx, groupID, requestedModel) {
slog.Warn("channel pricing restriction blocked request",
"group_id", derefGroupID(groupID),
@@ -1593,7 +1847,7 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
}
}
if s.concurrencyService == nil || !cfg.LoadBatchEnabled {
- account, err := s.selectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID)
+ account, err := s.selectAccountForModelWithExclusions(ctx, groupID, sessionHash, requestedModel, excludedIDs, requireCompact, stickyAccountID, requiredCapability)
if err != nil {
return nil, err
}
@@ -1646,8 +1900,8 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
if clearSticky {
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
}
- if !clearSticky && isOpenAIAccountEligibleForRequest(account, requestedModel, false) {
- account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact)
+ if !clearSticky && isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, false, requiredCapability) {
+ account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact, requiredCapability)
if account == nil {
_ = s.deleteStickySessionAccountID(ctx, groupID, sessionHash)
} else if s.isOpenAIAccountRuntimeBlocked(account) {
@@ -1691,15 +1945,12 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
// Scheduler snapshots can be temporarily stale (bucket rebuild is throttled);
// re-check schedulability here so recently rate-limited/overloaded accounts
// are not selected again before the bucket is rebuilt.
- if !acc.IsSchedulable() {
+ if !isOpenAIAccountEligibleForRequest(ctx, acc, requestedModel, false, requiredCapability) {
continue
}
if s.isOpenAIAccountRuntimeBlocked(acc) {
continue
}
- if requestedModel != "" && !acc.IsModelSupported(requestedModel) {
- continue
- }
if needsUpstreamCheck && s.isUpstreamModelRestrictedByChannel(ctx, *groupID, acc, requestedModel, requireCompact) {
continue
}
@@ -1779,11 +2030,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
}
for _, item := range selectionOrder {
- fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, item.account, requestedModel, false)
+ fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, item.account, requestedModel, false, requiredCapability)
if fresh == nil {
continue
}
- fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact)
+ fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact, requiredCapability)
if fresh == nil {
continue
}
@@ -1813,11 +2064,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
ordered = prioritizeOpenAICompactAccounts(ordered)
}
for _, acc := range ordered {
- fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false)
+ fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false, requiredCapability)
if fresh == nil {
continue
}
- fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact)
+ fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact, requiredCapability)
if fresh == nil {
continue
}
@@ -1858,11 +2109,11 @@ func (s *OpenAIGatewayService) selectAccountWithLoadAwareness(ctx context.Contex
candidates = prioritizeOpenAICompactAccounts(candidates)
}
for _, acc := range candidates {
- fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false)
+ fresh := s.resolveFreshSchedulableOpenAIAccount(ctx, acc, requestedModel, false, requiredCapability)
if fresh == nil {
continue
}
- fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact)
+ fresh = s.recheckSelectedOpenAIAccountFromDB(ctx, fresh, requestedModel, requireCompact, requiredCapability)
if fresh == nil {
continue
}
@@ -1910,7 +2161,7 @@ func (s *OpenAIGatewayService) tryAcquireAccountSlot(ctx context.Context, accoun
return s.concurrencyService.AcquireAccountSlot(ctx, accountID, maxConcurrency)
}
-func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.Context, account *Account, requestedModel string, requireCompact bool) *Account {
+func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.Context, account *Account, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account {
if account == nil {
return nil
}
@@ -1924,7 +2175,7 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.
fresh = current
}
- if !isOpenAIAccountEligibleForRequest(fresh, requestedModel, requireCompact) {
+ if !isOpenAIAccountEligibleForRequest(ctx, fresh, requestedModel, requireCompact, requiredCapability) {
return nil
}
if s.isOpenAIAccountRuntimeBlocked(fresh) {
@@ -1933,12 +2184,12 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context.
return fresh
}
-func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, requestedModel string, requireCompact bool) *Account {
+func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Context, account *Account, requestedModel string, requireCompact bool, requiredCapability OpenAIEndpointCapability) *Account {
if account == nil {
return nil
}
if s.schedulerSnapshot == nil || s.accountRepo == nil {
- if !isOpenAIAccountEligibleForRequest(account, requestedModel, requireCompact) {
+ if !isOpenAIAccountEligibleForRequest(ctx, account, requestedModel, requireCompact, requiredCapability) {
return nil
}
return account
@@ -1948,7 +2199,7 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co
if err != nil || latest == nil {
return nil
}
- if !isOpenAIAccountEligibleForRequest(latest, requestedModel, requireCompact) {
+ if !isOpenAIAccountEligibleForRequest(ctx, latest, requestedModel, requireCompact, requiredCapability) {
return nil
}
if s.isOpenAIAccountRuntimeBlocked(latest) {
@@ -2067,8 +2318,46 @@ func (s *OpenAIGatewayService) shouldFailoverOpenAIUpstreamResponse(statusCode i
return isOpenAITransientProcessingError(statusCode, upstreamMsg, upstreamBody)
}
-func (s *OpenAIGatewayService) handleFailoverSideEffects(ctx context.Context, resp *http.Response, account *Account) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+func marshalOpenAIUpstreamJSON(v any) ([]byte, error) {
+ var buf bytes.Buffer
+ enc := json.NewEncoder(&buf)
+ enc.SetEscapeHTML(false)
+ if err := enc.Encode(v); err != nil {
+ return nil, err
+ }
+ out := buf.Bytes()
+ if len(out) > 0 && out[len(out)-1] == '\n' {
+ out = out[:len(out)-1]
+ }
+ return out, nil
+}
+
+func openAIUpstreamErrorBodyReadLimitForConfig(cfg *config.Config) int64 {
+ limit := openAIUpstreamErrorBodyReadLimit
+ if cfg != nil && cfg.Gateway.LogUpstreamErrorBody && cfg.Gateway.LogUpstreamErrorBodyMaxBytes > int(limit) {
+ limit = int64(cfg.Gateway.LogUpstreamErrorBodyMaxBytes)
+ }
+ return limit
+}
+
+func (s *OpenAIGatewayService) readUpstreamErrorBody(resp *http.Response) []byte {
+ if resp == nil || resp.Body == nil {
+ return nil
+ }
+ cfg := (*config.Config)(nil)
+ if s != nil {
+ cfg = s.cfg
+ }
+ body, _ := io.ReadAll(io.LimitReader(resp.Body, openAIUpstreamErrorBodyReadLimitForConfig(cfg)))
+ return body
+}
+
+func (s *OpenAIGatewayService) handleFailoverSideEffects(ctx context.Context, resp *http.Response, account *Account, requestedModel ...string) {
+ body := s.readUpstreamErrorBody(resp)
+ if len(requestedModel) > 0 {
+ s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, requestedModel[0])
+ return
+ }
s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body)
}
@@ -2091,7 +2380,8 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
}
originalBody := body
- reqModel, reqStream, promptCacheKey := extractOpenAIRequestMetaFromBody(body)
+ requestView := newOpenAIRequestView(body)
+ reqModel, reqStream, promptCacheKey := requestView.Model, requestView.Stream, requestView.PromptCacheKey
originalModel := reqModel
if account.Type == AccountTypeAPIKey && !openai_compat.ShouldUseResponsesAPI(account.Extra) {
@@ -2141,172 +2431,83 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
return s.forwardOpenAIPassthrough(ctx, c, account, originalBody, reqModel, reasoningEffort, reqStream, startTime)
}
- reqBody, err := getOpenAIRequestBodyMap(c, body)
- if err != nil {
- return nil, err
+ bodyModified := false
+ var reqBody map[string]any
+ ensureReqBody := func() (map[string]any, error) {
+ if requestView.HasPatches() {
+ patchedBody, patchErr := requestView.ApplyPatches()
+ if patchErr != nil {
+ return nil, patchErr
+ }
+ body = patchedBody
+ requestView = newOpenAIRequestView(body)
+ reqBody = nil
+ bodyModified = false
+ }
+ if reqBody != nil {
+ return reqBody, nil
+ }
+ decoded, decodeErr := requestView.Decode(c)
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ reqBody = decoded
+ return reqBody, nil
+ }
+ markPatchSet := func(path string, value any) {
+ bodyModified = true
+ if requestView.patchesDisabled {
+ if reqBody != nil {
+ setOpenAIRequestMapPath(reqBody, path, value)
+ }
+ return
+ }
+ requestView.MarkPatchSet(path, value)
+ }
+ markPatchDelete := func(path string) {
+ bodyModified = true
+ if requestView.patchesDisabled {
+ if reqBody != nil {
+ deleteOpenAIRequestMapPath(reqBody, path)
+ }
+ return
+ }
+ requestView.MarkPatchDelete(path)
+ }
+ disablePatch := func() {
+ requestView.DisablePatches()
+ }
+ markDecodedModified := func() {
+ bodyModified = true
+ disablePatch()
}
- if v, ok := reqBody["model"].(string); ok {
- reqModel = v
- originalModel = reqModel
- }
- if v, ok := reqBody["stream"].(bool); ok {
- reqStream = v
- }
- if promptCacheKey == "" {
- if v, ok := reqBody["prompt_cache_key"].(string); ok {
- promptCacheKey = strings.TrimSpace(v)
- }
- }
apiKey := getAPIKeyFromContext(c)
imageGenerationAllowed := GroupAllowsImageGeneration(nil)
if apiKey != nil {
imageGenerationAllowed = GroupAllowsImageGeneration(apiKey.Group)
}
codexImageGenerationBridgeEnabled := isCodexCLI && imageGenerationAllowed && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey)
- if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) && !imageGenerationAllowed {
+ imageIntent := IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body)
+ if imageIntent && !imageGenerationAllowed {
MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate)
- c.JSON(http.StatusForbidden, gin.H{
- "error": gin.H{
- "type": "permission_error",
- "message": ImageGenerationPermissionMessage(),
- },
- })
+ c.JSON(http.StatusForbidden, gin.H{"error": gin.H{"type": "permission_error", "message": ImageGenerationPermissionMessage()}})
return nil, errors.New("image generation disabled for group")
}
- // Track if body needs re-serialization
- bodyModified := false
- // 单字段补丁快速路径:只要整个变更集最终可归约为同一路径的 set/delete,就避免全量 Marshal。
- patchDisabled := false
- patchHasOp := false
- patchDelete := false
- patchPath := ""
- var patchValue any
- markPatchSet := func(path string, value any) {
- if strings.TrimSpace(path) == "" {
- patchDisabled = true
- return
- }
- if patchDisabled {
- return
- }
- if !patchHasOp {
- patchHasOp = true
- patchDelete = false
- patchPath = path
- patchValue = value
- return
- }
- if patchDelete || patchPath != path {
- patchDisabled = true
- return
- }
- patchValue = value
- }
- markPatchDelete := func(path string) {
- if strings.TrimSpace(path) == "" {
- patchDisabled = true
- return
- }
- if patchDisabled {
- return
- }
- if !patchHasOp {
- patchHasOp = true
- patchDelete = true
- patchPath = path
- return
- }
- if !patchDelete || patchPath != path {
- patchDisabled = true
- }
- }
- disablePatch := func() {
- patchDisabled = true
- }
-
- // 非透传模式下,instructions 为空时注入默认指令。
- if isInstructionsEmpty(reqBody) && !compatMessagesBridge {
- reqBody["instructions"] = "You are a helpful coding assistant."
- bodyModified = true
+ instructions := gjson.GetBytes(body, "instructions")
+ instructionsEmpty := !instructions.Exists() || instructions.Type != gjson.String || strings.TrimSpace(instructions.String()) == ""
+ if instructionsEmpty && !compatMessagesBridge {
markPatchSet("instructions", "You are a helpful coding assistant.")
}
- if codexImageGenerationBridgeEnabled && ensureOpenAIResponsesImageGenerationTool(reqBody) {
- bodyModified = true
- disablePatch()
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Injected /responses image_generation tool for Codex client")
- }
-
- if normalizeOpenAIResponsesImageGenerationTools(reqBody) {
- bodyModified = true
- disablePatch()
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image_generation tool payload")
- }
- if codexImageGenerationBridgeEnabled && applyCodexImageGenerationBridgeInstructions(reqBody) {
- bodyModified = true
- disablePatch()
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Added Codex image_generation bridge instructions")
- }
-
- // 对所有请求执行模型映射(包含 Codex CLI)。
billingModel := account.GetMappedModel(reqModel)
if billingModel != reqModel {
logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Model mapping applied: %s -> %s (account: %s, isCodexCLI: %v)", reqModel, billingModel, account.Name, isCodexCLI)
- reqBody["model"] = billingModel
- bodyModified = true
+ reqModel = billingModel
markPatchSet("model", billingModel)
}
upstreamModel := billingModel
- if imageGenerationAllowed && normalizeOpenAIResponsesImageOnlyModel(reqBody) {
- bodyModified = true
- disablePatch()
- if model, ok := reqBody["model"].(string); ok {
- upstreamModel = strings.TrimSpace(model)
- }
- logger.LegacyPrintf(
- "service.openai_gateway",
- "[OpenAI] Normalized /responses image-only model request inbound_model=%s image_model=%s upstream_model=%s",
- reqModel,
- billingModel,
- upstreamModel,
- )
- }
- if err := validateOpenAIResponsesImageModel(reqBody, upstreamModel); err != nil {
- setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "")
- c.JSON(http.StatusBadRequest, gin.H{
- "error": gin.H{
- "type": "invalid_request_error",
- "message": err.Error(),
- "param": "model",
- },
- })
- return nil, err
- }
- if hasOpenAIImageGenerationTool(reqBody) {
- logger.LegacyPrintf(
- "service.openai_gateway",
- "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s",
- reqModel,
- upstreamModel,
- account.Type,
- )
- }
- if err := validateCodexSparkInput(reqBody, upstreamModel); err != nil {
- setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "")
- c.JSON(http.StatusBadRequest, gin.H{
- "error": gin.H{
- "type": "invalid_request_error",
- "message": err.Error(),
- "param": "input",
- },
- })
- return nil, err
- }
-
- // Compact-only model 映射:仅在 /responses/compact 路径生效,且优先级高于
- // OAuth 模型规范化(避免 OAuth 规范化覆盖 compact-only 自定义模型)。
isCompactRequest := isOpenAIResponsesCompactPath(c)
compactMapped := false
if isCompactRequest {
@@ -2314,65 +2515,100 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
if compactMappedModel != "" && compactMappedModel != billingModel {
compactMapped = true
upstreamModel = compactMappedModel
- reqBody["model"] = compactMappedModel
- bodyModified = true
+ reqModel = compactMappedModel
markPatchSet("model", compactMappedModel)
logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Compact model mapping applied: %s -> %s (account: %s, isCodexCLI: %v)", billingModel, compactMappedModel, account.Name, isCodexCLI)
}
}
-
- // OpenAI OAuth 账号走 ChatGPT internal Codex endpoint,需要将模型名规范化为
- // 上游可识别的 Codex/GPT 系列。API Key 账号则应保留原始/映射后的模型名,
- // 以兼容自定义 base_url 的 OpenAI-compatible 上游。
- if model, ok := reqBody["model"].(string); ok {
- if !compactMapped {
- upstreamModel = normalizeOpenAIModelForUpstream(account, model)
- if upstreamModel != "" && upstreamModel != model {
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Upstream model resolved: %s -> %s (account: %s, type: %s, isCodexCLI: %v)",
- model, upstreamModel, account.Name, account.Type, isCodexCLI)
- reqBody["model"] = upstreamModel
- bodyModified = true
- markPatchSet("model", upstreamModel)
- }
+ if !compactMapped {
+ modelForNormalize := reqModel
+ if modelForNormalize == "" {
+ modelForNormalize = requestView.Model
}
-
- // 移除 gpt-5.2-codex 以下的版本 verbosity 参数
- // 确保高版本模型向低版本模型映射不报错
- if !SupportsVerbosity(upstreamModel) {
- if text, ok := reqBody["text"].(map[string]any); ok {
- delete(text, "verbosity")
- }
+ upstreamModel = normalizeOpenAIModelForUpstream(account, modelForNormalize)
+ if upstreamModel != "" && upstreamModel != modelForNormalize {
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Upstream model resolved: %s -> %s (account: %s, type: %s, isCodexCLI: %v)", modelForNormalize, upstreamModel, account.Name, account.Type, isCodexCLI)
+ reqModel = upstreamModel
+ markPatchSet("model", upstreamModel)
}
}
+ if strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()) == "minimal" {
+ markPatchSet("reasoning.effort", "none")
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized reasoning.effort: minimal -> none (account: %s)", account.Name)
+ }
- // 规范化 reasoning.effort 参数(minimal -> none),与上游允许值对齐。
- if reasoning, ok := reqBody["reasoning"].(map[string]any); ok {
- if effort, ok := reasoning["effort"].(string); ok && effort == "minimal" {
- reasoning["effort"] = "none"
- bodyModified = true
- markPatchSet("reasoning.effort", "none")
- logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized reasoning.effort: minimal -> none (account: %s)", account.Name)
+ imageIntent = imageIntent || IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, nil) || isOpenAIImageGenerationModel(upstreamModel)
+ if imageIntent && !imageGenerationAllowed {
+ MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate)
+ c.JSON(http.StatusForbidden, gin.H{"error": gin.H{"type": "permission_error", "message": ImageGenerationPermissionMessage()}})
+ return nil, errors.New("image generation disabled for group")
+ }
+
+ if imageGenerationAllowed && (codexImageGenerationBridgeEnabled || isOpenAIImageGenerationModel(requestView.Model) || openAIRequestBodyImageGenerationToolNeedsNormalization(body) || isOpenAIImageGenerationModel(upstreamModel)) {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ if codexImageGenerationBridgeEnabled && ensureOpenAIResponsesImageGenerationTool(decoded) {
+ markDecodedModified()
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Injected /responses image_generation tool for Codex client")
+ }
+ if normalizeOpenAIResponsesImageGenerationTools(decoded) {
+ markDecodedModified()
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image_generation tool payload")
+ }
+ if normalizeOpenAIResponsesImageOnlyModel(decoded) {
+ markDecodedModified()
+ if model, ok := decoded["model"].(string); ok {
+ upstreamModel = strings.TrimSpace(model)
+ }
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized /responses image-only model request inbound_model=%s image_model=%s upstream_model=%s", requestView.Model, billingModel, upstreamModel)
+ }
+ if err := validateOpenAIResponsesImageModel(decoded, upstreamModel); err != nil {
+ setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "")
+ c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": err.Error(), "param": "model"}})
+ return nil, err
+ }
+ if hasOpenAIImageGenerationTool(decoded) {
+ imageIntent = true
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", requestView.Model, upstreamModel, account.Type)
+ }
+ if codexImageGenerationBridgeEnabled && applyCodexImageGenerationBridgeInstructions(decoded) {
+ markDecodedModified()
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Added Codex image_generation bridge instructions")
+ }
+ } else if imageGenerationAllowed && imageIntent && openAIRequestBodyHasImageGenerationTool(body) {
+ // 完整 image_generation tool 只做 raw 计费读取,校验/桥接/旧字段迁移命中时才展开大 input map。
+ logger.LegacyPrintf("service.openai_gateway", "[OpenAI] /responses image_generation request inbound_model=%s mapped_model=%s account_type=%s", requestView.Model, upstreamModel, account.Type)
+ }
+
+ if isCodexSparkModel(upstreamModel) && openAIRequestBodyMayContainImageInput(body) {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ if err := validateCodexSparkInput(decoded, upstreamModel); err != nil {
+ setOpsUpstreamError(c, http.StatusBadRequest, err.Error(), "")
+ c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": err.Error(), "param": "input"}})
+ return nil, err
}
}
if account.Type == AccountTypeOAuth {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
codexResult := codexTransformResult{}
if compatMessagesBridge {
- codexResult = applyCodexOAuthTransformWithOptions(reqBody, codexOAuthTransformOptions{
- IsCodexCLI: isCodexCLI,
- IsCompact: isCompactRequest,
- SkipDefaultInstructions: true,
- PreserveToolCallIDs: true,
- })
- ensureCodexOAuthInstructionsField(reqBody)
- bodyModified = true
- disablePatch()
+ codexResult = applyCodexOAuthTransformWithOptions(decoded, codexOAuthTransformOptions{IsCodexCLI: isCodexCLI, IsCompact: isCompactRequest, SkipDefaultInstructions: true, PreserveToolCallIDs: true})
+ ensureCodexOAuthInstructionsField(decoded)
+ markDecodedModified()
} else {
- codexResult = applyCodexOAuthTransform(reqBody, isCodexCLI, isCompactRequest)
+ codexResult = applyCodexOAuthTransform(decoded, isCodexCLI, isCompactRequest)
}
if codexResult.Modified {
- bodyModified = true
- disablePatch()
+ markDecodedModified()
}
if codexResult.NormalizedModel != "" {
upstreamModel = codexResult.NormalizedModel
@@ -2382,90 +2618,57 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
}
}
- // Handle max_output_tokens based on platform and account type
+ if !SupportsVerbosity(upstreamModel) && gjson.GetBytes(body, "text.verbosity").Exists() {
+ markPatchDelete("text.verbosity")
+ }
+
if !isCodexCLI {
- if maxOutputTokens, hasMaxOutputTokens := reqBody["max_output_tokens"]; hasMaxOutputTokens {
+ maxOutputTokens := gjson.GetBytes(body, "max_output_tokens")
+ if maxOutputTokens.Exists() {
switch account.Platform {
case PlatformOpenAI:
- // For OpenAI API Key, remove max_output_tokens (not supported)
- // For OpenAI OAuth (Responses API), keep it (supported)
if account.Type == AccountTypeAPIKey {
- delete(reqBody, "max_output_tokens")
- bodyModified = true
markPatchDelete("max_output_tokens")
}
case PlatformAnthropic:
- // For Anthropic (Claude), convert to max_tokens
- delete(reqBody, "max_output_tokens")
- markPatchDelete("max_output_tokens")
- if _, hasMaxTokens := reqBody["max_tokens"]; !hasMaxTokens {
- reqBody["max_tokens"] = maxOutputTokens
- disablePatch()
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
}
- bodyModified = true
+ delete(decoded, "max_output_tokens")
+ if _, hasMaxTokens := decoded["max_tokens"]; !hasMaxTokens {
+ decoded["max_tokens"] = maxOutputTokens.Value()
+ }
+ markDecodedModified()
case PlatformGemini:
- // For Gemini, remove (will be handled by Gemini-specific transform)
- delete(reqBody, "max_output_tokens")
- bodyModified = true
markPatchDelete("max_output_tokens")
default:
- // For unknown platforms, remove to be safe
- delete(reqBody, "max_output_tokens")
- bodyModified = true
markPatchDelete("max_output_tokens")
}
}
-
- // Also handle max_completion_tokens (similar logic)
- if _, hasMaxCompletionTokens := reqBody["max_completion_tokens"]; hasMaxCompletionTokens {
- if account.Type == AccountTypeAPIKey || account.Platform != PlatformOpenAI {
- delete(reqBody, "max_completion_tokens")
- bodyModified = true
- markPatchDelete("max_completion_tokens")
- }
+ if gjson.GetBytes(body, "max_completion_tokens").Exists() && (account.Type == AccountTypeAPIKey || account.Platform != PlatformOpenAI) {
+ markPatchDelete("max_completion_tokens")
}
-
- // Remove unsupported fields (not supported by upstream OpenAI API)
- unsupportedFields := []string{"prompt_cache_retention", "safety_identifier"}
- for _, unsupportedField := range unsupportedFields {
- if _, has := reqBody[unsupportedField]; has {
- delete(reqBody, unsupportedField)
- bodyModified = true
+ for _, unsupportedField := range []string{"prompt_cache_retention", "safety_identifier"} {
+ if gjson.GetBytes(body, unsupportedField).Exists() {
markPatchDelete(unsupportedField)
}
}
}
-
- // 仅在 WSv2 模式保留 previous_response_id,其他模式(HTTP/WSv1)统一过滤。
- // 注意:该规则同样适用于 Codex CLI 请求,避免 WSv1 向上游透传不支持字段。
- if wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
- if _, has := reqBody["previous_response_id"]; has {
- delete(reqBody, "previous_response_id")
- bodyModified = true
- markPatchDelete("previous_response_id")
+ if wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 && gjson.GetBytes(body, "previous_response_id").Exists() {
+ markPatchDelete("previous_response_id")
+ }
+ if openAIRequestBodyMayContainEmptyBase64InputImage(body) {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ if sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(decoded) {
+ markDecodedModified()
}
}
- if sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody) {
- bodyModified = true
- disablePatch()
- }
-
- // Apply OpenAI fast policy (参照 Claude BetaPolicy 的 fast-mode 过滤):
- // 针对 body 的 service_tier 字段("priority" 即 fast,"flex"),按策略
- // 执行 filter(删除字段)或 block(拒绝请求)。对 gpt-5.5 等模型屏蔽
- // fast 时在此生效。
- //
- // 注意:
- // 1. 此处统一使用 upstreamModel(已经过 GetMappedModel +
- // normalizeOpenAIModelForUpstream + Codex OAuth normalize),与
- // chat-completions / messages 入口保持一致,避免不同入口因为模型
- // 维度不同而出现 whitelist 命中差异。
- // 2. action=pass 时也要把 raw "fast" 归一化为 "priority" 写回 body,
- // 否则 native /responses 入口透传 "fast" 给上游会被拒。chat-
- // completions 入口由 normalizeResponsesBodyServiceTier 完成同一
- // 行为,这里手工实现等效逻辑。
- if rawTier, ok := reqBody["service_tier"].(string); ok {
+ if rawTier := requestView.ServiceTier; rawTier != "" {
if normTier := normalizedOpenAIServiceTierValue(rawTier); normTier != "" {
action, errMsg := s.evaluateOpenAIFastPolicy(ctx, account, upstreamModel, normTier)
switch action {
@@ -2478,46 +2681,51 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
writeOpenAIFastPolicyBlockedResponse(c, blocked)
return nil, blocked
case BetaPolicyActionFilter:
- delete(reqBody, "service_tier")
- bodyModified = true
- disablePatch()
+ markPatchDelete("service_tier")
default:
- // pass:若客户端传的是别名 "fast",归一化为 "priority"
- // 后写回 body,确保上游收到的是其能识别的规范值。
if normTier != rawTier {
- reqBody["service_tier"] = normTier
- bodyModified = true
markPatchSet("service_tier", normTier)
}
}
}
}
- if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) && !imageGenerationAllowed {
- MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalFeatureGate)
- c.JSON(http.StatusForbidden, gin.H{
- "error": gin.H{
- "type": "permission_error",
- "message": ImageGenerationPermissionMessage(),
- },
- })
- return nil, errors.New("image generation disabled for group")
+ if bodyModified {
+ if requestView.HasPatches() {
+ if patchedBody, patchErr := requestView.ApplyPatches(); patchErr == nil {
+ body = patchedBody
+ requestView = newOpenAIRequestView(body)
+ reqBody = nil
+ bodyModified = false
+ }
+ }
+ if bodyModified {
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ var marshalErr error
+ body, marshalErr = marshalOpenAIUpstreamJSON(decoded)
+ if marshalErr != nil {
+ return nil, fmt.Errorf("serialize request body: %w", marshalErr)
+ }
+ requestView = newOpenAIRequestView(body)
+ }
}
imageBillingModel := ""
imageSizeTier := ""
imageInputSize := ""
- if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) {
+ if imageIntent {
+ var imageCfg OpenAIResponsesImageBillingConfig
var imageCfgErr error
- imageCfg, imageCfgErr := resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, billingModel)
+ if reqBody != nil {
+ imageCfg, imageCfgErr = resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, billingModel)
+ } else {
+ imageCfg, imageCfgErr = resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, billingModel)
+ }
if imageCfgErr != nil {
setOpsUpstreamError(c, http.StatusBadRequest, imageCfgErr.Error(), "")
- c.JSON(http.StatusBadRequest, gin.H{
- "error": gin.H{
- "type": "invalid_request_error",
- "message": imageCfgErr.Error(),
- "param": "size",
- },
- })
+ c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": imageCfgErr.Error(), "param": "size"}})
return nil, imageCfgErr
}
imageBillingModel = imageCfg.Model
@@ -2525,29 +2733,6 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
imageInputSize = imageCfg.InputSize
}
- // Re-serialize body only if modified
- if bodyModified {
- serializedByPatch := false
- if !patchDisabled && patchHasOp {
- var patchErr error
- if patchDelete {
- body, patchErr = sjson.DeleteBytes(body, patchPath)
- } else {
- body, patchErr = sjson.SetBytes(body, patchPath, patchValue)
- }
- if patchErr == nil {
- serializedByPatch = true
- }
- }
- if !serializedByPatch {
- var marshalErr error
- body, marshalErr = json.Marshal(reqBody)
- if marshalErr != nil {
- return nil, fmt.Errorf("serialize request body: %w", marshalErr)
- }
- }
- }
-
// Get access token
token, _, err := s.GetAccessToken(ctx, account)
if err != nil {
@@ -2556,12 +2741,10 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
// 命中 WS 时仅走 WebSocket Mode;不再自动回退 HTTP。
if wsDecision.Transport == OpenAIUpstreamTransportResponsesWebsocketV2 {
- wsReqBody := reqBody
- if len(reqBody) > 0 {
- wsReqBody = make(map[string]any, len(reqBody))
- for k, v := range reqBody {
- wsReqBody[k] = v
- }
+ // WS 分支需要结构化 payload 与重连恢复,命中后再触发 full-map decode。
+ wsReqBody, err := ensureReqBody()
+ if err != nil {
+ return nil, err
}
_, hasPreviousResponseID := wsReqBody["previous_response_id"]
logOpenAIWSModeDebug(
@@ -2813,7 +2996,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
// Handle error response
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
@@ -2821,8 +3004,12 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
upstreamCode := extractUpstreamErrorCode(respBody)
if !httpInvalidEncryptedContentRetryTried && resp.StatusCode == http.StatusBadRequest && upstreamCode == "invalid_encrypted_content" {
- if trimOpenAIEncryptedReasoningItems(reqBody) {
- body, err = json.Marshal(reqBody)
+ decoded, decodeErr := ensureReqBody()
+ if decodeErr != nil {
+ return nil, decodeErr
+ }
+ if trimOpenAIEncryptedReasoningItems(decoded) {
+ body, err = marshalOpenAIUpstreamJSON(decoded)
if err != nil {
return nil, fmt.Errorf("serialize invalid_encrypted_content retry body: %w", err)
}
@@ -2852,20 +3039,21 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
Detail: upstreamDetail,
})
- s.handleFailoverSideEffects(ctx, resp, account)
+ s.handleFailoverSideEffects(ctx, resp, account, upstreamModel)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && (isPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
+ RetryableOnSameAccount: account.IsPoolMode() && (account.IsPoolModeRetryableStatus(resp.StatusCode) || isOpenAITransientProcessingError(resp.StatusCode, upstreamMsg, respBody)),
}
}
- return s.handleErrorResponse(ctx, resp, c, account, body)
+ return s.handleErrorResponse(ctx, resp, c, account, body, billingModel)
}
defer func() { _ = resp.Body.Close() }()
- reasoningEffort := extractOpenAIReasoningEffort(reqBody, originalModel)
- serviceTier := extractOpenAIServiceTier(reqBody)
- releaseOpenAIParsedRequestBody(c)
+ reasoningEffort := extractOpenAIReasoningEffortFromBody(body, originalModel)
+ serviceTier := extractOpenAIServiceTierFromBody(body)
+ // 上游接受后只保留计费需要的标量,避免响应处理期间继续保活完整 input/tools map。
+ reqBody = nil
// Handle normal response
var usage *OpenAIUsage
@@ -3330,7 +3518,7 @@ func (s *OpenAIGatewayService) handleFailoverErrorResponsePassthrough(
account *Account,
requestBody []byte,
) error {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
@@ -3344,7 +3532,8 @@ func (s *OpenAIGatewayService) handleFailoverErrorResponsePassthrough(
}
setOpsUpstreamError(c, resp.StatusCode, upstreamMsg, upstreamDetail)
logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body)
- _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body)
+ reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody)
+ _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel)
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
@@ -3371,7 +3560,7 @@ func (s *OpenAIGatewayService) handleErrorResponsePassthrough(
account *Account,
requestBody []byte,
) error {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
@@ -3387,7 +3576,8 @@ func (s *OpenAIGatewayService) handleErrorResponsePassthrough(
logOpenAIInstructionsRequiredDebug(ctx, c, account, resp.StatusCode, upstreamMsg, requestBody, body)
// 透传模式保留原始上游错误响应,但运行态账号状态仍需更新,
// 避免粘性路由继续复用刚被限流的账号。
- _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body)
+ reqModel, _, _ := extractOpenAIRequestMetaFromBody(requestBody)
+ _ = s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel)
appendOpsUpstreamError(c, OpsUpstreamErrorEvent{
Platform: account.Platform,
AccountID: account.ID,
@@ -4061,8 +4251,9 @@ func (s *OpenAIGatewayService) handleErrorResponse(
c *gin.Context,
account *Account,
requestBody []byte,
+ requestedModel ...string,
) (*OpenAIForwardResult, error) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
@@ -4137,7 +4328,14 @@ func (s *OpenAIGatewayService) handleErrorResponse(
}
// Handle upstream error (mark account status)
- shouldDisable := s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body)
+ var reqModel string
+ if len(requestedModel) > 0 {
+ reqModel = strings.TrimSpace(requestedModel[0])
+ }
+ if reqModel == "" {
+ reqModel, _, _ = extractOpenAIRequestMetaFromBody(requestBody)
+ }
+ shouldDisable := s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, body, reqModel)
kind := "http_error"
if shouldDisable {
kind = "failover"
@@ -4156,7 +4354,7 @@ func (s *OpenAIGatewayService) handleErrorResponse(
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: body,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
@@ -4214,8 +4412,9 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse(
c *gin.Context,
account *Account,
writeError compatErrorWriter,
+ requestedModel ...string,
) (*OpenAIForwardResult, error) {
- body, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ body := s.readUpstreamErrorBody(resp)
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(body))
if upstreamMsg == "" {
@@ -4269,8 +4468,12 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse(
}
// Track rate limits and decide whether to trigger secondary failover.
+ var modelForCooldown string
+ if len(requestedModel) > 0 {
+ modelForCooldown = requestedModel[0]
+ }
shouldDisable := s.handleOpenAIAccountUpstreamError(
- c.Request.Context(), account, resp.StatusCode, resp.Header, body,
+ c.Request.Context(), account, resp.StatusCode, resp.Header, body, modelForCooldown,
)
kind := "http_error"
if shouldDisable {
@@ -4290,7 +4493,7 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse(
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: body,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
@@ -4435,6 +4638,9 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp
}
needModelReplace := originalModel != mappedModel
+ streamOutputAccumulator := apicompat.NewBufferedResponseAccumulator()
+ streamImageOutputs := make([]json.RawMessage, 0, 1)
+ streamSeenImages := make(map[string]struct{})
resultWithUsage := func() *openaiStreamingResult {
return &openaiStreamingResult{
usage: usage,
@@ -4513,13 +4719,6 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp
}
// Extract data from SSE line (supports both "data: " and "data:" formats)
if data, ok := extractOpenAISSEDataLine(line); ok {
-
- // Replace model in response if needed.
- // Fast path: most events do not contain model field values.
- if needModelReplace && mappedModel != "" && strings.Contains(data, mappedModel) {
- line = s.replaceModelInSSELine(line, mappedModel, originalModel)
- }
-
dataBytes := []byte(data)
if openAIStreamEventIsTerminal(data) {
sawTerminalEvent = true
@@ -4545,6 +4744,26 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp
line = "data: " + data
eventType = strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String())
}
+ if imageOutput, ok := extractImageGenerationOutputFromSSEData(dataBytes, streamSeenImages); ok {
+ streamImageOutputs = append(streamImageOutputs, imageOutput)
+ }
+ if responsesStreamEventMayContributeToOutput(eventType) {
+ var streamEvent apicompat.ResponsesStreamEvent
+ if err := json.Unmarshal(dataBytes, &streamEvent); err == nil {
+ streamOutputAccumulator.ProcessEvent(&streamEvent)
+ }
+ }
+ if normalizedData, normalized := normalizeResponsesStreamingTerminalOutput(dataBytes, streamOutputAccumulator, streamImageOutputs); normalized {
+ dataBytes = normalizedData
+ data = string(normalizedData)
+ line = "data: " + data
+ eventType = strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String())
+ }
+ // Replace model in response if needed.
+ // Fast path: most events do not contain model field values.
+ if needModelReplace && mappedModel != "" && strings.Contains(line, mappedModel) {
+ line = s.replaceModelInSSELine(line, mappedModel, originalModel)
+ }
startsClientOutput := forceFlushFailedEvent || openAIStreamDataStartsClientOutput(data, eventType)
// 写入客户端(客户端断开后继续 drain 上游)
@@ -4842,7 +5061,7 @@ func (s *OpenAIGatewayService) parseSSEUsageBytes(data []byte, usage *OpenAIUsag
return
}
eventType := gjson.GetBytes(data, "type").String()
- if eventType != "response.completed" && eventType != "response.done" &&
+ if eventType != "response.completed" && eventType != "response.done" && eventType != "response.failed" &&
eventType != "response.incomplete" && eventType != "response.cancelled" && eventType != "response.canceled" {
return
}
@@ -4903,20 +5122,22 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r
if isEventStreamResponse(resp.Header) {
return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel)
}
+ bodyLooksLikeSSE := bytes.Contains(body, []byte("data:")) || bytes.Contains(body, []byte("event:"))
+
// For OAuth accounts, also fall back to a body-content heuristic because
// the upstream may omit the Content-Type header while still sending SSE.
// This heuristic is NOT applied to API-key accounts to avoid false
// positives on JSON responses that coincidentally contain "data:" or
// "event:" in their text content.
- if account.Type == AccountTypeOAuth {
- bodyLooksLikeSSE := bytes.Contains(body, []byte("data:")) || bytes.Contains(body, []byte("event:"))
- if bodyLooksLikeSSE {
- return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel)
- }
+ if account.Type == AccountTypeOAuth && bodyLooksLikeSSE {
+ return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel)
}
usageValue, usageOK := extractOpenAIUsageFromJSONBytes(body)
if !usageOK {
+ if bodyLooksLikeSSE {
+ return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel)
+ }
return nil, fmt.Errorf("parse response: invalid json response")
}
usage := &usageValue
@@ -5078,6 +5299,45 @@ func extractCodexFinalResponse(body string) ([]byte, bool) {
return nil, false
}
+func normalizeResponsesStreamingTerminalOutput(data []byte, acc *apicompat.BufferedResponseAccumulator, imageOutputs []json.RawMessage) ([]byte, bool) {
+ eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String())
+ switch eventType {
+ case "response.completed", "response.done", "response.incomplete", "response.cancelled", "response.canceled":
+ default:
+ return data, false
+ }
+
+ output := gjson.GetBytes(data, "response.output")
+ hasAccumulatedOutput := (acc != nil && acc.HasContent()) || len(imageOutputs) > 0
+ if output.Exists() && output.IsArray() {
+ if len(output.Array()) > 0 || !hasAccumulatedOutput {
+ return data, false
+ }
+ }
+
+ outputJSON := []byte("[]")
+ if reconstructed, ok := buildResponsesOutputJSON(acc, imageOutputs); ok {
+ outputJSON = reconstructed
+ }
+ updated, err := sjson.SetRawBytes(data, "response.output", outputJSON)
+ if err != nil {
+ return data, false
+ }
+ return updated, true
+}
+
+func responsesStreamEventMayContributeToOutput(eventType string) bool {
+ switch eventType {
+ case "response.output_text.delta",
+ "response.output_item.added",
+ "response.function_call_arguments.delta",
+ "response.reasoning_summary_text.delta":
+ return true
+ default:
+ return false
+ }
+}
+
// reconstructResponseOutputFromSSE scans raw SSE body text for delta events and
// returns a JSON-encoded output array reconstructed from accumulated deltas.
// Returns (nil, false) if no content was found in deltas.
@@ -5089,17 +5349,23 @@ func reconstructResponseOutputFromSSE(bodyText string) ([]byte, bool) {
if imageOutput, ok := extractImageGenerationOutputFromSSEData(data, seenImages); ok {
imageOutputs = append(imageOutputs, imageOutput)
}
- var event apicompat.ResponsesStreamEvent
- if err := json.Unmarshal(data, &event); err == nil {
- acc.ProcessEvent(&event)
+ eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String())
+ if responsesStreamEventMayContributeToOutput(eventType) {
+ var event apicompat.ResponsesStreamEvent
+ if err := json.Unmarshal(data, &event); err == nil {
+ acc.ProcessEvent(&event)
+ }
}
})
- if !acc.HasContent() && len(imageOutputs) == 0 {
+ return buildResponsesOutputJSON(acc, imageOutputs)
+}
+
+func buildResponsesOutputJSON(acc *apicompat.BufferedResponseAccumulator, imageOutputs []json.RawMessage) ([]byte, bool) {
+ if (acc == nil || !acc.HasContent()) && len(imageOutputs) == 0 {
return nil, false
}
-
var output []json.RawMessage
- if acc.HasContent() {
+ if acc != nil && acc.HasContent() {
outputJSON, err := json.Marshal(acc.BuildOutput())
if err == nil {
_ = json.Unmarshal(outputJSON, &output)
@@ -5523,6 +5789,11 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
durationMs := int(result.Duration.Milliseconds())
accountRateMultiplier := account.BillingRateMultiplier()
requestID := resolveUsageBillingRequestID(ctx, result.RequestID)
+ if result.OpenAIWSMode {
+ if upstreamRequestID := strings.TrimSpace(result.RequestID); upstreamRequestID != "" {
+ requestID = upstreamRequestID
+ }
+ }
// 确定 RequestedModel(渠道映射前的原始模型)
requestedModel := result.Model
@@ -6002,15 +6273,163 @@ func deriveOpenAIReasoningEffortFromModel(model string) string {
return normalizeOpenAIReasoningEffort(parts[len(parts)-1])
}
-func extractOpenAIRequestMetaFromBody(body []byte) (model string, stream bool, promptCacheKey string) {
- if len(body) == 0 {
- return "", false, ""
- }
+type openAIRequestView struct {
+ body []byte
+ Model string
+ Stream bool
+ PromptCacheKey string
+ PreviousResponseID string
+ ServiceTier string
+ ReasoningEffort string
+ patches []openAIRequestPatch
+ patchesDisabled bool
+}
- model = strings.TrimSpace(gjson.GetBytes(body, "model").String())
- stream = gjson.GetBytes(body, "stream").Bool()
- promptCacheKey = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String())
- return model, stream, promptCacheKey
+type openAIRequestPatch struct {
+ path string
+ delete bool
+ value any
+}
+
+func newOpenAIRequestView(body []byte) openAIRequestView {
+ if len(body) == 0 {
+ return openAIRequestView{}
+ }
+ return openAIRequestView{
+ body: body,
+ Model: strings.TrimSpace(gjson.GetBytes(body, "model").String()),
+ Stream: gjson.GetBytes(body, "stream").Bool(),
+ PromptCacheKey: strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()),
+ PreviousResponseID: strings.TrimSpace(gjson.GetBytes(body, "previous_response_id").String()),
+ ServiceTier: strings.TrimSpace(gjson.GetBytes(body, "service_tier").String()),
+ ReasoningEffort: strings.TrimSpace(gjson.GetBytes(body, "reasoning.effort").String()),
+ }
+}
+
+// Decode 保留阶段一既有 full-map 行为;后续阶段会把调用点下沉到复杂分支。
+func (v openAIRequestView) Decode(c *gin.Context) (map[string]any, error) {
+ return getOpenAIRequestBodyMap(c, v.body)
+}
+
+func (v *openAIRequestView) MarkPatchSet(path string, value any) {
+ if v == nil || v.patchesDisabled {
+ return
+ }
+ path = strings.TrimSpace(path)
+ if !isSimpleOpenAIRequestPatchPath(path) {
+ v.DisablePatches()
+ return
+ }
+ v.patches = append(v.patches, openAIRequestPatch{path: path, value: value})
+}
+
+func (v *openAIRequestView) MarkPatchDelete(path string) {
+ if v == nil || v.patchesDisabled {
+ return
+ }
+ path = strings.TrimSpace(path)
+ if !isSimpleOpenAIRequestPatchPath(path) {
+ v.DisablePatches()
+ return
+ }
+ v.patches = append(v.patches, openAIRequestPatch{path: path, delete: true})
+}
+
+func isSimpleOpenAIRequestPatchPath(path string) bool {
+ if path == "" || strings.ContainsRune(path, '\\') {
+ return false
+ }
+ for _, part := range strings.Split(path, ".") {
+ if strings.TrimSpace(part) == "" {
+ return false
+ }
+ }
+ return true
+}
+
+func (v *openAIRequestView) DisablePatches() {
+ if v == nil {
+ return
+ }
+ v.patchesDisabled = true
+ v.patches = nil
+}
+
+func (v openAIRequestView) HasPatches() bool {
+ return !v.patchesDisabled && len(v.patches) > 0
+}
+
+func (v openAIRequestView) ApplyPatches() ([]byte, error) {
+ if v.patchesDisabled || len(v.patches) == 0 {
+ return nil, errors.New("openai request patches disabled")
+ }
+ body := v.body
+ for _, patch := range v.patches {
+ var err error
+ if patch.delete {
+ body, err = sjson.DeleteBytes(body, patch.path)
+ } else {
+ body, err = sjson.SetBytes(body, patch.path, patch.value)
+ }
+ if err != nil {
+ return nil, err
+ }
+ }
+ return body, nil
+}
+
+func setOpenAIRequestMapPath(reqBody map[string]any, path string, value any) {
+ path = strings.TrimSpace(path)
+ if reqBody == nil || path == "" {
+ return
+ }
+ parts := strings.Split(path, ".")
+ current := reqBody
+ for _, part := range parts[:len(parts)-1] {
+ part = strings.TrimSpace(part)
+ if part == "" {
+ return
+ }
+ next, _ := current[part].(map[string]any)
+ if next == nil {
+ next = map[string]any{}
+ current[part] = next
+ }
+ current = next
+ }
+ last := strings.TrimSpace(parts[len(parts)-1])
+ if last != "" {
+ current[last] = value
+ }
+}
+
+func deleteOpenAIRequestMapPath(reqBody map[string]any, path string) {
+ path = strings.TrimSpace(path)
+ if reqBody == nil || path == "" {
+ return
+ }
+ parts := strings.Split(path, ".")
+ current := reqBody
+ for _, part := range parts[:len(parts)-1] {
+ part = strings.TrimSpace(part)
+ if part == "" {
+ return
+ }
+ next, _ := current[part].(map[string]any)
+ if next == nil {
+ return
+ }
+ current = next
+ }
+ last := strings.TrimSpace(parts[len(parts)-1])
+ if last != "" {
+ delete(current, last)
+ }
+}
+
+func extractOpenAIRequestMetaFromBody(body []byte) (model string, stream bool, promptCacheKey string) {
+ view := newOpenAIRequestView(body)
+ return view.Model, view.Stream, view.PromptCacheKey
}
// normalizeOpenAIPassthroughOAuthBody 将透传 OAuth 请求体收敛为旧链路关键行为:
@@ -6461,8 +6880,84 @@ func buildOpenAIFastPolicyBlockedWSEvent(err *OpenAIFastBlockedError) []byte {
return payload
}
+func openAIRequestBodyMayContainImageInput(body []byte) bool {
+ if len(body) == 0 {
+ return false
+ }
+ input := gjson.GetBytes(body, "input")
+ messages := gjson.GetBytes(body, "messages.#-1")
+ return openAIJSONValueMayContainImageInput(input) || openAIJSONValueMayContainImageInput(messages)
+}
+
+func openAIJSONValueMayContainImageInput(value gjson.Result) bool {
+ if !value.Exists() {
+ return false
+ }
+ if value.IsArray() {
+ found := false
+ value.ForEach(func(_, item gjson.Result) bool {
+ if openAIJSONValueMayContainImageInput(item) {
+ found = true
+ return false
+ }
+ return true
+ })
+ return found
+ }
+ if value.IsObject() {
+ if strings.TrimSpace(value.Get("type").String()) == "input_image" || value.Get("image_url").Exists() {
+ return true
+ }
+ return openAIJSONValueMayContainImageInput(value.Get("content"))
+ }
+ return false
+}
+
+func openAIRequestBodyMayContainEmptyBase64InputImage(body []byte) bool {
+ if len(body) == 0 || !openAIRequestBodyMayContainInputImageToken(body) {
+ return false
+ }
+ input := gjson.GetBytes(body, "input")
+ if !input.Exists() {
+ return false
+ }
+ return openAIJSONValueMayContainEmptyBase64InputImage(input)
+}
+
+func openAIRequestBodyMayContainInputImageToken(body []byte) bool {
+ if bytes.Contains(body, []byte("input_image")) {
+ return true
+ }
+ // JSON 字符串任意字符都可能被 unicode escape,遇到 \u 时交给 gjson 解码后的结构扫描兜底。
+ return bytes.Contains(body, []byte("\\u"))
+}
+
+func openAIJSONValueMayContainEmptyBase64InputImage(value gjson.Result) bool {
+ if !value.Exists() {
+ return false
+ }
+ if value.IsArray() {
+ found := false
+ value.ForEach(func(_, item gjson.Result) bool {
+ if openAIJSONValueMayContainEmptyBase64InputImage(item) {
+ found = true
+ return false
+ }
+ return true
+ })
+ return found
+ }
+ if value.IsObject() {
+ if strings.TrimSpace(value.Get("type").String()) == "input_image" && isEmptyBase64DataURI(value.Get("image_url").String()) {
+ return true
+ }
+ return openAIJSONValueMayContainEmptyBase64InputImage(value.Get("content"))
+ }
+ return false
+}
+
func sanitizeEmptyBase64InputImagesInOpenAIBody(body []byte) ([]byte, bool, error) {
- if len(body) == 0 || !bytes.Contains(body, []byte(`"image_url"`)) || !bytes.Contains(body, []byte(`base64,`)) {
+ if !openAIRequestBodyMayContainEmptyBase64InputImage(body) {
return body, false, nil
}
@@ -6473,7 +6968,7 @@ func sanitizeEmptyBase64InputImagesInOpenAIBody(body []byte) ([]byte, bool, erro
if !sanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(reqBody) {
return body, false, nil
}
- normalized, err := json.Marshal(reqBody)
+ normalized, err := marshalOpenAIUpstreamJSON(reqBody)
if err != nil {
return body, false, fmt.Errorf("serialize sanitized request body: %w", err)
}
@@ -6578,32 +7073,14 @@ func isEmptyBase64DataURI(raw string) bool {
return strings.TrimSpace(strings.TrimPrefix(rest, "base64,")) == ""
}
-func getOpenAIRequestBodyMap(c *gin.Context, body []byte) (map[string]any, error) {
- if c != nil {
- if cached, ok := c.Get(OpenAIParsedRequestBodyKey); ok {
- if reqBody, ok := cached.(map[string]any); ok && reqBody != nil {
- return reqBody, nil
- }
- }
- }
-
+func getOpenAIRequestBodyMap(_ *gin.Context, body []byte) (map[string]any, error) {
var reqBody map[string]any
if err := json.Unmarshal(body, &reqBody); err != nil {
return nil, fmt.Errorf("parse request: %w", err)
}
- if c != nil {
- c.Set(OpenAIParsedRequestBodyKey, reqBody)
- }
return reqBody, nil
}
-func releaseOpenAIParsedRequestBody(c *gin.Context) {
- if c == nil {
- return
- }
- delete(c.Keys, OpenAIParsedRequestBodyKey)
-}
-
func extractOpenAIReasoningEffort(reqBody map[string]any, requestedModel string) *string {
if value, present := getOpenAIReasoningEffortFromReqBody(reqBody); present {
if value == "" {
diff --git a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go
index 17a874ea..10d58654 100644
--- a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go
+++ b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go
@@ -18,7 +18,7 @@ type stubCodexRestrictionDetector struct {
result CodexClientRestrictionDetectionResult
}
-func (s *stubCodexRestrictionDetector) Detect(_ *gin.Context, _ *Account) CodexClientRestrictionDetectionResult {
+func (s *stubCodexRestrictionDetector) Detect(_ *gin.Context, _ *Account, _ []string) CodexClientRestrictionDetectionResult {
return s.result
}
@@ -52,7 +52,7 @@ func TestOpenAIGatewayService_GetCodexClientRestrictionDetector(t *testing.T) {
c.Request.Header.Set("User-Agent", "curl/8.0")
account := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth, Extra: map[string]any{"codex_cli_only": true}}
- result := got.Detect(c, account)
+ result := got.Detect(c, account, nil)
require.True(t, result.Enabled)
require.True(t, result.Matched)
require.Equal(t, CodexClientRestrictionReasonForceCodexCLI, result.Reason)
diff --git a/backend/internal/service/openai_gateway_service_codex_snapshot_test.go b/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
index 654dd4ca..27208b58 100644
--- a/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
+++ b/backend/internal/service/openai_gateway_service_codex_snapshot_test.go
@@ -104,6 +104,39 @@ func TestBuildCodexUsageExtraUpdates_UsesSnapshotUpdatedAt(t *testing.T) {
}
}
+// TestBuildCodexUsageExtraUpdates_FreshAccountUsedPercentNotInverted_Issue2994 locks in the
+// canonical "used %" semantics for the 5h window. A fresh account reports a tiny
+// secondary-used-percent (~1%); the stored codex_5h_used_percent must equal that value
+// directly and must NOT be inverted to ~99%. Regression guard for issue #2994 / the reverted
+// commit b65dde63 (PR #2918), which applied `100 - used` and made fresh accounts look
+// exhausted, tripping auto-pause and excluding them from scheduling.
+func TestBuildCodexUsageExtraUpdates_FreshAccountUsedPercentNotInverted_Issue2994(t *testing.T) {
+ secondaryUsed := 1.0 // 5h window: barely used
+ secondaryWindow := 300
+ primaryUsed := 2.0 // 7d window: barely used
+ primaryWindow := 10080
+
+ snapshot := &OpenAICodexUsageSnapshot{
+ PrimaryUsedPercent: &primaryUsed,
+ PrimaryWindowMinutes: &primaryWindow,
+ SecondaryUsedPercent: &secondaryUsed,
+ SecondaryWindowMinutes: &secondaryWindow,
+ UpdatedAt: "2026-02-16T10:00:00Z",
+ }
+
+ updates := buildCodexUsageExtraUpdates(snapshot, time.Date(2026, 2, 16, 10, 0, 0, 0, time.UTC))
+ if updates == nil {
+ t.Fatal("expected non-nil updates")
+ }
+
+ if got := updates["codex_5h_used_percent"]; got != 1.0 {
+ t.Fatalf("codex_5h_used_percent = %v, want 1.0 (direct used%%, NOT inverted to 99)", got)
+ }
+ if got := updates["codex_7d_used_percent"]; got != 2.0 {
+ t.Fatalf("codex_7d_used_percent = %v, want 2.0 (direct used%%, NOT inverted to 98)", got)
+ }
+}
+
func TestBuildCodexUsageExtraUpdates_FallbackToNowWhenUpdatedAtInvalid(t *testing.T) {
primaryUsed := 15.0
primaryReset := 30
diff --git a/backend/internal/service/openai_gateway_service_hotpath_test.go b/backend/internal/service/openai_gateway_service_hotpath_test.go
index 234dee00..92a0d1ac 100644
--- a/backend/internal/service/openai_gateway_service_hotpath_test.go
+++ b/backend/internal/service/openai_gateway_service_hotpath_test.go
@@ -1,14 +1,533 @@
package service
import (
+ "context"
"encoding/json"
+ "io"
+ "net/http"
"net/http/httptest"
+ "strings"
"testing"
+ "github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
)
+func TestOpenAIRequestView_ExtractsRawScalars(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":" gpt-5 ","stream":true,"prompt_cache_key":" ses-1 ","previous_response_id":" resp-1 ","service_tier":" fast ","reasoning":{"effort":" medium "}}`))
+
+ require.Equal(t, "gpt-5", view.Model)
+ require.True(t, view.Stream)
+ require.Equal(t, "ses-1", view.PromptCacheKey)
+ require.Equal(t, "resp-1", view.PreviousResponseID)
+ require.Equal(t, "fast", view.ServiceTier)
+ require.Equal(t, "medium", view.ReasoningEffort)
+}
+
+func TestOpenAIRequestView_DecodeKeepsFullMapBehavior(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5","stream":true,"input":[{"type":"message","content":"hi"}]}`))
+
+ reqBody, err := view.Decode(nil)
+ require.NoError(t, err)
+ require.Equal(t, "gpt-5", reqBody["model"])
+ require.IsType(t, []any{}, reqBody["input"])
+}
+
+func TestOpenAIRequestView_ApplyPatches(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5","previous_response_id":"resp_1","reasoning":{"effort":"minimal"},"input":[{"type":"message","content":"hi"}]}`))
+ view.MarkPatchSet("model", "gpt-5.1")
+ view.MarkPatchDelete("previous_response_id")
+ view.MarkPatchSet("reasoning.effort", "none")
+
+ patched, err := view.ApplyPatches()
+ require.NoError(t, err)
+ require.JSONEq(t, `{"model":"gpt-5.1","reasoning":{"effort":"none"},"input":[{"type":"message","content":"hi"}]}`, string(patched))
+}
+
+func TestOpenAIRequestView_RejectsEscapedPatchPath(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"metadata":{"user.id":"old"}}`))
+ view.MarkPatchSet(`metadata.user\.id`, "new")
+
+ require.False(t, view.HasPatches())
+ _, err := view.ApplyPatches()
+ require.Error(t, err)
+}
+
+func TestOpenAIRequestView_ApplyPatchesDisabled(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5"}`))
+ view.MarkPatchSet("model", "gpt-5.1")
+ view.DisablePatches()
+
+ _, err := view.ApplyPatches()
+ require.Error(t, err)
+}
+
+func TestOpenAIRequestView_HasPatches(t *testing.T) {
+ view := newOpenAIRequestView([]byte(`{"model":"gpt-5"}`))
+ require.False(t, view.HasPatches())
+
+ view.MarkPatchSet("model", "gpt-5.1")
+ require.True(t, view.HasPatches())
+
+ view.DisablePatches()
+ require.False(t, view.HasPatches())
+}
+
+func TestOpenAIGatewayService_Forward_HTTPPatchPathKeepsLargeInputRaw(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(
+ `{"usage":{"input_tokens":1,"output_tokens":2,"input_tokens_details":{"cached_tokens":0}}}`,
+ )),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 1,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"reasoning":{"effort":"minimal"},"input":[{"type":"message","content":[{"type":"input_text","text":"hi","nonce":9007199254740993}]}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.NotNil(t, upstream.lastReq)
+ require.JSONEq(t, `{"model":"gpt-5","stream":false,"reasoning":{"effort":"none"},"instructions":"You are a helpful coding assistant.","input":[{"type":"message","content":[{"type":"input_text","text":"hi","nonce":9007199254740993}]}]}`, string(upstream.lastBody))
+ require.Equal(t, "9007199254740993", gjson.GetBytes(upstream.lastBody, "input.0.content.0.nonce").Raw)
+}
+
+func TestOpenAIGatewayService_Forward_DecodedMutationKeepsLaterFieldDeletes(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 2,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5.4","stream":false,"max_completion_tokens":12,"tools":[{"type":"image_generation","format":"png"}],"input":[{"type":"message","content":"draw"}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.False(t, gjson.GetBytes(upstream.lastBody, "max_completion_tokens").Exists())
+ require.False(t, gjson.GetBytes(upstream.lastBody, "tools.0.format").Exists())
+ require.Equal(t, "png", gjson.GetBytes(upstream.lastBody, "tools.0.output_format").String())
+}
+
+func TestOpenAIGatewayService_Forward_MappedImageModelUsesImageGate(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 3,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ "model_mapping": map[string]any{"draw-alias": "gpt-image-2"},
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ c.Set("api_key", &APIKey{Group: &Group{AllowImageGeneration: false}})
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"draw-alias","stream":false,"input":"draw"}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.Error(t, err)
+ require.Nil(t, result)
+ require.Nil(t, upstream.lastReq)
+ require.Equal(t, http.StatusForbidden, rec.Code)
+}
+
+func TestOpenAIGatewayService_Forward_TextDataImageDoesNotForceMapMarshal(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 4,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"input":[{"type":"message","content":[{"type":"input_text","text":"literal data:image/png;base64, only","nonce":1e1000000}]}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, "1e1000000", gjson.GetBytes(upstream.lastBody, "input.0.content.0.nonce").Raw)
+}
+
+func TestOpenAIGatewayService_Forward_ImageToolBillingDoesNotForceFullDecode(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(
+ `{"output":[{"id":"ig_1","type":"image_generation_call","result":"final-image"}],"usage":{"input_tokens":1,"output_tokens":2}}`,
+ )),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 9,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2","size":"2048x1152"}],"input":[{"type":"message","content":[{"type":"input_text","text":"draw","nonce":1e1000000}]}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, "1e1000000", gjson.GetBytes(upstream.lastBody, "input.0.content.0.nonce").Raw)
+ require.Equal(t, 1, result.ImageCount)
+ require.Equal(t, "2K", result.ImageSize)
+ require.Equal(t, "gpt-image-2", result.BillingModel)
+}
+
+func TestOpenAIGatewayService_Forward_ImageToolWithImageOnlyModelIsNormalized(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 11,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-image-2","stream":false,"tools":[{"type":"image_generation","model":"gpt-image-2"}],"input":"draw"}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, openAIImagesResponsesMainModel, gjson.GetBytes(upstream.lastBody, "model").String())
+}
+
+func TestOpenAIGatewayService_Forward_HTTPRetryRecoveryDoesNotDecodeBeforeError(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ responses: []*http.Response{
+ {
+ StatusCode: http.StatusBadRequest,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"error":{"code":"invalid_encrypted_content","type":"invalid_request_error","message":"bad encrypted content"}}`)),
+ },
+ {
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 10,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"input":[{"type":"reasoning","encrypted_content":"gAAA","summary":[{"type":"summary_text","text":"keep me"}]},{"type":"message","content":[{"type":"input_text","text":"hi","nonce":9007199254740993}]}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Len(t, upstream.bodies, 2)
+ require.Equal(t, "gAAA", gjson.GetBytes(upstream.bodies[0], "input.0.encrypted_content").String())
+ require.Equal(t, "9007199254740993", gjson.GetBytes(upstream.bodies[0], "input.1.content.0.nonce").Raw)
+ require.False(t, gjson.GetBytes(upstream.bodies[1], "input.0.encrypted_content").Exists())
+ require.Equal(t, "summary_text", gjson.GetBytes(upstream.bodies[1], "input.0.summary.0.type").String())
+}
+
+func TestOpenAIGatewayService_Forward_CodexSparkRejectsEscapedInputImage(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 5,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5.3-codex-spark","stream":false,"input":[{"type":"input_` + "\\u0069" + `mage","file_id":"file_1"}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.Error(t, err)
+ require.Nil(t, result)
+ require.Nil(t, upstream.lastReq)
+ require.Equal(t, http.StatusBadRequest, rec.Code)
+}
+
+func TestOpenAIGatewayService_Forward_CodexBridgeInjectionSetsImageBilling(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(
+ `{"output":[{"id":"ig_1","type":"image_generation_call","result":"final-image","size":"1024x1024"}],"usage":{"input_tokens":1,"output_tokens":2}}`,
+ )),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Gateway.ForceCodexCLI = true
+ cfg.Gateway.CodexImageGenerationBridgeEnabled = true
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 7,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ c.Set("api_key", &APIKey{Group: &Group{AllowImageGeneration: true}})
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-5","stream":false,"input":"draw if needed"}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, 1, result.ImageCount)
+ require.Equal(t, "2K", result.ImageSize)
+ require.Equal(t, "gpt-image-2", result.BillingModel)
+}
+
+func TestOpenAIGatewayService_Forward_HTTPDeletesPreviousResponseIDWhenPresent(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ account := &Account{
+ ID: 8,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+
+ for _, body := range [][]byte{
+ []byte(`{"model":"gpt-5","stream":false,"previous_response_id":"","input":"hi"}`),
+ []byte(`{"model":"gpt-5","stream":false,"previous_response_id":null,"input":"hi"}`),
+ } {
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.False(t, gjson.GetBytes(upstream.lastBody, "previous_response_id").Exists())
+ }
+}
+
+func TestOpenAIRequestBodyMayContainEmptyBase64InputImageSeesEscapedJSON(t *testing.T) {
+ body := []byte(`{"input":[{"type":"message","content":[{"type":"input_image","image_` + "\\u0075" + `rl":"data:image/png;base64` + "\\u002c" + ` "}]}]}`)
+
+ require.True(t, openAIRequestBodyMayContainEmptyBase64InputImage(body))
+}
+
+func TestOpenAIRequestBodyMayContainEmptyBase64InputImageSeesEscapedImageType(t *testing.T) {
+ body := []byte(`{"input":[{"type":"message","content":[{"type":"input_` + "\\u0069" + `mage","image_url":"data:image/png;base64, "}]}]}`)
+
+ require.True(t, openAIRequestBodyMayContainEmptyBase64InputImage(body))
+}
+
+func TestOpenAIRequestBodyMayContainEmptyBase64InputImageSeesEscapedInputPrefix(t *testing.T) {
+ body := []byte(`{"input":[{"type":"message","content":[{"type":"inp` + "\\u0075" + `t_image","image_url":"data:image/png;base64, "}]}]}`)
+
+ require.True(t, openAIRequestBodyMayContainEmptyBase64InputImage(body))
+}
+
+func TestOpenAIGatewayService_Forward_ImageOnlyModelKeepsSupportedVerbosity(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ upstream := &httpUpstreamRecorder{
+ resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(`{"usage":{"input_tokens":1,"output_tokens":2}}`)),
+ },
+ }
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ svc := &OpenAIGatewayService{cfg: cfg, httpUpstream: upstream}
+ account := &Account{
+ ID: 6,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "https://example.com",
+ },
+ Extra: map[string]any{"use_responses_api": true},
+ }
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ SetOpenAIClientTransport(c, OpenAIClientTransportHTTP)
+
+ body := []byte(`{"model":"gpt-image-2","stream":false,"text":{"verbosity":"low"},"input":"draw"}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, "low", gjson.GetBytes(upstream.lastBody, "text.verbosity").String())
+ require.Equal(t, openAIImagesResponsesMainModel, gjson.GetBytes(upstream.lastBody, "model").String())
+}
+
func TestExtractOpenAIRequestMetaFromBody(t *testing.T) {
tests := []struct {
name string
@@ -106,26 +625,13 @@ func TestExtractOpenAIReasoningEffortFromBody(t *testing.T) {
}
}
-func TestGetOpenAIRequestBodyMap_UsesContextCache(t *testing.T) {
- gin.SetMode(gin.TestMode)
- rec := httptest.NewRecorder()
- c, _ := gin.CreateTestContext(rec)
-
- cached := map[string]any{"model": "cached-model", "stream": true}
- c.Set(OpenAIParsedRequestBodyKey, cached)
-
- got, err := getOpenAIRequestBodyMap(c, []byte(`{invalid-json`))
- require.NoError(t, err)
- require.Equal(t, cached, got)
-}
-
-func TestGetOpenAIRequestBodyMap_ParseErrorWithoutCache(t *testing.T) {
+func TestGetOpenAIRequestBodyMap_ParseError(t *testing.T) {
_, err := getOpenAIRequestBodyMap(nil, []byte(`{invalid-json`))
require.Error(t, err)
require.Contains(t, err.Error(), "parse request")
}
-func TestGetOpenAIRequestBodyMap_WriteBackContextCache(t *testing.T) {
+func TestGetOpenAIRequestBodyMap_DoesNotWriteContextCache(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
@@ -133,12 +639,7 @@ func TestGetOpenAIRequestBodyMap_WriteBackContextCache(t *testing.T) {
got, err := getOpenAIRequestBodyMap(c, []byte(`{"model":"gpt-5","stream":true}`))
require.NoError(t, err)
require.Equal(t, "gpt-5", got["model"])
-
- cached, ok := c.Get(OpenAIParsedRequestBodyKey)
- require.True(t, ok)
- cachedMap, ok := cached.(map[string]any)
- require.True(t, ok)
- require.Equal(t, got, cachedMap)
+ require.Empty(t, c.Keys)
}
func TestSanitizeEmptyBase64InputImagesInOpenAIRequestBodyMap(t *testing.T) {
diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go
index ef35aa1a..e57777d2 100644
--- a/backend/internal/service/openai_gateway_service_test.go
+++ b/backend/internal/service/openai_gateway_service_test.go
@@ -1233,6 +1233,85 @@ func TestOpenAIStreamingPreambleKeepaliveUsesDownstreamIdle(t *testing.T) {
require.Contains(t, rec.Body.String(), "response.completed")
}
+func TestOpenAIStreamingNormalizesTerminalOutputFromDeltas(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ cfg := &config.Config{
+ Gateway: config.GatewayConfig{
+ StreamDataIntervalTimeout: 0,
+ StreamKeepaliveInterval: 0,
+ MaxLineSize: defaultMaxLineSize,
+ },
+ }
+ svc := &OpenAIGatewayService{cfg: cfg}
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
+
+ resp := &http.Response{
+ StatusCode: http.StatusOK,
+ Body: io.NopCloser(strings.NewReader(strings.Join([]string{
+ `data: {"type":"response.created","response":{"id":"resp_sdk_parse"}}`,
+ "",
+ `data: {"type":"response.output_text.delta","delta":"pon"}`,
+ "",
+ `data: {"type":"response.output_text.delta","delta":"g"}`,
+ "",
+ `data: {"type":"response.completed","response":{"id":"resp_sdk_parse","status":"completed","output":null,"usage":{"input_tokens":1,"output_tokens":1}}}`,
+ "",
+ }, "\n"))),
+ Header: http.Header{"X-Request-Id": []string{"rid-sdk-parse"}},
+ }
+
+ result, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Name: "acc"}, time.Now(), "model", "model")
+ require.NoError(t, err)
+ require.NotNil(t, result)
+
+ terminalType, terminalPayload, ok := extractOpenAISSETerminalEvent(rec.Body.String())
+ require.True(t, ok)
+ require.Equal(t, "response.completed", terminalType)
+ output := gjson.GetBytes(terminalPayload, "response.output")
+ require.True(t, output.IsArray())
+ require.Len(t, output.Array(), 1)
+ require.Equal(t, "pong", gjson.GetBytes(terminalPayload, "response.output.0.content.0.text").String())
+}
+
+func TestOpenAIStreamingNormalizesTerminalOutputToEmptyArray(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ cfg := &config.Config{
+ Gateway: config.GatewayConfig{
+ StreamDataIntervalTimeout: 0,
+ StreamKeepaliveInterval: 0,
+ MaxLineSize: defaultMaxLineSize,
+ },
+ }
+ svc := &OpenAIGatewayService{cfg: cfg}
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
+
+ resp := &http.Response{
+ StatusCode: http.StatusOK,
+ Body: io.NopCloser(strings.NewReader(strings.Join([]string{
+ `data: {"type":"response.completed","response":{"id":"resp_empty","status":"completed","output":null,"usage":{"input_tokens":1,"output_tokens":0}}}`,
+ "",
+ }, "\n"))),
+ Header: http.Header{"X-Request-Id": []string{"rid-empty-output"}},
+ }
+
+ result, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Name: "acc"}, time.Now(), "model", "model")
+ require.NoError(t, err)
+ require.NotNil(t, result)
+
+ terminalType, terminalPayload, ok := extractOpenAISSETerminalEvent(rec.Body.String())
+ require.True(t, ok)
+ require.Equal(t, "response.completed", terminalType)
+ output := gjson.GetBytes(terminalPayload, "response.output")
+ require.True(t, output.IsArray())
+ require.Len(t, output.Array(), 0)
+}
+
func TestOpenAIStreamingPolicyResponseFailedBeforeOutputPassesThrough(t *testing.T) {
gin.SetMode(gin.TestMode)
cfg := &config.Config{
@@ -2218,6 +2297,12 @@ func TestParseSSEUsage_SelectiveParsing(t *testing.T) {
require.Equal(t, 15, usage.OutputTokens)
require.Equal(t, 4, usage.CacheReadInputTokens)
+ // failed 事件在部分上游路径也会携带已消耗 usage,应与 WS/passthrough 保持一致
+ svc.parseSSEUsage(`{"type":"response.failed","response":{"usage":{"input_tokens":17,"output_tokens":19,"input_tokens_details":{"cached_tokens":6}}}}`, usage)
+ require.Equal(t, 17, usage.InputTokens)
+ require.Equal(t, 19, usage.OutputTokens)
+ require.Equal(t, 6, usage.CacheReadInputTokens)
+
svc.parseSSEUsage(`{"type":"response.completed","response":{"usage":{"prompt_tokens":21,"completion_tokens":8,"prompt_tokens_details":{"cached_tokens":6}}}}`, usage)
require.Equal(t, 21, usage.InputTokens)
require.Equal(t, 8, usage.OutputTokens)
@@ -2281,6 +2366,35 @@ func TestHandleSSEToJSON_CompletedEventReturnsJSON(t *testing.T) {
require.NotContains(t, rec.Body.String(), "data:")
}
+func TestHandleNonStreamingResponse_APIKeyFallsBackToSSEBodyWhenContentTypeIsWrong(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
+
+ svc := &OpenAIGatewayService{cfg: &config.Config{}}
+ resp := &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{"Content-Type": []string{"application/json"}},
+ Body: io.NopCloser(strings.NewReader(strings.Join([]string{
+ `data: {"type":"response.output_text.delta","delta":"hel"}`,
+ `data: {"type":"response.output_text.delta","delta":"lo"}`,
+ `data: {"type":"response.completed","response":{"id":"resp_api_key_sse","object":"response","model":"gpt-5.4","status":"completed","output":[],"usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}}`,
+ `data: [DONE]`,
+ }, "\n"))),
+ }
+ account := &Account{ID: 1, Type: AccountTypeAPIKey}
+
+ result, err := svc.handleNonStreamingResponse(context.Background(), resp, c, account, "gpt-5.4", "gpt-5.4")
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, 3, result.InputTokens)
+ require.Equal(t, 2, result.OutputTokens)
+ require.NotContains(t, rec.Body.String(), "data:")
+ require.Equal(t, "resp_api_key_sse", gjson.Get(rec.Body.String(), "id").String())
+ require.Equal(t, "hello", gjson.Get(rec.Body.String(), "output.0.content.0.text").String())
+}
+
func TestHandleSSEToJSON_ReconstructsImageGenerationOutputItemDone(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
diff --git a/backend/internal/service/openai_images.go b/backend/internal/service/openai_images.go
index 19066f1d..beb34780 100644
--- a/backend/internal/service/openai_images.go
+++ b/backend/internal/service/openai_images.go
@@ -622,7 +622,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesAPIKey(
return nil, fmt.Errorf("upstream request failed: %s", safeErr)
}
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody))
@@ -638,11 +638,11 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesAPIKey(
Kind: "failover",
Message: upstreamMsg,
})
- s.handleFailoverSideEffects(upstreamCtx, resp, account)
+ s.handleFailoverSideEffects(upstreamCtx, resp, account, upstreamModel)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
return s.handleErrorResponse(upstreamCtx, resp, c, account, forwardBody)
@@ -1276,21 +1276,22 @@ func resolveOpenAIImageBytes(
headers http.Header,
conversationID string,
pointer openAIImagePointerInfo,
+ errorBodyReadLimit int64,
) ([]byte, error) {
if normalized := normalizeOpenAIImageBase64(pointer.B64JSON); normalized != "" {
return base64.StdEncoding.DecodeString(normalized)
}
if downloadURL := strings.TrimSpace(pointer.DownloadURL); downloadURL != "" {
- return downloadOpenAIImageBytes(ctx, client, headers, downloadURL)
+ return downloadOpenAIImageBytes(ctx, client, headers, downloadURL, errorBodyReadLimit)
}
if strings.TrimSpace(pointer.Pointer) == "" {
return nil, fmt.Errorf("image asset is missing pointer, url, and base64 data")
}
- downloadURL, err := fetchOpenAIImageDownloadURL(ctx, client, headers, conversationID, pointer.Pointer)
+ downloadURL, err := fetchOpenAIImageDownloadURL(ctx, client, headers, conversationID, pointer.Pointer, errorBodyReadLimit)
if err != nil {
return nil, err
}
- return downloadOpenAIImageBytes(ctx, client, headers, downloadURL)
+ return downloadOpenAIImageBytes(ctx, client, headers, downloadURL, errorBodyReadLimit)
}
func normalizeOpenAIImageBase64(raw string) string {
@@ -1395,6 +1396,7 @@ func fetchOpenAIImageDownloadURL(
headers http.Header,
conversationID string,
pointer string,
+ errorBodyReadLimit int64,
) (string, error) {
url := ""
allowConversationRetry := false
@@ -1425,7 +1427,7 @@ func fetchOpenAIImageDownloadURL(
} else if resp.IsSuccessState() && strings.TrimSpace(result.DownloadURL) != "" {
return strings.TrimSpace(result.DownloadURL), nil
} else {
- statusErr := newOpenAIImageStatusError(resp, "fetch image download url failed")
+ statusErr := newOpenAIImageStatusError(resp, "fetch image download url failed", errorBodyReadLimit)
if !allowConversationRetry || !isOpenAIImageTransientConversationNotFoundError(statusErr) {
return "", statusErr
}
@@ -1450,7 +1452,7 @@ func fetchOpenAIImageDownloadURL(
return "", lastErr
}
-func downloadOpenAIImageBytes(ctx context.Context, client *req.Client, headers http.Header, downloadURL string) ([]byte, error) {
+func downloadOpenAIImageBytes(ctx context.Context, client *req.Client, headers http.Header, downloadURL string, errorBodyReadLimit int64) ([]byte, error) {
request := client.R().
SetContext(ctx).
DisableAutoReadResponse()
@@ -1478,7 +1480,7 @@ func downloadOpenAIImageBytes(ctx context.Context, client *req.Client, headers h
}
}()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
- return nil, newOpenAIImageStatusError(resp, "download image bytes failed")
+ return nil, newOpenAIImageStatusError(resp, "download image bytes failed", errorBodyReadLimit)
}
return io.ReadAll(io.LimitReader(resp.Body, openAIImageMaxDownloadBytes))
}
@@ -1505,7 +1507,7 @@ func (e *openAIImageStatusError) Error() string {
return "openai image backend request failed"
}
-func newOpenAIImageStatusError(resp *req.Response, fallback string) error {
+func newOpenAIImageStatusError(resp *req.Response, fallback string, errorBodyReadLimit int64) error {
if resp == nil {
if strings.TrimSpace(fallback) == "" {
fallback = "openai image backend request failed"
@@ -1526,7 +1528,10 @@ func newOpenAIImageStatusError(resp *req.Response, fallback string) error {
requestURL = resp.Request.URL.String()
}
if resp.Body != nil {
- body, _ = io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ if errorBodyReadLimit <= 0 {
+ errorBodyReadLimit = openAIUpstreamErrorBodyReadLimit
+ }
+ body, _ = io.ReadAll(io.LimitReader(resp.Body, errorBodyReadLimit))
_ = resp.Body.Close()
}
}
diff --git a/backend/internal/service/openai_images_responses.go b/backend/internal/service/openai_images_responses.go
index b39fa609..db9c7b16 100644
--- a/backend/internal/service/openai_images_responses.go
+++ b/backend/internal/service/openai_images_responses.go
@@ -1172,7 +1172,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth(
return nil, fmt.Errorf("upstream request failed: %s", safeErr)
}
if resp.StatusCode >= 400 {
- respBody, _ := io.ReadAll(io.LimitReader(resp.Body, 2<<20))
+ respBody := s.readUpstreamErrorBody(resp)
_ = resp.Body.Close()
resp.Body = io.NopCloser(bytes.NewReader(respBody))
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(respBody))
@@ -1188,11 +1188,11 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth(
Kind: "failover",
Message: upstreamMsg,
})
- s.handleFailoverSideEffects(upstreamCtx, resp, account)
+ s.handleFailoverSideEffects(upstreamCtx, resp, account, requestModel)
return nil, &UpstreamFailoverError{
StatusCode: resp.StatusCode,
ResponseBody: respBody,
- RetryableOnSameAccount: account.IsPoolMode() && isPoolModeRetryableStatus(resp.StatusCode),
+ RetryableOnSameAccount: account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode),
}
}
return s.handleErrorResponse(upstreamCtx, resp, c, account, responsesBody)
diff --git a/backend/internal/service/openai_images_test.go b/backend/internal/service/openai_images_test.go
index 854e9f6d..c3efdc93 100644
--- a/backend/internal/service/openai_images_test.go
+++ b/backend/internal/service/openai_images_test.go
@@ -4,6 +4,7 @@ import (
"bytes"
"context"
"errors"
+ "fmt"
"io"
"mime/multipart"
"net/http"
@@ -14,6 +15,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
+ "github.com/imroc/req/v3"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
@@ -398,11 +400,38 @@ func TestCollectOpenAIImagePointers_RecognizesDirectAssets(t *testing.T) {
func TestResolveOpenAIImageBytes_PrefersInlineBase64(t *testing.T) {
data, err := resolveOpenAIImageBytes(context.Background(), nil, nil, "", openAIImagePointerInfo{
B64JSON: "data:image/png;base64,QUJD",
- })
+ }, openAIUpstreamErrorBodyReadLimit)
require.NoError(t, err)
require.Equal(t, []byte("ABC"), data)
}
+func TestNewOpenAIImageStatusError_UsesProvidedReadLimit(t *testing.T) {
+ padding := strings.Repeat("x", int(openAIUpstreamErrorBodyReadLimit)+1024)
+ body := fmt.Sprintf(`{"error":{"padding":"%s","message":"diagnostic-marker"}}`, padding)
+ resp := &req.Response{Response: &http.Response{
+ StatusCode: http.StatusBadGateway,
+ Header: http.Header{},
+ Body: io.NopCloser(strings.NewReader(body)),
+ }}
+
+ err := newOpenAIImageStatusError(resp, "download image bytes failed", int64(len(body)))
+ require.Error(t, err)
+ require.Equal(t, "diagnostic-marker", err.Error())
+
+ var statusErr *openAIImageStatusError
+ require.ErrorAs(t, err, &statusErr)
+ require.Len(t, statusErr.ResponseBody, len(body))
+}
+
+func TestOpenAIUpstreamErrorBodyReadLimitForConfig_RespectsDiagnosticLimit(t *testing.T) {
+ cfg := &config.Config{Gateway: config.GatewayConfig{
+ LogUpstreamErrorBody: true,
+ LogUpstreamErrorBodyMaxBytes: int(openAIUpstreamErrorBodyReadLimit) + 1024,
+ }}
+
+ require.Equal(t, int64(cfg.Gateway.LogUpstreamErrorBodyMaxBytes), openAIUpstreamErrorBodyReadLimitForConfig(cfg))
+}
+
func TestAccountSupportsOpenAIImageCapability_OAuthSupportsNative(t *testing.T) {
account := &Account{
Platform: PlatformOpenAI,
@@ -413,6 +442,79 @@ func TestAccountSupportsOpenAIImageCapability_OAuthSupportsNative(t *testing.T)
require.True(t, account.SupportsOpenAIImageCapability(OpenAIImagesCapabilityNative))
}
+func TestAccountSupportsOpenAIEndpointCapability(t *testing.T) {
+ t.Run("OpenAI APIKey 默认兼容 chat 和 embeddings", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ }
+
+ require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions))
+ require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings))
+ })
+
+ t.Run("OpenAI OAuth 默认仅兼容 chat", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ }
+
+ require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions))
+ require.False(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings))
+ })
+
+ t.Run("显式列表支持同时声明 chat 和 embeddings", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{
+ "openai_capabilities": []any{"chat_completions", "embeddings"},
+ },
+ }
+
+ require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions))
+ require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings))
+ })
+
+ t.Run("显式列表只声明 chat 时不支持 embeddings", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{
+ "openai_capabilities": []any{"chat_completions"},
+ },
+ }
+
+ require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions))
+ require.False(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings))
+ })
+
+ t.Run("显式 map 支持单独关闭 chat 并开启 embeddings", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{
+ "openai_capabilities": map[string]any{
+ "chat_completions": false,
+ "embeddings": true,
+ },
+ },
+ }
+
+ require.False(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityChatCompletions))
+ require.True(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityEmbeddings))
+ })
+
+ t.Run("未知能力不应默认放行", func(t *testing.T) {
+ account := &Account{
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ }
+
+ require.False(t, account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapability("unknown")))
+ })
+}
+
func TestBuildOpenAIImagesURL_HandlesVersionedBaseURL(t *testing.T) {
require.Equal(t,
"https://image-upstream.example/v1/images/generations",
diff --git a/backend/internal/service/openai_oauth_passthrough_test.go b/backend/internal/service/openai_oauth_passthrough_test.go
index 398cbb85..2710c696 100644
--- a/backend/internal/service/openai_oauth_passthrough_test.go
+++ b/backend/internal/service/openai_oauth_passthrough_test.go
@@ -14,6 +14,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/openai_compat"
"github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
@@ -101,6 +102,57 @@ func TestOpenAIGatewayService_ResponsesUnknownModelDoesNotFallbackToGPT54(t *tes
require.True(t, rec.Code >= http.StatusBadRequest)
}
+func TestOpenAIGatewayService_NativeResponsesBodyModificationPreservesHTMLChars(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ payloadText := strings.Repeat(`&value`, 128)
+ originalBody := []byte(fmt.Sprintf(`{"model":"gpt-5.5","stream":false,"max_output_tokens":100,"previous_response_id":"resp_prev","input":[{"type":"message","role":"user","content":[{"type":"input_text","text":%q}]}]}`, payloadText))
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(originalBody))
+ c.Request.Header.Set("Content-Type", "application/json")
+
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusBadRequest,
+ Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid_native_reencode"}},
+ Body: io.NopCloser(strings.NewReader(`{"error":{"type":"invalid_request_error","message":"stop after capture"}}`)),
+ }}
+ svc := &OpenAIGatewayService{
+ cfg: &config.Config{Security: config.SecurityConfig{URLAllowlist: config.URLAllowlistConfig{
+ Enabled: false,
+ AllowInsecureHTTP: true,
+ }}},
+ httpUpstream: upstream,
+ }
+ account := &Account{
+ ID: 456,
+ Name: "openai-apikey",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": "http://upstream.example",
+ },
+ Extra: map[string]any{
+ openai_compat.ExtraKeyResponsesMode: string(openai_compat.ResponsesSupportModeAuto),
+ openai_compat.ExtraKeyResponsesSupported: true,
+ },
+ Status: StatusActive,
+ Schedulable: true,
+ }
+
+ result, err := svc.Forward(context.Background(), c, account, originalBody)
+ require.Error(t, err)
+ require.Nil(t, result)
+ require.NotNil(t, upstream.lastReq)
+ require.Equal(t, "http://upstream.example/v1/responses", upstream.lastReq.URL.String())
+ require.Contains(t, string(upstream.lastBody), payloadText)
+ require.NotContains(t, string(upstream.lastBody), `\\u003c`)
+ require.NotContains(t, string(upstream.lastBody), `\\u003e`)
+ require.NotContains(t, string(upstream.lastBody), `\\u0026`)
+}
+
func TestOpenAIGatewayService_OAuthMessagesBridgeDoesNotInjectDefaultInstructions(t *testing.T) {
gin.SetMode(gin.TestMode)
diff --git a/backend/internal/service/openai_oauth_service.go b/backend/internal/service/openai_oauth_service.go
index dc094d43..0ee357a9 100644
--- a/backend/internal/service/openai_oauth_service.go
+++ b/backend/internal/service/openai_oauth_service.go
@@ -278,11 +278,29 @@ func (s *OpenAIOAuthService) enrichTokenInfo(ctx context.Context, tokenInfo *Ope
tokenInfo.Email = info.Email
}
}
+ if strings.TrimSpace(tokenInfo.SubscriptionExpiresAt) == "" {
+ if expiresAt := fetchChatGPTSubscriptionExpiresAt(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL, resolveChatGPTSubscriptionAccountID(tokenInfo, orgID)); expiresAt != "" {
+ tokenInfo.SubscriptionExpiresAt = expiresAt
+ }
+ }
// 尝试设置隐私(关闭训练数据共享),best-effort
tokenInfo.PrivacyMode = disableOpenAITraining(ctx, s.privacyClientFactory, tokenInfo.AccessToken, proxyURL)
}
+func resolveChatGPTSubscriptionAccountID(tokenInfo *OpenAITokenInfo, orgID string) string {
+ for _, candidate := range []string{
+ tokenInfo.ChatGPTAccountID,
+ tokenInfo.OrganizationID,
+ orgID,
+ } {
+ if trimmed := strings.TrimSpace(candidate); trimmed != "" {
+ return trimmed
+ }
+ }
+ return ""
+}
+
// RefreshAccountToken refreshes token for an OpenAI OAuth account
func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *Account) (*OpenAITokenInfo, error) {
if account.Platform != PlatformOpenAI {
@@ -292,30 +310,6 @@ func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *A
return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_INVALID_ACCOUNT_TYPE", "account is not an OAuth account")
}
- refreshToken := account.GetCredential("refresh_token")
- if refreshToken == "" {
- accessToken := account.GetCredential("access_token")
- if accessToken != "" {
- tokenInfo := &OpenAITokenInfo{
- AccessToken: accessToken,
- RefreshToken: "",
- IDToken: account.GetCredential("id_token"),
- ClientID: account.GetCredential("client_id"),
- Email: account.GetCredential("email"),
- ChatGPTAccountID: account.GetCredential("chatgpt_account_id"),
- ChatGPTUserID: account.GetCredential("chatgpt_user_id"),
- OrganizationID: account.GetCredential("organization_id"),
- PlanType: account.GetCredential("plan_type"),
- }
- if expiresAt := account.GetCredentialAsTime("expires_at"); expiresAt != nil {
- tokenInfo.ExpiresAt = expiresAt.Unix()
- tokenInfo.ExpiresIn = int64(time.Until(*expiresAt).Seconds())
- }
- return tokenInfo, nil
- }
- return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_NO_REFRESH_TOKEN", "no refresh token available")
- }
-
var proxyURL string
if account.ProxyID != nil {
proxy, err := s.proxyRepo.GetByID(ctx, *account.ProxyID)
@@ -324,6 +318,32 @@ func (s *OpenAIOAuthService) RefreshAccountToken(ctx context.Context, account *A
}
}
+ refreshToken := account.GetCredential("refresh_token")
+ if refreshToken == "" {
+ accessToken := account.GetCredential("access_token")
+ if accessToken != "" {
+ tokenInfo := &OpenAITokenInfo{
+ AccessToken: accessToken,
+ RefreshToken: "",
+ IDToken: account.GetCredential("id_token"),
+ ClientID: account.GetCredential("client_id"),
+ Email: account.GetCredential("email"),
+ ChatGPTAccountID: account.GetCredential("chatgpt_account_id"),
+ ChatGPTUserID: account.GetCredential("chatgpt_user_id"),
+ OrganizationID: account.GetCredential("organization_id"),
+ PlanType: account.GetCredential("plan_type"),
+ SubscriptionExpiresAt: account.GetCredential("subscription_expires_at"),
+ }
+ if expiresAt := account.GetCredentialAsTime("expires_at"); expiresAt != nil {
+ tokenInfo.ExpiresAt = expiresAt.Unix()
+ tokenInfo.ExpiresIn = int64(time.Until(*expiresAt).Seconds())
+ }
+ s.enrichTokenInfo(ctx, tokenInfo, proxyURL)
+ return tokenInfo, nil
+ }
+ return nil, infraerrors.New(http.StatusBadRequest, "OPENAI_OAUTH_NO_REFRESH_TOKEN", "no refresh token available")
+ }
+
clientID := account.GetCredential("client_id")
return s.RefreshTokenWithClientID(ctx, refreshToken, proxyURL, clientID)
}
diff --git a/backend/internal/service/openai_oauth_service_refresh_test.go b/backend/internal/service/openai_oauth_service_refresh_test.go
index 84b68ea6..75588c8d 100644
--- a/backend/internal/service/openai_oauth_service_refresh_test.go
+++ b/backend/internal/service/openai_oauth_service_refresh_test.go
@@ -8,6 +8,7 @@ import (
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
+ "github.com/imroc/req/v3"
"github.com/stretchr/testify/require"
)
@@ -32,6 +33,11 @@ func (s *openaiOAuthClientRefreshStub) RefreshTokenWithClientID(ctx context.Cont
func TestOpenAIOAuthService_RefreshAccountToken_NoRefreshTokenUsesExistingAccessToken(t *testing.T) {
client := &openaiOAuthClientRefreshStub{}
svc := NewOpenAIOAuthService(nil, client)
+ var privacyClientCalls int32
+ svc.SetPrivacyClientFactory(func(proxyURL string) (*req.Client, error) {
+ atomic.AddInt32(&privacyClientCalls, 1)
+ return nil, errors.New("stop before request")
+ })
expiresAt := time.Now().Add(30 * time.Minute).UTC().Format(time.RFC3339)
account := &Account{
@@ -51,6 +57,7 @@ func TestOpenAIOAuthService_RefreshAccountToken_NoRefreshTokenUsesExistingAccess
require.Equal(t, "existing-access-token", info.AccessToken)
require.Equal(t, "client-id-1", info.ClientID)
require.Zero(t, atomic.LoadInt32(&client.refreshCalls), "existing access token should be reused without calling refresh")
+ require.Positive(t, atomic.LoadInt32(&privacyClientCalls), "existing access token should still run enrichment")
}
func TestOpenAITokenRefresher_NeedsRefresh_SkipsAccountWithoutRefreshToken(t *testing.T) {
diff --git a/backend/internal/service/openai_privacy_service.go b/backend/internal/service/openai_privacy_service.go
index da6dbefc..99cbb726 100644
--- a/backend/internal/service/openai_privacy_service.go
+++ b/backend/internal/service/openai_privacy_service.go
@@ -95,6 +95,8 @@ type ChatGPTAccountInfo struct {
const chatGPTAccountsCheckURL = "https://chatgpt.com/backend-api/accounts/check/v4-2023-04-27"
+var chatGPTSubscriptionsURL = "https://chatgpt.com/backend-api/subscriptions"
+
// fetchChatGPTAccountInfo calls ChatGPT backend-api to get account info (plan_type, etc.).
// Used as fallback when id_token doesn't contain these fields (e.g., Mobile RT).
// orgID is used to match the correct account when multiple accounts exist (e.g., personal + team).
@@ -199,6 +201,62 @@ func fetchChatGPTAccountInfo(ctx context.Context, clientFactory PrivacyClientFac
return info
}
+// fetchChatGPTSubscriptionExpiresAt reads the lightweight subscription endpoint used by
+// ChatGPT/Codex clients. Some Plus accounts no longer expose entitlement.expires_at in
+// accounts/check, but this endpoint still returns active_until.
+func fetchChatGPTSubscriptionExpiresAt(ctx context.Context, clientFactory PrivacyClientFactory, accessToken, proxyURL, accountID string) string {
+ accountID = strings.TrimSpace(accountID)
+ if accessToken == "" || accountID == "" || clientFactory == nil {
+ return ""
+ }
+
+ ctx, cancel := context.WithTimeout(ctx, 15*time.Second)
+ defer cancel()
+
+ client, err := clientFactory(proxyURL)
+ if err != nil {
+ slog.Debug("chatgpt_subscription_client_error", "error", err.Error())
+ return ""
+ }
+
+ var result struct {
+ PlanType string `json:"plan_type"`
+ ActiveUntil string `json:"active_until"`
+ WillRenew bool `json:"will_renew"`
+ ID string `json:"id"`
+ }
+ resp, err := client.R().
+ SetContext(ctx).
+ SetHeader("Authorization", "Bearer "+accessToken).
+ SetHeader("Origin", "https://chatgpt.com").
+ SetHeader("Referer", "https://chatgpt.com/").
+ SetHeader("Accept", "application/json").
+ SetSuccessResult(&result).
+ SetQueryParam("account_id", accountID).
+ Get(chatGPTSubscriptionsURL)
+ if err != nil {
+ slog.Debug("chatgpt_subscription_request_error", "error", err.Error())
+ return ""
+ }
+ if !resp.IsSuccessState() {
+ slog.Debug("chatgpt_subscription_failed", "status", resp.StatusCode, "body", truncate(resp.String(), 200))
+ return ""
+ }
+
+ activeUntil := strings.TrimSpace(result.ActiveUntil)
+ if activeUntil == "" {
+ slog.Debug("chatgpt_subscription_no_active_until", "plan_type", result.PlanType, "has_subscription_id", strings.TrimSpace(result.ID) != "", "will_renew", result.WillRenew)
+ return ""
+ }
+ if _, err := time.Parse(time.RFC3339, activeUntil); err != nil {
+ slog.Debug("chatgpt_subscription_bad_active_until", "active_until", activeUntil, "error", err.Error())
+ return ""
+ }
+
+ slog.Info("chatgpt_subscription_success", "plan_type", result.PlanType, "subscription_expires_at", activeUntil, "account_id", accountID)
+ return activeUntil
+}
+
// fillAccountInfo 从单个 account 对象中提取 plan_type 和 subscription_expires_at
func fillAccountInfo(info *ChatGPTAccountInfo, acct map[string]any) {
info.PlanType = extractPlanType(acct)
diff --git a/backend/internal/service/openai_subscription_test.go b/backend/internal/service/openai_subscription_test.go
new file mode 100644
index 00000000..89df54db
--- /dev/null
+++ b/backend/internal/service/openai_subscription_test.go
@@ -0,0 +1,42 @@
+package service
+
+import (
+ "context"
+ "encoding/json"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+
+ "github.com/imroc/req/v3"
+ "github.com/stretchr/testify/require"
+)
+
+func TestFetchChatGPTSubscriptionExpiresAt(t *testing.T) {
+ const wantExpiresAt = "2026-06-10T02:52:15Z"
+
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ require.Equal(t, "/backend-api/subscriptions", r.URL.Path)
+ require.Equal(t, "acc_123", r.URL.Query().Get("account_id"))
+ require.Equal(t, "Bearer access-token", r.Header.Get("Authorization"))
+
+ w.Header().Set("Content-Type", "application/json")
+ _ = json.NewEncoder(w).Encode(map[string]any{
+ "plan_type": "plus",
+ "active_until": wantExpiresAt,
+ "will_renew": true,
+ "id": "sub_123",
+ })
+ }))
+ defer server.Close()
+
+ oldURL := chatGPTSubscriptionsURL
+ chatGPTSubscriptionsURL = server.URL + "/backend-api/subscriptions"
+ t.Cleanup(func() { chatGPTSubscriptionsURL = oldURL })
+
+ got := fetchChatGPTSubscriptionExpiresAt(context.Background(), func(proxyURL string) (*req.Client, error) {
+ return req.C().SetTimeout(5 * time.Second), nil
+ }, "access-token", "", "acc_123")
+
+ require.Equal(t, wantExpiresAt, got)
+}
diff --git a/backend/internal/service/openai_tool_continuation.go b/backend/internal/service/openai_tool_continuation.go
index 7d503f5a..6515c0c4 100644
--- a/backend/internal/service/openai_tool_continuation.go
+++ b/backend/internal/service/openai_tool_continuation.go
@@ -1,6 +1,10 @@
package service
-import "strings"
+import (
+ "strings"
+
+ "github.com/tidwall/gjson"
+)
// ToolContinuationSignals 聚合工具续链相关信号,避免重复遍历 input。
type ToolContinuationSignals struct {
@@ -150,6 +154,67 @@ func AnalyzeToolContinuationSignals(reqBody map[string]any) ToolContinuationSign
return signals
}
+// ValidateFunctionCallOutputContextBytes 基于 raw JSON 校验工具输出续链,避免 handler 预校验阶段全量解码大 input。
+func ValidateFunctionCallOutputContextBytes(body []byte) FunctionCallOutputValidation {
+ result := FunctionCallOutputValidation{}
+ if len(body) == 0 {
+ return result
+ }
+ // handler 热路径只读扫描 input,避免 GetBytes 为大 Responses body 复制整段 JSON。
+ input := parseRawJSONView(body).Get("input")
+ if !input.IsArray() {
+ return result
+ }
+
+ var callIDs map[string]struct{}
+ var referenceIDs map[string]struct{}
+ input.ForEach(func(_, item gjson.Result) bool {
+ if !item.IsObject() {
+ return true
+ }
+ itemType := item.Get("type").String()
+ switch {
+ case isCodexToolCallOutputItemType(itemType):
+ result.HasFunctionCallOutput = true
+ callID := strings.TrimSpace(item.Get("call_id").String())
+ if callID == "" {
+ result.HasFunctionCallOutputMissingCallID = true
+ return true
+ }
+ if callIDs == nil {
+ callIDs = make(map[string]struct{})
+ }
+ callIDs[callID] = struct{}{}
+ case isCodexToolCallContextItemType(itemType):
+ if strings.TrimSpace(item.Get("call_id").String()) != "" {
+ result.HasToolCallContext = true
+ }
+ case itemType == "item_reference":
+ idValue := strings.TrimSpace(item.Get("id").String())
+ if idValue == "" {
+ return true
+ }
+ if referenceIDs == nil {
+ referenceIDs = make(map[string]struct{})
+ }
+ referenceIDs[idValue] = struct{}{}
+ }
+ return !result.HasFunctionCallOutput || !result.HasToolCallContext
+ })
+ if !result.HasFunctionCallOutput || result.HasToolCallContext || len(callIDs) == 0 || len(referenceIDs) == 0 {
+ return result
+ }
+ allReferenced := true
+ for callID := range callIDs {
+ if _, ok := referenceIDs[callID]; !ok {
+ allReferenced = false
+ break
+ }
+ }
+ result.HasItemReferenceForAllCallIDs = allReferenced
+ return result
+}
+
// ValidateFunctionCallOutputContext 为 handler 提供低开销校验结果:
// 1) 无工具输出直接返回
// 2) 若已存在工具调用上下文则提前返回
diff --git a/backend/internal/service/openai_tool_continuation_test.go b/backend/internal/service/openai_tool_continuation_test.go
index 0e0552f6..4610652b 100644
--- a/backend/internal/service/openai_tool_continuation_test.go
+++ b/backend/internal/service/openai_tool_continuation_test.go
@@ -1,6 +1,7 @@
package service
import (
+ "encoding/json"
"testing"
"github.com/stretchr/testify/require"
@@ -118,3 +119,68 @@ func TestHasItemReferenceForCallIDs(t *testing.T) {
require.True(t, HasItemReferenceForCallIDs(req, []string{"call_1", "call_2"}))
require.False(t, HasItemReferenceForCallIDs(req, []string{"call_1", "call_3"}))
}
+
+func TestValidateFunctionCallOutputContextBytesMatchesMapValidation(t *testing.T) {
+ // handler 预校验走 raw JSON 扫描,语义必须与 service 内部 map 校验保持一致。
+ cases := []struct {
+ name string
+ body map[string]any
+ }{
+ {
+ name: "no_input",
+ body: map[string]any{"model": "gpt-5.4"},
+ },
+ {
+ name: "missing_call_id",
+ body: map[string]any{"input": []any{map[string]any{"type": "function_call_output"}}},
+ },
+ {
+ name: "call_id_without_reference",
+ body: map[string]any{"input": []any{map[string]any{"type": "function_call_output", "call_id": "call_1"}}},
+ },
+ {
+ name: "matching_reference",
+ body: map[string]any{"input": []any{
+ map[string]any{"type": "function_call_output", "call_id": "call_1"},
+ map[string]any{"type": "item_reference", "id": "call_1"},
+ }},
+ },
+ {
+ name: "partial_reference",
+ body: map[string]any{"input": []any{
+ map[string]any{"type": "function_call_output", "call_id": "call_1"},
+ map[string]any{"type": "tool_search_output", "call_id": "call_2"},
+ map[string]any{"type": "item_reference", "id": "call_1"},
+ }},
+ },
+ {
+ name: "tool_context",
+ body: map[string]any{"input": []any{
+ map[string]any{"type": "function_call_output", "call_id": "call_1"},
+ map[string]any{"type": "function_call", "call_id": "call_1"},
+ }},
+ },
+ {
+ name: "all_codex_tool_outputs",
+ body: map[string]any{"input": []any{
+ map[string]any{"type": "function_call_output", "call_id": "call_function"},
+ map[string]any{"type": "tool_search_output", "call_id": "call_search"},
+ map[string]any{"type": "custom_tool_call_output", "call_id": "call_custom"},
+ map[string]any{"type": "mcp_tool_call_output", "call_id": "call_mcp"},
+ map[string]any{"type": "item_reference", "id": "call_function"},
+ map[string]any{"type": "item_reference", "id": "call_search"},
+ map[string]any{"type": "item_reference", "id": "call_custom"},
+ map[string]any{"type": "item_reference", "id": "call_mcp"},
+ }},
+ },
+ }
+
+ for _, tt := range cases {
+ t.Run(tt.name, func(t *testing.T) {
+ bodyBytes, err := json.Marshal(tt.body)
+ require.NoError(t, err)
+
+ require.Equal(t, ValidateFunctionCallOutputContext(tt.body), ValidateFunctionCallOutputContextBytes(bodyBytes))
+ })
+ }
+}
diff --git a/backend/internal/service/openai_ws_account_sticky_test.go b/backend/internal/service/openai_ws_account_sticky_test.go
index 4005a921..6fc44298 100644
--- a/backend/internal/service/openai_ws_account_sticky_test.go
+++ b/backend/internal/service/openai_ws_account_sticky_test.go
@@ -48,6 +48,46 @@ func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_Hit(t *testing.T
}
}
+func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_QuotaAutoPausedMiss(t *testing.T) {
+ ctx := context.Background()
+ groupID := int64(23)
+ account := Account{
+ ID: 77,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 2,
+ Extra: map[string]any{
+ "openai_apikey_responses_websockets_v2_enabled": true,
+ "codex_5h_used_percent": 96.0,
+ "auto_pause_5h_threshold": 0.95,
+ },
+ }
+ cache := &stubGatewayCache{}
+ store := NewOpenAIWSStateStore(cache)
+ cfg := newOpenAIWSV2TestConfig()
+ svc := &OpenAIGatewayService{
+ accountRepo: stubOpenAIAccountRepo{accounts: []Account{account}},
+ cache: cache,
+ cfg: cfg,
+ concurrencyService: NewConcurrencyService(stubConcurrencyCache{}),
+ openaiWSStateStore: store,
+ }
+
+ require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_prev_quota", account.ID, time.Hour))
+
+ selection, err := svc.SelectAccountByPreviousResponseID(ctx, &groupID, "resp_prev_quota", "gpt-5.1", nil, false)
+ require.NoError(t, err)
+ require.Nil(t, selection, "超过 5h 配额阈值的账号不应继续命中 previous_response_id 粘连")
+
+ // Auto-pause is transient, so the binding is preserved: the chain can resume on the
+ // same account once the quota window resets.
+ boundAccountID, getErr := store.GetResponseAccount(ctx, groupID, "resp_prev_quota")
+ require.NoError(t, getErr)
+ require.Equal(t, account.ID, boundAccountID)
+}
+
func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_RateLimitedMiss(t *testing.T) {
ctx := context.Background()
groupID := int64(23)
@@ -268,6 +308,52 @@ func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_BusyKeepsSticky(
require.Equal(t, int64(21), selection.WaitPlan.AccountID)
}
+func TestOpenAIGatewayService_SelectAccountByPreviousResponseID_CapabilityMismatchKeepsSticky(t *testing.T) {
+ ctx := context.Background()
+ groupID := int64(25)
+ account := Account{
+ ID: 31,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "openai_capabilities": []any{"chat_completions"},
+ },
+ Extra: map[string]any{
+ "openai_apikey_responses_websockets_v2_enabled": true,
+ },
+ }
+ cache := &stubGatewayCache{}
+ store := NewOpenAIWSStateStore(cache)
+ cfg := newOpenAIWSV2TestConfig()
+ svc := &OpenAIGatewayService{
+ accountRepo: stubOpenAIAccountRepo{accounts: []Account{account}},
+ cache: cache,
+ cfg: cfg,
+ concurrencyService: NewConcurrencyService(stubConcurrencyCache{}),
+ openaiWSStateStore: store,
+ }
+
+ require.NoError(t, store.BindResponseAccount(ctx, groupID, "resp_prev_capability", account.ID, time.Hour))
+
+ selection, err := svc.selectAccountByPreviousResponseIDForCapability(
+ ctx,
+ &groupID,
+ "resp_prev_capability",
+ "text-embedding-3-small",
+ nil,
+ OpenAIEndpointCapabilityEmbeddings,
+ false,
+ )
+ require.NoError(t, err)
+ require.Nil(t, selection)
+ boundAccountID, getErr := store.GetResponseAccount(ctx, groupID, "resp_prev_capability")
+ require.NoError(t, getErr)
+ require.Equal(t, account.ID, boundAccountID)
+}
+
func newOpenAIWSV2TestConfig() *config.Config {
cfg := &config.Config{}
cfg.Gateway.OpenAIWS.Enabled = true
diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go
index d7452467..43538e81 100644
--- a/backend/internal/service/openai_ws_forwarder.go
+++ b/backend/internal/service/openai_ws_forwarder.go
@@ -369,7 +369,12 @@ func openAIWSEventMayContainToolCalls(eventType string) bool {
}
func openAIWSEventShouldParseUsage(eventType string) bool {
- return eventType == "response.completed" || strings.TrimSpace(eventType) == "response.completed"
+ switch strings.TrimSpace(eventType) {
+ case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled":
+ return true
+ default:
+ return false
+ }
}
func parseOpenAIWSEventEnvelope(message []byte) (eventType string, responseID string, response gjson.Result) {
@@ -1555,6 +1560,38 @@ func openAIWSRawItemsHasFunctionCallOutput(items []json.RawMessage) bool {
return false
}
+func openAIWSRawItemsHaveToolCallContextForOutputs(items []json.RawMessage) bool {
+ if len(items) == 0 {
+ return false
+ }
+ contextCallIDs := make(map[string]struct{})
+ outputCallIDs := make(map[string]struct{})
+ for _, item := range items {
+ itemType := gjson.GetBytes(item, "type").String()
+ callID := strings.TrimSpace(gjson.GetBytes(item, "call_id").String())
+ switch {
+ case isCodexToolCallContextItemType(itemType):
+ if callID != "" {
+ contextCallIDs[callID] = struct{}{}
+ }
+ case isCodexToolCallOutputItemType(itemType):
+ if callID == "" {
+ return false
+ }
+ outputCallIDs[callID] = struct{}{}
+ }
+ }
+ if len(outputCallIDs) == 0 || len(contextCallIDs) == 0 {
+ return false
+ }
+ for callID := range outputCallIDs {
+ if _, ok := contextCallIDs[callID]; !ok {
+ return false
+ }
+ }
+ return true
+}
+
func openAIWSRawPayloadHasToolCallOutput(payload []byte) bool {
if len(payload) == 0 {
return false
@@ -2472,6 +2509,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
wsPath = normalizeOpenAIWSLogValue(parsedURL.Path)
}
debugEnabled := isOpenAIWSModeDebugEnabled()
+ isCodexCLI := openai.IsCodexOfficialClientByHeaders(c.GetHeader("User-Agent"), c.GetHeader("originator")) || (s.cfg != nil && s.cfg.Gateway.ForceCodexCLI)
type openAIWSClientPayload struct {
payloadRaw []byte
@@ -2484,6 +2522,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
imageInputSize string
payloadBytes int
}
+ ingressSessionOriginalModel := ""
applyPayloadMutation := func(current []byte, path string, value any) ([]byte, error) {
next, err := sjson.SetBytes(current, path, value)
@@ -2547,12 +2586,21 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
originalModel := strings.TrimSpace(values[1].String())
+ modelMissing := originalModel == ""
if originalModel == "" {
- return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(
- coderws.StatusPolicyViolation,
- "model is required in response.create payload",
- nil,
- )
+ // 入站 WS 长会话里,部分客户端只在第一轮 response.create 上声明
+ // model,后续 turn 复用同一 session-level model。为避免因省略
+ // model 直接断开用户连接,这里回落到上一轮已通过校验的客户端模型,
+ // 并在下方写回上游 payload,保证账号模型映射/fast policy/图片权限
+ // 仍按同一模型执行。
+ originalModel = ingressSessionOriginalModel
+ if originalModel == "" {
+ return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(
+ coderws.StatusPolicyViolation,
+ "model is required in response.create payload",
+ nil,
+ )
+ }
}
promptCacheKey := strings.TrimSpace(values[2].String())
previousResponseID := strings.TrimSpace(values[3].String())
@@ -2571,8 +2619,36 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
normalized = next
}
+ apiKey := getAPIKeyFromContext(c)
+ imageGenerationAllowed := GroupAllowsImageGeneration(apiKeyGroup(apiKey))
+ codexBridgeEnabled := isCodexCLI && imageGenerationAllowed && s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey)
+ if codexBridgeEnabled {
+ payloadMap := make(map[string]any)
+ if err := json.Unmarshal(normalized, &payloadMap); err != nil {
+ return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", err)
+ }
+ bridgeModified := false
+ if ensureOpenAIResponsesImageGenerationTool(payloadMap) {
+ bridgeModified = true
+ logOpenAIWSModeInfo("ingress_ws_codex_image_tool_injected account_id=%d", account.ID)
+ }
+ if normalizeOpenAIResponsesImageGenerationTools(payloadMap) {
+ bridgeModified = true
+ }
+ if applyCodexImageGenerationBridgeInstructions(payloadMap) {
+ bridgeModified = true
+ logOpenAIWSModeInfo("ingress_ws_codex_image_bridge_instructions_added account_id=%d", account.ID)
+ }
+ if bridgeModified {
+ rebuilt, marshalErr := json.Marshal(payloadMap)
+ if marshalErr != nil {
+ return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", marshalErr)
+ }
+ normalized = rebuilt
+ }
+ }
upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
- if upstreamModel != originalModel {
+ if modelMissing || upstreamModel != originalModel {
next, setErr := applyPayloadMutation(normalized, "model", upstreamModel)
if setErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", setErr)
@@ -2580,7 +2656,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
normalized = next
}
imageIntent := IsImageGenerationIntent(openAIResponsesEndpoint, originalModel, normalized)
- if imageIntent && !GroupAllowsImageGeneration(apiKeyGroup(getAPIKeyFromContext(c))) {
+ if imageIntent && !imageGenerationAllowed {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, ImageGenerationPermissionMessage(), nil)
}
imageBillingModel := ""
@@ -2602,11 +2678,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
// single integration point for all WS ingress turns (first + follow-up
// frames flow through here).
//
- // Model fallback: parseClientPayload above rejects any frame whose
- // "model" field is missing (line ~2493-2500), so by the time we
- // reach this point upstreamModel is always derived from a non-empty
- // per-frame model. The capturedSessionModel fallback used in the
- // passthrough adapter is therefore not needed in this path.
+ // Model fallback: first turn still requires model at the handler layer;
+ // follow-up response.create frames may omit it and then reuse
+ // ingressSessionOriginalModel. We always write a concrete upstream model
+ // before evaluating policy, so whitelist / filter behavior remains stable.
policyApplied, blocked, policyErr := s.applyOpenAIFastPolicyToWSResponseCreate(ctx, account, upstreamModel, normalized)
if policyErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", policyErr)
@@ -2635,6 +2710,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
)
}
normalized = policyApplied
+ ingressSessionOriginalModel = originalModel
return openAIWSClientPayload{
payloadRaw: normalized,
@@ -2649,6 +2725,27 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}, nil
}
+ writeClientMessage := func(message []byte) error {
+ writeCtx, cancel := context.WithTimeout(ctx, s.openAIWSWriteTimeout())
+ defer cancel()
+ return clientConn.Write(writeCtx, coderws.MessageText, message)
+ }
+
+ readClientMessage := func() ([]byte, error) {
+ msgType, payload, readErr := clientConn.Read(ctx)
+ if readErr != nil {
+ return nil, readErr
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ return nil, NewOpenAIWSClientCloseError(
+ coderws.StatusPolicyViolation,
+ fmt.Sprintf("unsupported websocket client message type: %s", msgType.String()),
+ nil,
+ )
+ }
+ return payload, nil
+ }
+
firstPayload, err := parseClientPayload(firstClientMessage)
if err != nil {
return err
@@ -2657,29 +2754,155 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
turnState := strings.TrimSpace(c.GetHeader(openAIWSTurnStateHeader))
stateStore := s.getOpenAIWSStateStore()
groupID := getOpenAIGroupIDFromContext(c)
- sessionHash := s.GenerateSessionHash(c, firstPayload.rawForHash)
- if turnState == "" && stateStore != nil && sessionHash != "" {
- if savedTurnState, ok := stateStore.GetSessionTurnState(groupID, sessionHash); ok {
- turnState = savedTurnState
- }
- }
-
- preferredConnID := ""
- if stateStore != nil && firstPayload.previousResponseID != "" {
- if connID, ok := stateStore.GetResponseConn(firstPayload.previousResponseID); ok {
- preferredConnID = connID
- }
- }
-
- storeDisabled := s.isOpenAIWSStoreDisabledInRequestRaw(firstPayload.payloadRaw, account)
storeDisabledConnMode := s.openAIWSStoreDisabledConnMode()
- if stateStore != nil && storeDisabled && firstPayload.previousResponseID == "" && sessionHash != "" {
- if connID, ok := stateStore.GetSessionConn(groupID, sessionHash); ok {
- preferredConnID = connID
+ sessionHash := ""
+ preferredConnID := ""
+ storeDisabled := false
+ refreshIngressRouteState := func(payload openAIWSClientPayload) {
+ sessionHash = s.GenerateSessionHash(c, payload.rawForHash)
+ if turnState == "" && stateStore != nil && sessionHash != "" {
+ if savedTurnState, ok := stateStore.GetSessionTurnState(groupID, sessionHash); ok {
+ turnState = savedTurnState
+ }
+ }
+
+ preferredConnID = ""
+ if stateStore != nil && payload.previousResponseID != "" {
+ if connID, ok := stateStore.GetResponseConn(payload.previousResponseID); ok {
+ preferredConnID = connID
+ }
+ }
+
+ storeDisabled = s.isOpenAIWSStoreDisabledInRequestRaw(payload.payloadRaw, account)
+ if stateStore != nil && storeDisabled && payload.previousResponseID == "" && sessionHash != "" {
+ if connID, ok := stateStore.GetSessionConn(groupID, sessionHash); ok {
+ preferredConnID = connID
+ }
+ }
+ }
+ refreshIngressRouteState(firstPayload)
+
+ if s.shouldBridgeOpenAIWSHTTP(firstPayload.payloadBytes, firstPayload.previousResponseID) {
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_start account_id=%d account_type=%s payload_bytes=%d threshold_bytes=%d has_session_hash=%v store_disabled=%v",
+ account.ID,
+ account.Type,
+ firstPayload.payloadBytes,
+ s.openAIWSHTTPBridgeThresholdBytes(),
+ sessionHash != "",
+ storeDisabled,
+ )
+ currentBridgePayload := firstPayload
+ var bridgeReplayInput []json.RawMessage
+ bridgeReplayInputExists := false
+ for turn := 1; ; turn++ {
+ if turn > 1 && hooks != nil && hooks.BeforeRequest != nil {
+ if err := hooks.BeforeRequest(turn, currentBridgePayload.payloadRaw, currentBridgePayload.originalModel); err != nil {
+ return err
+ }
+ }
+ if hooks != nil && hooks.BeforeTurn != nil {
+ if err := hooks.BeforeTurn(turn); err != nil {
+ return err
+ }
+ }
+ if turnState != "" && c != nil && c.Request != nil {
+ c.Request.Header.Set(openAIWSTurnStateHeader, turnState)
+ }
+ bridgePayloadRaw := currentBridgePayload.payloadRaw
+ bridgePayloadBytes := currentBridgePayload.payloadBytes
+ needsBridgeReplay := currentBridgePayload.previousResponseID != "" || openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw)
+ turnReplayInput, turnReplayInputExists, replayInputErr := buildOpenAIWSReplayInputSequence(
+ bridgeReplayInput,
+ bridgeReplayInputExists,
+ currentBridgePayload.payloadRaw,
+ needsBridgeReplay,
+ )
+ if replayInputErr != nil {
+ return fmt.Errorf("build websocket http bridge replay input: %w", replayInputErr)
+ }
+ if needsBridgeReplay && turnReplayInputExists {
+ updatedPayload, setInputErr := setOpenAIWSPayloadInputSequence(
+ currentBridgePayload.payloadRaw,
+ turnReplayInput,
+ true,
+ )
+ if setInputErr != nil {
+ return fmt.Errorf("set websocket http bridge replay input: %w", setInputErr)
+ }
+ bridgePayloadRaw = updatedPayload
+ bridgePayloadBytes = len(updatedPayload)
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_replay_input account_id=%d turn=%d input_items=%d previous_response_id_present=%v has_tool_output=%v",
+ account.ID,
+ turn,
+ len(turnReplayInput),
+ currentBridgePayload.previousResponseID != "",
+ openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw),
+ )
+ }
+ result, bridgeErr := s.proxyOpenAIWSHTTPBridgeTurn(
+ ctx,
+ c,
+ account,
+ token,
+ bridgePayloadRaw,
+ bridgePayloadBytes,
+ currentBridgePayload.originalModel,
+ currentBridgePayload.imageBillingModel,
+ currentBridgePayload.imageSizeTier,
+ currentBridgePayload.imageInputSize,
+ turn,
+ writeClientMessage,
+ )
+ if hooks != nil && hooks.AfterTurn != nil {
+ hooks.AfterTurn(turn, result, bridgeErr)
+ }
+ if bridgeErr != nil {
+ return bridgeErr
+ }
+ if result == nil {
+ return errors.New("websocket http bridge turn result is nil")
+ }
+ bridgeReplayInput = cloneOpenAIWSRawMessages(turnReplayInput)
+ bridgeReplayInputExists = turnReplayInputExists
+ if result.wsReplayInputExists {
+ bridgeReplayInput = append(bridgeReplayInput, cloneOpenAIWSRawMessages(result.wsReplayInput)...)
+ bridgeReplayInputExists = true
+ }
+ if bridgeTurnState := strings.TrimSpace(result.ResponseHeaders.Get(openAIWSTurnStateHeader)); bridgeTurnState != "" {
+ turnState = bridgeTurnState
+ if stateStore != nil && sessionHash != "" {
+ stateStore.BindSessionTurnState(groupID, sessionHash, bridgeTurnState, s.openAIWSSessionStickyTTL())
+ }
+ }
+ responseID := strings.TrimSpace(result.RequestID)
+ if responseID != "" && stateStore != nil {
+ ttl := s.openAIWSResponseStickyTTL()
+ logOpenAIWSBindResponseAccountWarn(groupID, account.ID, responseID, stateStore.BindResponseAccount(ctx, groupID, responseID, account.ID, ttl))
+ }
+ nextClientMessage, readErr := readClientMessage()
+ if readErr != nil {
+ if isOpenAIWSClientDisconnectError(readErr) {
+ closeStatus, closeReason := summarizeOpenAIWSReadCloseError(readErr)
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_client_closed account_id=%d close_status=%s close_reason=%s",
+ account.ID,
+ closeStatus,
+ truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
+ )
+ return nil
+ }
+ return fmt.Errorf("read client websocket request: %w", readErr)
+ }
+ nextPayload, parseErr := parseClientPayload(nextClientMessage)
+ if parseErr != nil {
+ return parseErr
+ }
+ currentBridgePayload = nextPayload
}
}
- isCodexCLI := openai.IsCodexOfficialClientByHeaders(c.GetHeader("User-Agent"), c.GetHeader("originator")) || (s.cfg != nil && s.cfg.Gateway.ForceCodexCLI)
wsHeaders, _ := s.buildOpenAIWSHeaders(c, account, token, wsDecision, isCodexCLI, turnState, strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)), firstPayload.promptCacheKey)
baseAcquireReq := openAIWSAcquireRequest{
Account: account,
@@ -2782,6 +3005,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
var dialErr *openAIWSDialError
if errors.As(acquireErr, &dialErr) && dialErr != nil && dialErr.StatusCode == http.StatusTooManyRequests {
s.persistOpenAIWSRateLimitSignal(ctx, account, dialErr.ResponseHeaders, nil, "rate_limit_exceeded", "rate_limit_error", strings.TrimSpace(acquireErr.Error()))
+ return nil, &UpstreamFailoverError{
+ StatusCode: http.StatusTooManyRequests,
+ ResponseHeaders: cloneHeader(dialErr.ResponseHeaders),
+ }
}
if errors.Is(acquireErr, errOpenAIWSPreferredConnUnavailable) {
return nil, NewOpenAIWSClientCloseError(
@@ -2825,27 +3052,6 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
return lease, nil
}
- writeClientMessage := func(message []byte) error {
- writeCtx, cancel := context.WithTimeout(ctx, s.openAIWSWriteTimeout())
- defer cancel()
- return clientConn.Write(writeCtx, coderws.MessageText, message)
- }
-
- readClientMessage := func() ([]byte, error) {
- msgType, payload, readErr := clientConn.Read(ctx)
- if readErr != nil {
- return nil, readErr
- }
- if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
- return nil, NewOpenAIWSClientCloseError(
- coderws.StatusPolicyViolation,
- fmt.Sprintf("unsupported websocket client message type: %s", msgType.String()),
- nil,
- )
- }
- return payload, nil
- }
-
sendAndRelay := func(turn int, lease *openAIWSConnLease, payload []byte, payloadBytes int, originalModel string, imageBillingModel string, imageSizeTier string, imageInputSize string) (*OpenAIForwardResult, error) {
if lease == nil {
return nil, errors.New("upstream websocket lease is nil")
@@ -2882,6 +3088,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
eventCount := 0
tokenEventCount := 0
terminalEventCount := 0
+ replayCollector := &openAIWSToolCallReplayCollector{}
firstEventType := ""
lastEventType := ""
needModelReplace := false
@@ -2977,6 +3184,14 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
false,
)
}
+ if !wroteDownstream && isOpenAIWSRateLimitError(errCodeRaw, errTypeRaw, errMsgRaw) {
+ lease.MarkBroken()
+ return nil, &UpstreamFailoverError{
+ StatusCode: http.StatusTooManyRequests,
+ ResponseBody: append([]byte(nil), upstreamMessage...),
+ ResponseHeaders: cloneHeader(lease.HandshakeHeaders()),
+ }
+ }
}
isTokenEvent := isOpenAIWSTokenEvent(eventType)
if isTokenEvent {
@@ -3004,6 +3219,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
upstreamMessage = corrected
}
}
+ replayCollector.AddEvent(eventType, upstreamMessage)
if err := writeClientMessage(upstreamMessage); err != nil {
if isOpenAIWSClientDisconnectError(err) {
clientDisconnected = true
@@ -3067,6 +3283,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
Duration: time.Since(turnStart),
FirstTokenMs: firstTokenMs,
}
+ if replayInput := replayCollector.Items(); len(replayInput) > 0 {
+ result.wsReplayInput = replayInput
+ result.wsReplayInputExists = true
+ }
if imageCount > 0 {
result.ImageCount = imageCount
result.ImageSize = imageSizeTier
@@ -3460,9 +3680,12 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
if forcePreferredConn {
// 携带 function_call_output 的请求不能丢弃 previous_response_id:
// 上游 API 需要 response chain 来匹配 tool_result 与之前的 tool_use,
- // 丢弃后会导致 "No tool call found for function call output" 400 错误。
+ // 除非 replay input 已经包含与每个 tool_result 匹配的 tool_use 上下文。
hasFCOutput := hasFunctionCallOutput
- if !turnPrevRecoveryTried && currentPreviousResponseID != "" && !hasFCOutput {
+ hasReplayToolContext := hasFCOutput &&
+ currentTurnReplayInputExists &&
+ openAIWSRawItemsHaveToolCallContextForOutputs(currentTurnReplayInput)
+ if !turnPrevRecoveryTried && currentPreviousResponseID != "" && (!hasFCOutput || hasReplayToolContext) {
updatedPayload, removed, dropErr := dropPreviousResponseIDFromRawPayload(currentPayload)
if dropErr != nil || !removed {
reason := "not_removed"
@@ -3494,11 +3717,13 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
)
} else {
logOpenAIWSModeInfo(
- "ingress_ws_preflight_ping_recovery account_id=%d turn=%d conn_id=%s action=drop_previous_response_id_retry previous_response_id=%s",
+ "ingress_ws_preflight_ping_recovery account_id=%d turn=%d conn_id=%s action=drop_previous_response_id_retry previous_response_id=%s has_function_call_output=%v has_replay_tool_context=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
+ hasFCOutput,
+ hasReplayToolContext,
)
turnPrevRecoveryTried = true
currentPayload = updatedWithInput
@@ -3510,12 +3735,18 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
}
if hasFCOutput && currentPreviousResponseID != "" {
+ reason := "function_call_output_missing_replay_context"
+ if hasReplayToolContext {
+ reason = "function_call_output_replay_not_applied"
+ }
logOpenAIWSModeInfo(
- "ingress_ws_preflight_ping_recovery_skip account_id=%d turn=%d conn_id=%s reason=function_call_output action=fail_close previous_response_id=%s",
+ "ingress_ws_preflight_ping_recovery_skip account_id=%d turn=%d conn_id=%s reason=%s action=fail_close previous_response_id=%s has_replay_tool_context=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
+ reason,
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
+ hasReplayToolContext,
)
}
resetSessionLease(true)
@@ -3595,6 +3826,10 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
lastTurnPayload = cloneOpenAIWSPayloadBytes(currentPayload)
lastTurnReplayInput = cloneOpenAIWSRawMessages(currentTurnReplayInput)
lastTurnReplayInputExists = currentTurnReplayInputExists
+ if result.wsReplayInputExists {
+ lastTurnReplayInput = append(lastTurnReplayInput, cloneOpenAIWSRawMessages(result.wsReplayInput)...)
+ lastTurnReplayInputExists = true
+ }
nextStrictState, strictStateErr := buildOpenAIWSIngressPreviousTurnStrictState(currentPayload)
if strictStateErr != nil {
lastTurnStrictState = nil
@@ -3903,7 +4138,10 @@ func isOpenAIWSTokenEvent(eventType string) bool {
if strings.HasPrefix(eventType, "response.output") {
return true
}
- return eventType == "response.completed" || eventType == "response.done"
+ // 终止事件(response.completed/done/failed/...)由 isOpenAIWSTerminalEvent 单独处理。
+ // 不能把它们当作 token event,否则当上游没有可识别的 delta 时,
+ // firstTokenMs 会被填到终止时刻,等于把"总耗时"误报为"首 token 延迟"。
+ return false
}
func replaceOpenAIWSMessageModel(message []byte, fromModel, toModel string) []byte {
@@ -3975,6 +4213,18 @@ func (s *OpenAIGatewayService) SelectAccountByPreviousResponseID(
requestedModel string,
excludedIDs map[int64]struct{},
requireCompact bool,
+) (*AccountSelectionResult, error) {
+ return s.selectAccountByPreviousResponseIDForCapability(ctx, groupID, previousResponseID, requestedModel, excludedIDs, "", requireCompact)
+}
+
+func (s *OpenAIGatewayService) selectAccountByPreviousResponseIDForCapability(
+ ctx context.Context,
+ groupID *int64,
+ previousResponseID string,
+ requestedModel string,
+ excludedIDs map[int64]struct{},
+ requiredCapability OpenAIEndpointCapability,
+ requireCompact bool,
) (*AccountSelectionResult, error) {
if s == nil {
return nil, nil
@@ -4015,12 +4265,41 @@ func (s *OpenAIGatewayService) SelectAccountByPreviousResponseID(
if requestedModel != "" && !account.IsModelSupported(requestedModel) {
return nil, nil
}
- account = s.recheckSelectedOpenAIAccountFromDB(ctx, account, requestedModel, requireCompact)
- if account == nil {
- _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
+ if !account.SupportsOpenAIEndpointCapability(requiredCapability) {
return nil, nil
}
- // 兜底:若上游 compact 能力刚被探测为不支持,但 sticky 还在,需要主动放弃。
+ // Quota auto-pause must also gate the previous_response_id sticky path; otherwise an
+ // account over its 5h/7d threshold keeps serving the same response chain even though
+ // normal scheduling skips it. Pause is transient, so fall through to normal scheduling
+ // without deleting the binding (the window may reset before the next turn).
+ if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, account); paused {
+ return nil, nil
+ }
+ if s.schedulerSnapshot != nil && s.accountRepo != nil {
+ latest, latestErr := s.accountRepo.GetByID(ctx, account.ID)
+ if latestErr != nil || latest == nil {
+ _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
+ return nil, nil
+ }
+ if shouldClearStickySession(latest, requestedModel) || !latest.IsOpenAI() || !latest.IsSchedulable() {
+ _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
+ return nil, nil
+ }
+ if requestedModel != "" && !latest.IsModelSupported(requestedModel) {
+ return nil, nil
+ }
+ if !latest.SupportsOpenAIEndpointCapability(requiredCapability) {
+ return nil, nil
+ }
+ if paused, _ := shouldAutoPauseOpenAIAccountByQuota(ctx, latest); paused {
+ return nil, nil
+ }
+ if s.isOpenAIAccountRuntimeBlocked(latest) {
+ _ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
+ return nil, nil
+ }
+ account = latest
+ }
if requireCompact && openAICompactSupportTier(account) == 0 {
_ = store.DeleteResponseAccount(ctx, derefGroupID(groupID), responseID)
return nil, nil
diff --git a/backend/internal/service/openai_ws_forwarder_hotpath_optimization_test.go b/backend/internal/service/openai_ws_forwarder_hotpath_optimization_test.go
index 0350bde9..2622f7f2 100644
--- a/backend/internal/service/openai_ws_forwarder_hotpath_optimization_test.go
+++ b/backend/internal/service/openai_ws_forwarder_hotpath_optimization_test.go
@@ -39,6 +39,24 @@ func TestParseOpenAIWSResponseUsageFromCompletedEvent(t *testing.T) {
require.Equal(t, 4, usage.CacheReadInputTokens)
}
+func TestOpenAIWSEventShouldParseUsageTerminalEvents(t *testing.T) {
+ t.Parallel()
+
+ for _, eventType := range []string{
+ "response.completed",
+ "response.done",
+ "response.failed",
+ "response.incomplete",
+ "response.cancelled",
+ "response.canceled",
+ } {
+ require.True(t, openAIWSEventShouldParseUsage(eventType), eventType)
+ require.True(t, openAIWSEventShouldParseUsage(" "+eventType+" "), eventType)
+ }
+ require.False(t, openAIWSEventShouldParseUsage("response.output_text.delta"))
+ require.False(t, openAIWSEventShouldParseUsage(""))
+}
+
func TestOpenAIWSErrorEventHelpers_ConsistentWithWrapper(t *testing.T) {
message := []byte(`{"type":"error","error":{"type":"invalid_request_error","code":"invalid_request","message":"invalid input"}}`)
codeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(message)
diff --git a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go
index edb6fbcd..dc48b990 100644
--- a/backend/internal/service/openai_ws_forwarder_ingress_session_test.go
+++ b/backend/internal/service/openai_ws_forwarder_ingress_session_test.go
@@ -164,6 +164,276 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_KeepLeaseAcrossT
require.Len(t, captureConn.writes, 2, "应向同一上游连接发送两轮 response.create")
}
+func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_FollowupCreateCanOmitModel(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.OAuthEnabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
+ cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
+ cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
+ cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
+
+ captureConn := &openAIWSCaptureConn{
+ events: [][]byte{
+ []byte(`{"type":"response.completed","response":{"id":"resp_omit_model_1","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
+ []byte(`{"type":"response.completed","response":{"id":"resp_omit_model_2","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
+ },
+ }
+ captureDialer := &openAIWSCaptureDialer{conn: captureConn}
+ pool := newOpenAIWSConnPool(cfg)
+ pool.setClientDialerForTest(captureDialer)
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: &httpUpstreamRecorder{},
+ cache: &stubGatewayCache{},
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ openaiWSPool: pool,
+ }
+ account := &Account{
+ ID: 115,
+ Name: "openai-ingress-omit-model",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "model_mapping": map[string]any{
+ "client-model": "gpt-5.1",
+ },
+ },
+ Extra: map[string]any{
+ "responses_websockets_v2_enabled": true,
+ },
+ }
+
+ serverErrCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
+ CompressionMode: coderws.CompressionContextTakeover,
+ })
+ if err != nil {
+ serverErrCh <- err
+ return
+ }
+ defer func() {
+ _ = conn.CloseNow()
+ }()
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "unit-test-agent/1.0")
+ ginCtx.Request = req
+
+ readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
+ msgType, firstMessage, readErr := conn.Read(readCtx)
+ cancel()
+ if readErr != nil {
+ serverErrCh <- readErr
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ serverErrCh <- errors.New("unsupported websocket client message type")
+ return
+ }
+
+ serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() {
+ _ = clientConn.CloseNow()
+ }()
+
+ writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
+ err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"client-model","stream":false}`))
+ cancelWrite()
+ require.NoError(t, err)
+
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
+ _, firstEvent, readErr := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, readErr)
+ require.Equal(t, "resp_omit_model_1", gjson.GetBytes(firstEvent, "response.id").String())
+
+ writeCtx, cancelWrite = context.WithTimeout(context.Background(), 3*time.Second)
+ err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","stream":false,"previous_response_id":"resp_omit_model_1"}`))
+ cancelWrite()
+ require.NoError(t, err)
+
+ readCtx, cancelRead = context.WithTimeout(context.Background(), 3*time.Second)
+ _, secondEvent, readErr := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, readErr)
+ require.Equal(t, "resp_omit_model_2", gjson.GetBytes(secondEvent, "response.id").String())
+ _ = clientConn.Close(coderws.StatusNormalClosure, "done")
+
+ select {
+ case serverErr := <-serverErrCh:
+ require.NoError(t, serverErr)
+ case <-time.After(5 * time.Second):
+ t.Fatal("等待 ingress websocket 结束超时")
+ }
+
+ require.Len(t, captureConn.writes, 2)
+ require.Equal(t, "gpt-5.1", gjson.Get(requestToJSONString(captureConn.writes[0]), "model").String())
+ require.Equal(t, "gpt-5.1", gjson.Get(requestToJSONString(captureConn.writes[1]), "model").String())
+ require.Equal(t, "resp_omit_model_1", gjson.Get(requestToJSONString(captureConn.writes[1]), "previous_response_id").String())
+}
+
+func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_InjectsCodexImageBridge(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.OAuthEnabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
+ cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
+ cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
+ cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
+
+ captureConn := &openAIWSCaptureConn{
+ events: [][]byte{
+ []byte(`{"type":"response.completed","response":{"id":"resp_codex_image_bridge","model":"gpt-5.5","usage":{"input_tokens":1,"output_tokens":1}}}`),
+ },
+ }
+ captureDialer := &openAIWSCaptureDialer{conn: captureConn}
+ pool := newOpenAIWSConnPool(cfg)
+ pool.setClientDialerForTest(captureDialer)
+
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: &httpUpstreamRecorder{},
+ cache: &stubGatewayCache{},
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ openaiWSPool: pool,
+ }
+
+ groupID := int64(3)
+ apiKey := &APIKey{
+ ID: 1,
+ UserID: 1,
+ GroupID: &groupID,
+ Group: &Group{
+ ID: groupID,
+ AllowImageGeneration: true,
+ },
+ }
+ account := &Account{
+ ID: 31,
+ Name: "openai-codex-image-ws",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "test-token",
+ },
+ Extra: map[string]any{
+ "openai_oauth_responses_websockets_v2_enabled": true,
+ "codex_image_generation_bridge": true,
+ },
+ }
+
+ serverErrCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
+ CompressionMode: coderws.CompressionContextTakeover,
+ })
+ if err != nil {
+ serverErrCh <- err
+ return
+ }
+ defer func() {
+ _ = conn.CloseNow()
+ }()
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "codex_cli_rs/0.98.0")
+ ginCtx.Request = req
+ ginCtx.Set("api_key", apiKey)
+
+ readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
+ msgType, firstMessage, readErr := conn.Read(readCtx)
+ cancel()
+ if readErr != nil {
+ serverErrCh <- readErr
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ serverErrCh <- errors.New("unsupported websocket client message type")
+ return
+ }
+
+ serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "test-token", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() {
+ _ = clientConn.CloseNow()
+ }()
+
+ writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
+ err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.5","stream":false,"input":"draw a cat"}`))
+ cancelWrite()
+ require.NoError(t, err)
+
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
+ msgType, message, err := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, err)
+ require.Equal(t, coderws.MessageText, msgType)
+ require.Equal(t, "resp_codex_image_bridge", gjson.GetBytes(message, "response.id").String())
+
+ _ = clientConn.Close(coderws.StatusNormalClosure, "done")
+
+ select {
+ case serverErr := <-serverErrCh:
+ require.NoError(t, serverErr)
+ case <-time.After(5 * time.Second):
+ t.Fatal("等待 ingress websocket 结束超时")
+ }
+
+ require.Len(t, captureConn.writes, 1)
+ upstreamPayload := requestToJSONString(captureConn.writes[0])
+ require.True(t, gjson.Get(upstreamPayload, `tools.#(type=="image_generation")`).Exists())
+ require.Equal(t, "png", gjson.Get(upstreamPayload, `tools.#(type=="image_generation").output_format`).String())
+ require.Contains(t, gjson.Get(upstreamPayload, "instructions").String(), "image_generation")
+}
+
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_DedicatedModeDoesNotReuseConnAcrossSessions(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -441,6 +711,124 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_PassthroughModeR
require.Len(t, upstreamConn.writes, 1, "passthrough 模式应透传首条 response.create")
}
+func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_PassthroughHeadersUsePromptCacheAndTurnState(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.OAuthEnabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true
+ cfg.Gateway.OpenAIWS.IngressModeDefault = OpenAIWSIngressModeCtxPool
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
+
+ upstreamConn := &openAIWSCaptureConn{
+ events: [][]byte{
+ []byte(`{"type":"response.completed","response":{"id":"resp_passthrough_headers","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
+ },
+ }
+ captureDialer := &openAIWSCaptureDialer{conn: upstreamConn}
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: &httpUpstreamRecorder{},
+ cache: &stubGatewayCache{},
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ openaiWSPassthroughDialer: captureDialer,
+ }
+ account := &Account{
+ ID: 453,
+ Name: "openai-ingress-passthrough-headers",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "access_token": "oauth-token",
+ },
+ Extra: map[string]any{
+ "openai_oauth_responses_websockets_v2_mode": OpenAIWSIngressModePassthrough,
+ },
+ }
+
+ serverErrCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
+ CompressionMode: coderws.CompressionContextTakeover,
+ })
+ if err != nil {
+ serverErrCh <- err
+ return
+ }
+ defer func() {
+ _ = conn.CloseNow()
+ }()
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "codex_cli_rs/0.98.0")
+ req.Header.Set(openAIWSTurnStateHeader, "turn-state-1")
+ req.Header.Set(openAIWSTurnMetadataHeader, "turn-meta-1")
+ ginCtx.Request = req
+
+ readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
+ msgType, firstMessage, readErr := conn.Read(readCtx)
+ cancel()
+ if readErr != nil {
+ serverErrCh <- readErr
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ serverErrCh <- errors.New("unsupported websocket client message type")
+ return
+ }
+
+ serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "oauth-token", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() {
+ _ = clientConn.CloseNow()
+ }()
+
+ writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
+ err = clientConn.Write(writeCtx, coderws.MessageText, []byte(`{"type":"response.create","model":"gpt-5.1","stream":false,"prompt_cache_key":"pcache_passthrough"}`))
+ cancelWrite()
+ require.NoError(t, err)
+
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
+ _, event, readErr := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, readErr)
+ require.Equal(t, "resp_passthrough_headers", gjson.GetBytes(event, "response.id").String())
+ _ = clientConn.Close(coderws.StatusNormalClosure, "done")
+
+ select {
+ case serverErr := <-serverErrCh:
+ if serverErr != nil {
+ require.Contains(t, serverErr.Error(), "StatusNormalClosure")
+ }
+ case <-time.After(5 * time.Second):
+ t.Fatal("等待 passthrough websocket 结束超时")
+ }
+
+ require.Equal(t, isolateOpenAISessionID(0, "pcache_passthrough"), captureDialer.lastHeaders.Get("session_id"))
+ require.Equal(t, "turn-state-1", captureDialer.lastHeaders.Get(openAIWSTurnStateHeader))
+ require.Equal(t, "turn-meta-1", captureDialer.lastHeaders.Get(openAIWSTurnMetadataHeader))
+}
+
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_ModeOffReturnsPolicyViolation(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -2053,6 +2441,161 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledStr
require.Equal(t, "world", gjson.Get(secondWrite, "input.1.text").String())
}
+func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledPreflightPingFailReplaysFunctionCallOutputWithContext(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ prevPreflightPingIdle := openAIWSIngressPreflightPingIdle
+ openAIWSIngressPreflightPingIdle = 0
+ defer func() {
+ openAIWSIngressPreflightPingIdle = prevPreflightPingIdle
+ }()
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.OAuthEnabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2
+ cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
+ cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2
+ cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
+
+ firstConn := &openAIWSPreflightFailConn{
+ events: [][]byte{
+ []byte(`{"type":"response.completed","response":{"id":"resp_turn_ping_replay_ctx_1","model":"gpt-5.1","output":[{"type":"function_call","id":"fc_replay_1","call_id":"call_replay_1","name":"shell","arguments":"{}"}],"usage":{"input_tokens":1,"output_tokens":1}}}`),
+ },
+ }
+ secondConn := &openAIWSCaptureConn{
+ events: [][]byte{
+ []byte(`{"type":"response.completed","response":{"id":"resp_turn_ping_replay_ctx_2","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`),
+ },
+ }
+ dialer := &openAIWSQueueDialer{
+ conns: []openAIWSClientConn{firstConn, secondConn},
+ }
+ pool := newOpenAIWSConnPool(cfg)
+ pool.setClientDialerForTest(dialer)
+
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: &httpUpstreamRecorder{},
+ cache: &stubGatewayCache{},
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ openaiWSPool: pool,
+ }
+
+ account := &Account{
+ ID: 128,
+ Name: "openai-ingress-preflight-replay-function-output-with-context",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ },
+ Extra: map[string]any{
+ "responses_websockets_v2_enabled": true,
+ },
+ }
+
+ serverErrCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{
+ CompressionMode: coderws.CompressionContextTakeover,
+ })
+ if err != nil {
+ serverErrCh <- err
+ return
+ }
+ defer func() {
+ _ = conn.CloseNow()
+ }()
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "unit-test-agent/1.0")
+ ginCtx.Request = req
+
+ readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
+ msgType, firstMessage, readErr := conn.Read(readCtx)
+ cancel()
+ if readErr != nil {
+ serverErrCh <- readErr
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ serverErrCh <- errors.New("unsupported websocket client message type")
+ return
+ }
+
+ serverErrCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() {
+ _ = clientConn.CloseNow()
+ }()
+
+ writeMessage := func(payload string) {
+ writeCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancel()
+ require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
+ }
+ readMessage := func() []byte {
+ readCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancel()
+ msgType, message, readErr := clientConn.Read(readCtx)
+ require.NoError(t, readErr)
+ require.Equal(t, coderws.MessageText, msgType)
+ return message
+ }
+
+ writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false,"input":[{"type":"message","role":"user","content":[{"type":"input_text","text":"call tool"}]}]}`)
+ firstTurn := readMessage()
+ require.Equal(t, "resp_turn_ping_replay_ctx_1", gjson.GetBytes(firstTurn, "response.id").String())
+
+ writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false,"store":false,"previous_response_id":"resp_turn_ping_replay_ctx_1","input":[{"type":"function_call_output","call_id":"call_replay_1","output":"ok"}]}`)
+ secondTurn := readMessage()
+ require.Equal(t, "resp_turn_ping_replay_ctx_2", gjson.GetBytes(secondTurn, "response.id").String())
+
+ require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
+ select {
+ case serverErr := <-serverErrCh:
+ require.NoError(t, serverErr)
+ case <-time.After(5 * time.Second):
+ t.Fatal("等待 ingress websocket function_call_output 自包含重放后结束超时")
+ }
+
+ require.Equal(t, 2, dialer.DialCount(), "带完整 tool 上下文的 function_call_output 应在 ping 失败后换新连接重放")
+ require.Equal(t, 1, firstConn.WriteCount())
+ require.GreaterOrEqual(t, firstConn.PingCount(), 1)
+ secondConn.mu.Lock()
+ secondWrites := append([]map[string]any(nil), secondConn.writes...)
+ secondConn.mu.Unlock()
+ require.Len(t, secondWrites, 1)
+ secondWrite := requestToJSONString(secondWrites[0])
+ require.False(t, gjson.Get(secondWrite, "previous_response_id").Exists())
+ require.Equal(t, 3, len(gjson.Get(secondWrite, "input").Array()))
+ require.Equal(t, "message", gjson.Get(secondWrite, "input.0.type").String())
+ require.Equal(t, "function_call", gjson.Get(secondWrite, "input.1.type").String())
+ require.Equal(t, "call_replay_1", gjson.Get(secondWrite, "input.1.call_id").String())
+ require.Equal(t, "function_call_output", gjson.Get(secondWrite, "input.2.type").String())
+ require.Equal(t, "call_replay_1", gjson.Get(secondWrite, "input.2.call_id").String())
+}
+
func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_StoreDisabledPreflightPingFailClosesWhenFunctionCallOutputNeedsPreviousResponseID(t *testing.T) {
gin.SetMode(gin.TestMode)
prevPreflightPingIdle := openAIWSIngressPreflightPingIdle
diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go
index cd816533..bd262207 100644
--- a/backend/internal/service/openai_ws_forwarder_success_test.go
+++ b/backend/internal/service/openai_ws_forwarder_success_test.go
@@ -171,6 +171,91 @@ func TestOpenAIGatewayService_Forward_WSv2_SuccessAndBindSticky(t *testing.T) {
require.Equal(t, "resp_new_1", gjson.GetBytes(responseBody, "id").String())
}
+func TestOpenAIGatewayService_Forward_WSv2_UsesPatchedBodyAfterValidationDecode(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ type receivedPayload struct {
+ MaxCompletionTokensExists bool
+ }
+ receivedCh := make(chan receivedPayload, 1)
+
+ upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }}
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := upgrader.Upgrade(w, r, nil)
+ if err != nil {
+ t.Errorf("upgrade websocket failed: %v", err)
+ return
+ }
+ defer func() { _ = conn.Close() }()
+
+ var request map[string]any
+ if err := conn.ReadJSON(&request); err != nil {
+ t.Errorf("read ws request failed: %v", err)
+ return
+ }
+ requestJSON := requestToJSONString(request)
+ receivedCh <- receivedPayload{MaxCompletionTokensExists: gjson.Get(requestJSON, "max_completion_tokens").Exists()}
+
+ if err := conn.WriteJSON(map[string]any{
+ "type": "response.completed",
+ "response": map[string]any{
+ "id": "resp_patched_ws_1",
+ "model": "gpt-5.3-codex-spark",
+ "usage": map[string]any{"input_tokens": 1, "output_tokens": 1},
+ },
+ }); err != nil {
+ t.Errorf("write response.completed failed: %v", err)
+ return
+ }
+ }))
+ defer wsServer.Close()
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ c.Request.Header.Set("User-Agent", "unit-test-agent/1.0")
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 30
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 10
+
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ }
+
+ account := &Account{
+ ID: 10,
+ Name: "openai-ws",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ "base_url": wsServer.URL,
+ },
+ Extra: map[string]any{"responses_websockets_v2_enabled": true},
+ }
+
+ body := []byte(`{"model":"gpt-5.4","stream":false,"max_completion_tokens":12,"tools":[{"type":"image_generation"}],"input":[{"type":"input_text","text":"hello"}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.True(t, result.OpenAIWSMode)
+
+ received := <-receivedCh
+ require.False(t, received.MaxCompletionTokensExists)
+}
+
func TestOpenAIGatewayService_Forward_WSv2_ImageGenerationCountsOutputs(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -727,6 +812,70 @@ func TestOpenAIGatewayService_Forward_WSv2_HeaderSessionFallbackFromPromptCacheK
require.True(t, gjson.Get(requestToJSONString(captureConn.lastWrite), "stream").Exists())
}
+func TestOpenAIGatewayService_Forward_WSv2_ResponseDoneUsageParsed(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ rec := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(rec)
+ c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil)
+ c.Request.Header.Set("User-Agent", "unit-test-agent/1.0")
+
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.OAuthEnabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
+ cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
+ cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
+
+ captureConn := &openAIWSCaptureConn{
+ events: [][]byte{
+ []byte(`{"type":"response.done","response":{"id":"resp_done_usage","model":"gpt-5.1","usage":{"input_tokens":13,"output_tokens":8,"input_tokens_details":{"cached_tokens":5},"cache_creation_input_tokens":2,"output_tokens_details":{"image_tokens":4}}}}`),
+ },
+ }
+ captureDialer := &openAIWSCaptureDialer{conn: captureConn}
+ pool := newOpenAIWSConnPool(cfg)
+ pool.setClientDialerForTest(captureDialer)
+
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: &httpUpstreamRecorder{},
+ cache: &stubGatewayCache{},
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ openaiWSPool: pool,
+ }
+ account := &Account{
+ ID: 32,
+ Name: "openai-ws-done",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Credentials: map[string]any{
+ "api_key": "sk-test",
+ },
+ Extra: map[string]any{
+ "responses_websockets_v2_enabled": true,
+ },
+ }
+
+ body := []byte(`{"model":"gpt-5.1","stream":false,"input":[{"type":"input_text","text":"hi"}]}`)
+ result, err := svc.Forward(context.Background(), c, account, body)
+ require.NoError(t, err)
+ require.NotNil(t, result)
+ require.Equal(t, "resp_done_usage", result.RequestID)
+ require.Equal(t, 13, result.Usage.InputTokens)
+ require.Equal(t, 8, result.Usage.OutputTokens)
+ require.Equal(t, 5, result.Usage.CacheReadInputTokens)
+ require.Equal(t, 2, result.Usage.CacheCreationInputTokens)
+ require.Equal(t, 4, result.Usage.ImageOutputTokens)
+}
+
func TestOpenAIGatewayService_Forward_WSv1_Unsupported(t *testing.T) {
gin.SetMode(gin.TestMode)
diff --git a/backend/internal/service/openai_ws_forwarder_test.go b/backend/internal/service/openai_ws_forwarder_test.go
new file mode 100644
index 00000000..d817de6e
--- /dev/null
+++ b/backend/internal/service/openai_ws_forwarder_test.go
@@ -0,0 +1,75 @@
+package service
+
+import (
+ "testing"
+
+ "github.com/stretchr/testify/require"
+)
+
+// TestIsOpenAIWSTokenEvent_TerminalEventsExcluded 覆盖 isOpenAIWSTokenEvent 的回归用例。
+// 重点验证终止事件(response.completed / response.done)不再被当作 token event,
+// 否则当上游没有可识别的 delta 时,firstTokenMs 会被填到终止时刻,
+// 等于把"总耗时"误报为"首 token 延迟"(issue #2651)。
+func TestIsOpenAIWSTokenEvent_TerminalEventsExcluded(t *testing.T) {
+ cases := []struct {
+ name string
+ eventType string
+ want bool
+ }{
+ {name: "empty", eventType: "", want: false},
+ {name: "whitespace_trimmed_empty", eventType: " ", want: false},
+
+ {name: "response.created", eventType: "response.created", want: false},
+ {name: "response.in_progress", eventType: "response.in_progress", want: false},
+ {name: "response.output_item.added", eventType: "response.output_item.added", want: false},
+ {name: "response.output_item.done", eventType: "response.output_item.done", want: false},
+
+ {name: "terminal_response.completed", eventType: "response.completed", want: false},
+ {name: "terminal_response.done", eventType: "response.done", want: false},
+ {name: "terminal_response.completed_padded", eventType: " response.completed ", want: false},
+ {name: "terminal_response.done_padded", eventType: " response.done ", want: false},
+
+ {name: "delta_text", eventType: "response.output_text.delta", want: true},
+ {name: "delta_audio_transcript", eventType: "response.audio_transcript.delta", want: true},
+ {name: "delta_function_call_arguments", eventType: "response.function_call_arguments.delta", want: true},
+
+ {name: "output_text_done", eventType: "response.output_text.done", want: true},
+ {name: "output_text_annotation_added", eventType: "response.output_text.annotation.added", want: true},
+
+ {name: "output_audio_done", eventType: "response.output_audio.done", want: true},
+
+ {name: "reasoning_summary_delta", eventType: "response.reasoning_summary_text.delta", want: true},
+
+ {name: "unrelated_event_error", eventType: "error", want: false},
+ {name: "unknown_event_without_match", eventType: "response.reasoning_summary_part.added", want: false},
+ }
+
+ for _, tc := range cases {
+ tc := tc
+ t.Run(tc.name, func(t *testing.T) {
+ got := isOpenAIWSTokenEvent(tc.eventType)
+ require.Equal(t, tc.want, got, "isOpenAIWSTokenEvent(%q)", tc.eventType)
+ })
+ }
+}
+
+// TestIsOpenAIWSTokenEvent_DisjointWithTerminal 守护「token 事件集合与终止事件集合互斥」的不变量。
+// firstTokenMs 的计算依赖于 isTokenEvent && !isTerminalEvent;
+// 若两者再次出现交集,则 issue #2651 描述的 latency 误报会重现。
+func TestIsOpenAIWSTokenEvent_DisjointWithTerminal(t *testing.T) {
+ terminalEvents := []string{
+ "response.completed",
+ "response.done",
+ "response.failed",
+ "response.incomplete",
+ "response.cancelled",
+ "response.canceled",
+ }
+ for _, ev := range terminalEvents {
+ ev := ev
+ t.Run(ev, func(t *testing.T) {
+ require.True(t, isOpenAIWSTerminalEvent(ev), "expected terminal event %q to be classified as terminal", ev)
+ require.False(t, isOpenAIWSTokenEvent(ev), "terminal event %q must NOT be classified as token event (issue #2651)", ev)
+ })
+ }
+}
diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go
new file mode 100644
index 00000000..1f0f32a0
--- /dev/null
+++ b/backend/internal/service/openai_ws_http_bridge.go
@@ -0,0 +1,387 @@
+package service
+
+import (
+ "bufio"
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "net/http"
+ "strings"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/gin-gonic/gin"
+ "github.com/tidwall/gjson"
+)
+
+const (
+ openAIWSClientReadLimitBytesDefault int64 = 64 * 1024 * 1024
+ openAIWSHTTPBridgeThresholdBytesDefault int64 = 15 * 1024 * 1024
+ openAIWSHTTPBridgeErrorBodyLimitBytes = 64 * 1024
+)
+
+func ResolveOpenAIWSClientReadLimitBytes(cfg *config.Config) int64 {
+ if cfg == nil || cfg.Gateway.OpenAIWS.ClientReadLimitBytes <= 0 {
+ return openAIWSClientReadLimitBytesDefault
+ }
+ return cfg.Gateway.OpenAIWS.ClientReadLimitBytes
+}
+
+func (s *OpenAIGatewayService) openAIWSHTTPBridgeEnabled() bool {
+ return s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.HTTPBridgeEnabled
+}
+
+func (s *OpenAIGatewayService) openAIWSHTTPBridgeThresholdBytes() int64 {
+ if s == nil || s.cfg == nil || s.cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes <= 0 {
+ return openAIWSHTTPBridgeThresholdBytesDefault
+ }
+ return s.cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes
+}
+
+func (s *OpenAIGatewayService) shouldBridgeOpenAIWSHTTP(payloadBytes int, previousResponseID string) bool {
+ if !s.openAIWSHTTPBridgeEnabled() {
+ return false
+ }
+ if strings.TrimSpace(previousResponseID) != "" {
+ return false
+ }
+ threshold := s.openAIWSHTTPBridgeThresholdBytes()
+ return threshold > 0 && int64(payloadBytes) >= threshold
+}
+
+func prepareOpenAIWSHTTPBridgeBody(payload []byte) ([]byte, error) {
+ var body map[string]any
+ if err := json.Unmarshal(payload, &body); err != nil {
+ return nil, err
+ }
+ if body == nil {
+ return nil, errors.New("response.create payload must be a JSON object")
+ }
+ delete(body, "type")
+ delete(body, "generate")
+ delete(body, "previous_response_id")
+ body["stream"] = true
+ return json.Marshal(body)
+}
+
+type openAIWSToolCallReplayCollector struct {
+ items []json.RawMessage
+ seen map[string]struct{}
+}
+
+func (c *openAIWSToolCallReplayCollector) AddEvent(eventType string, message []byte) {
+ switch strings.TrimSpace(eventType) {
+ case "response.output_item.done":
+ c.addItem(gjson.GetBytes(message, "item"))
+ case "response.completed", "response.done":
+ output := gjson.GetBytes(message, "response.output")
+ if !output.IsArray() {
+ return
+ }
+ for _, item := range output.Array() {
+ c.addItem(item)
+ }
+ }
+}
+
+func (c *openAIWSToolCallReplayCollector) Items() []json.RawMessage {
+ return cloneOpenAIWSRawMessages(c.items)
+}
+
+func (c *openAIWSToolCallReplayCollector) addItem(item gjson.Result) {
+ if !item.Exists() || item.Type != gjson.JSON {
+ return
+ }
+ raw := strings.TrimSpace(item.Raw)
+ if raw == "" || !strings.HasPrefix(raw, "{") {
+ return
+ }
+ if !isCodexToolCallContextItemType(item.Get("type").String()) {
+ return
+ }
+ key := strings.TrimSpace(item.Get("id").String())
+ if key == "" {
+ key = strings.TrimSpace(item.Get("call_id").String())
+ }
+ if key == "" {
+ key = raw
+ }
+ if c.seen == nil {
+ c.seen = make(map[string]struct{})
+ }
+ if _, ok := c.seen[key]; ok {
+ return
+ }
+ c.seen[key] = struct{}{}
+ c.items = append(c.items, json.RawMessage(raw))
+}
+
+func buildOpenAIWSHTTPBridgeErrorEvent(statusCode int, message string) []byte {
+ message = strings.TrimSpace(message)
+ if message == "" {
+ message = http.StatusText(statusCode)
+ }
+ if message == "" {
+ message = "upstream request failed"
+ }
+ event := map[string]any{
+ "type": "error",
+ "status": statusCode,
+ "error": map[string]any{
+ "type": "upstream_error",
+ "message": message,
+ },
+ }
+ body, err := json.Marshal(event)
+ if err != nil {
+ return []byte(`{"type":"error","error":{"type":"upstream_error","message":"upstream request failed"}}`)
+ }
+ return body
+}
+
+func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
+ ctx context.Context,
+ c *gin.Context,
+ account *Account,
+ token string,
+ payload []byte,
+ payloadBytes int,
+ originalModel string,
+ imageBillingModel string,
+ imageSizeTier string,
+ imageInputSize string,
+ turn int,
+ writeClientMessage func([]byte) error,
+) (*OpenAIForwardResult, error) {
+ if s == nil {
+ return nil, errors.New("service is nil")
+ }
+ if s.httpUpstream == nil {
+ return nil, errors.New("openai http upstream is nil")
+ }
+ if account == nil {
+ return nil, errors.New("account is nil")
+ }
+ if writeClientMessage == nil {
+ return nil, errors.New("client websocket writer is nil")
+ }
+
+ body, err := prepareOpenAIWSHTTPBridgeBody(payload)
+ if err != nil {
+ return nil, fmt.Errorf("prepare http bridge body: %w", err)
+ }
+
+ upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
+ upstreamReq, err := s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token)
+ releaseUpstreamCtx()
+ if err != nil {
+ return nil, err
+ }
+
+ proxyURL := ""
+ if account.ProxyID != nil && account.Proxy != nil {
+ proxyURL = account.Proxy.URL()
+ }
+ if c != nil {
+ c.Set("openai_passthrough", true)
+ c.Set("openai_ws_http_bridge", true)
+ }
+
+ turnStart := time.Now()
+ resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency)
+ if err != nil {
+ safeErr := sanitizeUpstreamErrorMessage(err.Error())
+ _ = writeClientMessage(buildOpenAIWSHTTPBridgeErrorEvent(http.StatusBadGateway, "Upstream request failed"))
+ return nil, fmt.Errorf("upstream http bridge request failed: %s", safeErr)
+ }
+ defer func() { _ = resp.Body.Close() }()
+
+ if resp.StatusCode >= 400 {
+ respBody, _ := io.ReadAll(io.LimitReader(resp.Body, openAIWSHTTPBridgeErrorBodyLimitBytes))
+ upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody)))
+ if upstreamMsg == "" {
+ upstreamMsg = http.StatusText(resp.StatusCode)
+ }
+ _ = writeClientMessage(buildOpenAIWSHTTPBridgeErrorEvent(resp.StatusCode, upstreamMsg))
+ return nil, fmt.Errorf("upstream http bridge error: status=%d message=%s", resp.StatusCode, upstreamMsg)
+ }
+
+ responseID := ""
+ usage := OpenAIUsage{}
+ imageCounter := newOpenAIImageOutputCounter()
+ var firstTokenMs *int
+ reqStream := openAIWSPayloadBoolFromRaw(body, "stream", true)
+ eventCount := 0
+ tokenEventCount := 0
+ terminalEventCount := 0
+ replayCollector := &openAIWSToolCallReplayCollector{}
+ firstEventType := ""
+ lastEventType := ""
+ sawDone := false
+ wroteDownstream := false
+ clientDisconnected := false
+ mappedModel := ""
+ needModelReplace := false
+ var mappedModelBytes []byte
+ if originalModel != "" {
+ mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
+ needModelReplace = mappedModel != "" && mappedModel != originalModel
+ if needModelReplace {
+ mappedModelBytes = []byte(mappedModel)
+ }
+ }
+
+ resultWithUsage := func() *OpenAIForwardResult {
+ imageCount := imageCounter.Count()
+ result := &OpenAIForwardResult{
+ RequestID: responseID,
+ Usage: usage,
+ Model: originalModel,
+ UpstreamModel: mappedModel,
+ ServiceTier: extractOpenAIServiceTierFromBody(body),
+ ReasoningEffort: extractOpenAIReasoningEffortFromBody(body, originalModel),
+ Stream: reqStream,
+ OpenAIWSMode: true,
+ ResponseHeaders: cloneHeader(resp.Header),
+ Duration: time.Since(turnStart),
+ FirstTokenMs: firstTokenMs,
+ }
+ if replayInput := replayCollector.Items(); len(replayInput) > 0 {
+ result.wsReplayInput = replayInput
+ result.wsReplayInputExists = true
+ }
+ if imageCount > 0 {
+ result.ImageCount = imageCount
+ result.ImageSize = imageSizeTier
+ result.ImageInputSize = imageInputSize
+ result.ImageOutputSizes = imageCounter.Sizes()
+ result.BillingModel = imageBillingModel
+ }
+ return result
+ }
+
+ scanner := bufio.NewScanner(resp.Body)
+ maxLineSize := defaultMaxLineSize
+ if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
+ maxLineSize = s.cfg.Gateway.MaxLineSize
+ }
+ scanBuf := getSSEScannerBuf64K()
+ scanner.Buffer(scanBuf[:0], maxLineSize)
+ defer putSSEScannerBuf64K(scanBuf)
+
+ for scanner.Scan() {
+ line := scanner.Text()
+ data, ok := extractOpenAISSEDataLine(line)
+ if !ok {
+ continue
+ }
+ trimmedData := strings.TrimSpace(data)
+ if trimmedData == "" {
+ continue
+ }
+ if trimmedData == "[DONE]" {
+ sawDone = true
+ continue
+ }
+
+ upstreamMessage := []byte(trimmedData)
+ eventType, eventResponseID, _ := parseOpenAIWSEventEnvelope(upstreamMessage)
+ if responseID == "" && eventResponseID != "" {
+ responseID = eventResponseID
+ }
+ if eventType != "" {
+ eventCount++
+ if firstEventType == "" {
+ firstEventType = eventType
+ }
+ lastEventType = eventType
+ }
+ if isOpenAIWSTokenEvent(eventType) {
+ tokenEventCount++
+ if firstTokenMs == nil {
+ ms := int(time.Since(turnStart).Milliseconds())
+ firstTokenMs = &ms
+ }
+ }
+ if openAIWSEventShouldParseUsage(eventType) {
+ parseOpenAIWSResponseUsageFromCompletedEvent(upstreamMessage, &usage)
+ }
+ imageCounter.AddSSEData(upstreamMessage)
+
+ if needModelReplace && len(mappedModelBytes) > 0 && openAIWSEventMayContainModel(eventType) && strings.Contains(trimmedData, mappedModel) {
+ upstreamMessage = replaceOpenAIWSMessageModel(upstreamMessage, mappedModel, originalModel)
+ }
+ if s.toolCorrector != nil && openAIWSEventMayContainToolCalls(eventType) && openAIWSMessageLikelyContainsToolCalls(upstreamMessage) {
+ if corrected, changed := s.toolCorrector.CorrectToolCallsInSSEBytes(upstreamMessage); changed {
+ upstreamMessage = corrected
+ }
+ }
+ replayCollector.AddEvent(eventType, upstreamMessage)
+
+ if !clientDisconnected {
+ if err := writeClientMessage(upstreamMessage); err != nil {
+ if isOpenAIWSClientDisconnectError(err) {
+ clientDisconnected = true
+ closeStatus, closeReason := summarizeOpenAIWSReadCloseError(err)
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_client_disconnected_drain account_id=%d turn=%d close_status=%s close_reason=%s",
+ account.ID,
+ turn,
+ closeStatus,
+ truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
+ )
+ } else {
+ return nil, wrapOpenAIWSIngressTurnError(
+ "write_client",
+ fmt.Errorf("write client websocket event: %w", err),
+ wroteDownstream,
+ )
+ }
+ } else {
+ wroteDownstream = true
+ }
+ }
+
+ if eventType == "error" {
+ errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(upstreamMessage)
+ s.persistOpenAIWSRateLimitSignal(ctx, account, resp.Header, upstreamMessage, errCodeRaw, errTypeRaw, errMsgRaw)
+ errMessage := strings.TrimSpace(errMsgRaw)
+ if errMessage == "" {
+ errMessage = "upstream error event"
+ }
+ return resultWithUsage(), errors.New(errMessage)
+ }
+ if isOpenAIWSTerminalEvent(eventType) {
+ terminalEventCount++
+ firstTokenMsValue := -1
+ if firstTokenMs != nil {
+ firstTokenMsValue = *firstTokenMs
+ }
+ logOpenAIWSModeInfo(
+ "ingress_ws_http_bridge_turn_completed account_id=%d turn=%d response_id=%s payload_bytes=%d duration_ms=%d events=%d token_events=%d terminal_events=%d first_event=%s last_event=%s first_token_ms=%d client_disconnected=%v",
+ account.ID,
+ turn,
+ truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen),
+ payloadBytes,
+ time.Since(turnStart).Milliseconds(),
+ eventCount,
+ tokenEventCount,
+ terminalEventCount,
+ truncateOpenAIWSLogValue(firstEventType, openAIWSLogValueMaxLen),
+ truncateOpenAIWSLogValue(lastEventType, openAIWSLogValueMaxLen),
+ firstTokenMsValue,
+ clientDisconnected,
+ )
+ return resultWithUsage(), nil
+ }
+ }
+ if err := scanner.Err(); err != nil {
+ return resultWithUsage(), fmt.Errorf("read upstream http bridge stream: %w", err)
+ }
+ if sawDone && eventCount > 0 {
+ return resultWithUsage(), nil
+ }
+ return resultWithUsage(), errors.New("upstream http bridge stream ended before terminal event")
+}
diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go
new file mode 100644
index 00000000..0a1d6b56
--- /dev/null
+++ b/backend/internal/service/openai_ws_http_bridge_test.go
@@ -0,0 +1,461 @@
+package service
+
+import (
+ "context"
+ "io"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ coderws "github.com/coder/websocket"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/require"
+ "github.com/tidwall/gjson"
+)
+
+func TestPrepareOpenAIWSHTTPBridgeBodyStripsWSFields(t *testing.T) {
+ body, err := prepareOpenAIWSHTTPBridgeBody([]byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":false,"previous_response_id":"resp_prev","input":"hi"}`))
+ require.NoError(t, err)
+ require.False(t, gjson.GetBytes(body, "type").Exists())
+ require.False(t, gjson.GetBytes(body, "generate").Exists())
+ require.False(t, gjson.GetBytes(body, "previous_response_id").Exists())
+ require.Equal(t, "gpt-5", gjson.GetBytes(body, "model").String())
+ require.True(t, gjson.GetBytes(body, "stream").Bool())
+ require.Equal(t, "hi", gjson.GetBytes(body, "input").String())
+}
+
+func TestOpenAIWSHTTPBridgeDecisionKeepsSmallFramesOnWS(t *testing.T) {
+ svc := &OpenAIGatewayService{
+ cfg: &config.Config{
+ Gateway: config.GatewayConfig{
+ OpenAIWS: config.GatewayOpenAIWSConfig{
+ HTTPBridgeEnabled: true,
+ HTTPBridgeThresholdBytes: 100,
+ },
+ },
+ },
+ }
+
+ require.False(t, svc.shouldBridgeOpenAIWSHTTP(99, ""))
+ require.True(t, svc.shouldBridgeOpenAIWSHTTP(100, ""))
+ require.False(t, svc.shouldBridgeOpenAIWSHTTP(1000, "resp_existing"))
+
+ svc.cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = false
+ require.False(t, svc.shouldBridgeOpenAIWSHTTP(1000, ""))
+}
+
+func TestOpenAIWSHTTPBridgeRelaysSSEFramesAsWebSocketMessages(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ sseBody := strings.Join([]string{
+ `data: {"type":"response.created","response":{"id":"resp_bridge","model":"gpt-5"}}`,
+ "",
+ `data: {"type":"response.output_text.delta","response":{"id":"resp_bridge"},"delta":"ok"}`,
+ "",
+ `data: {"type":"response.completed","response":{"id":"resp_bridge","model":"gpt-5","usage":{"input_tokens":3,"output_tokens":2}}}`,
+ "",
+ }, "\n")
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "Content-Type": []string{"text/event-stream"},
+ "x-request-id": []string{"rid_bridge"},
+ },
+ Body: io.NopCloser(strings.NewReader(sseBody)),
+ }}
+ svc := &OpenAIGatewayService{
+ cfg: &config.Config{
+ Gateway: config.GatewayConfig{
+ MaxLineSize: defaultMaxLineSize,
+ OpenAIWS: config.GatewayOpenAIWSConfig{
+ HTTPBridgeEnabled: true,
+ HTTPBridgeThresholdBytes: 1,
+ },
+ },
+ },
+ httpUpstream: upstream,
+ toolCorrector: NewCodexToolCorrector(),
+ }
+ account := &Account{
+ ID: 7,
+ Name: "api-key",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Concurrency: 1,
+ Status: StatusActive,
+ }
+ payload := []byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":true,"input":"hi"}`)
+
+ type bridgeResult struct {
+ result *OpenAIForwardResult
+ err error
+ }
+ resultCh := make(chan bridgeResult, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
+ if err != nil {
+ resultCh <- bridgeResult{err: err}
+ return
+ }
+ defer func() { _ = conn.CloseNow() }()
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ ginCtx.Request = req
+
+ writeClient := func(message []byte) error {
+ writeCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second)
+ defer cancel()
+ return conn.Write(writeCtx, coderws.MessageText, message)
+ }
+ result, bridgeErr := svc.proxyOpenAIWSHTTPBridgeTurn(
+ r.Context(),
+ ginCtx,
+ account,
+ "sk-test",
+ payload,
+ len(payload),
+ "gpt-5",
+ "",
+ "",
+ "",
+ 1,
+ writeClient,
+ )
+ resultCh <- bridgeResult{result: result, err: bridgeErr}
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() { _ = clientConn.CloseNow() }()
+
+ readEvent := func() []byte {
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
+ msgType, event, readErr := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, readErr)
+ require.Equal(t, coderws.MessageText, msgType)
+ return event
+ }
+
+ created := readEvent()
+ delta := readEvent()
+ completed := readEvent()
+
+ require.Equal(t, "response.created", gjson.GetBytes(created, "type").String())
+ require.Equal(t, "response.output_text.delta", gjson.GetBytes(delta, "type").String())
+ require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String())
+
+ select {
+ case bridge := <-resultCh:
+ require.NoError(t, bridge.err)
+ require.NotNil(t, bridge.result)
+ require.Equal(t, "resp_bridge", bridge.result.RequestID)
+ require.Equal(t, 3, bridge.result.Usage.InputTokens)
+ require.Equal(t, 2, bridge.result.Usage.OutputTokens)
+ require.True(t, bridge.result.OpenAIWSMode)
+ case <-time.After(3 * time.Second):
+ t.Fatal("timed out waiting for bridge result")
+ }
+
+ require.NotNil(t, upstream.lastReq)
+ require.Equal(t, http.MethodPost, upstream.lastReq.Method)
+ require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists())
+ require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists())
+ require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
+}
+
+func TestOpenAIWSHTTPBridgeAcceptsFirstFrameAboveLegacy16MiB(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ sseBody := strings.Join([]string{
+ `data: {"type":"response.created","response":{"id":"resp_large_bridge","model":"gpt-5"}}`,
+ "",
+ `data: {"type":"response.completed","response":{"id":"resp_large_bridge","model":"gpt-5","usage":{"input_tokens":9,"output_tokens":1}}}`,
+ "",
+ }, "\n")
+ upstream := &httpUpstreamRecorder{resp: &http.Response{
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "Content-Type": []string{"text/event-stream"},
+ "x-request-id": []string{"rid_large_bridge"},
+ },
+ Body: io.NopCloser(strings.NewReader(sseBody)),
+ }}
+ cfg := &config.Config{
+ Gateway: config.GatewayConfig{
+ MaxLineSize: defaultMaxLineSize,
+ OpenAIWS: config.GatewayOpenAIWSConfig{
+ Enabled: true,
+ APIKeyEnabled: true,
+ ResponsesWebsocketsV2: true,
+ ClientReadLimitBytes: 64 * 1024 * 1024,
+ HTTPBridgeEnabled: true,
+ HTTPBridgeThresholdBytes: 15 * 1024 * 1024,
+ },
+ },
+ }
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: upstream,
+ toolCorrector: NewCodexToolCorrector(),
+ }
+ account := &Account{
+ ID: 9,
+ Name: "api-key",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{"api_key": "sk-upstream"},
+ Extra: map[string]any{
+ "openai_apikey_responses_websockets_v2_enabled": true,
+ },
+ Concurrency: 1,
+ Status: StatusActive,
+ }
+
+ payload := []byte(`{"type":"response.create","generate":true,"model":"gpt-5","stream":true,"input":"` + strings.Repeat("x", 17*1024*1024) + `"}`)
+ require.Greater(t, len(payload), 16*1024*1024)
+ require.Less(t, int64(len(payload)), ResolveOpenAIWSClientReadLimitBytes(cfg))
+
+ errCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
+ if err != nil {
+ errCh <- err
+ return
+ }
+ defer func() { _ = conn.CloseNow() }()
+ conn.SetReadLimit(ResolveOpenAIWSClientReadLimitBytes(cfg))
+
+ readCtx, cancelRead := context.WithTimeout(r.Context(), 10*time.Second)
+ msgType, firstMessage, err := conn.Read(readCtx)
+ cancelRead()
+ if err != nil {
+ errCh <- err
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ errCh <- NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "unexpected client websocket message type", nil)
+ return
+ }
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "codex_cli_rs/0.135.0")
+ ginCtx.Request = req
+
+ proxyCtx, cancelProxy := context.WithTimeout(r.Context(), 20*time.Second)
+ defer cancelProxy()
+ errCh <- svc.ProxyResponsesWebSocketFromClient(proxyCtx, ginCtx, conn, account, "sk-test", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 5*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() { _ = clientConn.CloseNow() }()
+
+ writeCtx, cancelWrite := context.WithTimeout(context.Background(), 20*time.Second)
+ err = clientConn.Write(writeCtx, coderws.MessageText, payload)
+ cancelWrite()
+ require.NoError(t, err)
+
+ var eventTypes []string
+ for {
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 10*time.Second)
+ msgType, event, readErr := clientConn.Read(readCtx)
+ cancelRead()
+ require.NoError(t, readErr)
+ require.Equal(t, coderws.MessageText, msgType)
+
+ eventType := gjson.GetBytes(event, "type").String()
+ eventTypes = append(eventTypes, eventType)
+ if eventType == "response.completed" {
+ break
+ }
+ }
+ require.Contains(t, eventTypes, "response.created")
+ require.Contains(t, eventTypes, "response.completed")
+
+ require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
+ select {
+ case proxyErr := <-errCh:
+ require.NoError(t, proxyErr)
+ case <-time.After(10 * time.Second):
+ t.Fatal("timed out waiting for websocket bridge proxy to finish")
+ }
+
+ require.NotNil(t, upstream.lastReq)
+ require.Equal(t, http.MethodPost, upstream.lastReq.Method)
+ require.Greater(t, len(upstream.lastBody), 16*1024*1024)
+ require.False(t, gjson.GetBytes(upstream.lastBody, "type").Exists())
+ require.False(t, gjson.GetBytes(upstream.lastBody, "generate").Exists())
+ require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool())
+ require.Equal(t, "gpt-5", gjson.GetBytes(upstream.lastBody, "model").String())
+}
+
+func TestOpenAIWSHTTPBridgeKeepsContinuationFramesOnHTTPWithoutPreviousResponseID(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+
+ firstSSEBody := strings.Join([]string{
+ `data: {"type":"response.completed","response":{"id":"resp_bridge_first","model":"gpt-5.1","output":[{"type":"function_call","id":"fc_bridge_1","call_id":"call_bridge_1","name":"shell","arguments":"{}"}],"usage":{"input_tokens":9,"output_tokens":1}}}`,
+ "",
+ }, "\n")
+ secondSSEBody := strings.Join([]string{
+ `data: {"type":"response.completed","response":{"id":"resp_bridge_second","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`,
+ "",
+ }, "\n")
+ upstream := &httpUpstreamRecorder{responses: []*http.Response{
+ {
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "Content-Type": []string{"text/event-stream"},
+ },
+ Body: io.NopCloser(strings.NewReader(firstSSEBody)),
+ },
+ {
+ StatusCode: http.StatusOK,
+ Header: http.Header{
+ "Content-Type": []string{"text/event-stream"},
+ },
+ Body: io.NopCloser(strings.NewReader(secondSSEBody)),
+ },
+ }}
+ cfg := &config.Config{}
+ cfg.Security.URLAllowlist.Enabled = false
+ cfg.Security.URLAllowlist.AllowInsecureHTTP = true
+ cfg.Gateway.OpenAIWS.Enabled = true
+ cfg.Gateway.OpenAIWS.OAuthEnabled = true
+ cfg.Gateway.OpenAIWS.APIKeyEnabled = true
+ cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true
+ cfg.Gateway.OpenAIWS.HTTPBridgeEnabled = true
+ cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes = 1
+ cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1
+ cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0
+ cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1
+ cfg.Gateway.OpenAIWS.QueueLimitPerConn = 8
+ cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3
+ cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3
+
+ captureConn := &openAIWSCaptureConn{}
+ captureDialer := &openAIWSCaptureDialer{conn: captureConn}
+ pool := newOpenAIWSConnPool(cfg)
+ pool.setClientDialerForTest(captureDialer)
+
+ svc := &OpenAIGatewayService{
+ cfg: cfg,
+ httpUpstream: upstream,
+ cache: &stubGatewayCache{},
+ openaiWSResolver: NewOpenAIWSProtocolResolver(cfg),
+ toolCorrector: NewCodexToolCorrector(),
+ openaiWSPool: pool,
+ }
+ account := &Account{
+ ID: 19,
+ Name: "api-key-bridge-handoff",
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Credentials: map[string]any{"api_key": "sk-upstream"},
+ Extra: map[string]any{
+ "responses_websockets_v2_enabled": true,
+ },
+ Concurrency: 1,
+ Status: StatusActive,
+ Schedulable: true,
+ }
+
+ errCh := make(chan error, 1)
+ wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ conn, err := coderws.Accept(w, r, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover})
+ if err != nil {
+ errCh <- err
+ return
+ }
+ defer func() { _ = conn.CloseNow() }()
+
+ readCtx, cancelRead := context.WithTimeout(r.Context(), 3*time.Second)
+ msgType, firstMessage, err := conn.Read(readCtx)
+ cancelRead()
+ if err != nil {
+ errCh <- err
+ return
+ }
+ if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
+ errCh <- NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "unexpected client websocket message type", nil)
+ return
+ }
+
+ rec := httptest.NewRecorder()
+ ginCtx, _ := gin.CreateTestContext(rec)
+ req := r.Clone(r.Context())
+ req.Header = req.Header.Clone()
+ req.Header.Set("User-Agent", "codex_cli_rs/0.135.0")
+ ginCtx.Request = req
+
+ errCh <- svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "sk-test", firstMessage, nil)
+ }))
+ defer wsServer.Close()
+
+ dialCtx, cancelDial := context.WithTimeout(context.Background(), 3*time.Second)
+ clientConn, _, err := coderws.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil)
+ cancelDial()
+ require.NoError(t, err)
+ defer func() { _ = clientConn.CloseNow() }()
+
+ writeMessage := func(payload string) {
+ writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancelWrite()
+ require.NoError(t, clientConn.Write(writeCtx, coderws.MessageText, []byte(payload)))
+ }
+ readMessage := func() []byte {
+ readCtx, cancelRead := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancelRead()
+ msgType, event, readErr := clientConn.Read(readCtx)
+ require.NoError(t, readErr)
+ require.Equal(t, coderws.MessageText, msgType)
+ return event
+ }
+
+ writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":true,"input":"first"}`)
+ firstTurnEvent := readMessage()
+ require.Equal(t, "response.completed", gjson.GetBytes(firstTurnEvent, "type").String())
+ require.Equal(t, "resp_bridge_first", gjson.GetBytes(firstTurnEvent, "response.id").String())
+
+ writeMessage(`{"type":"response.create","model":"gpt-5.1","stream":false,"previous_response_id":"resp_bridge_first","input":[{"type":"function_call_output","call_id":"call_bridge_1","output":"ok"}]}`)
+ secondTurnEvent := readMessage()
+ require.Equal(t, "response.completed", gjson.GetBytes(secondTurnEvent, "type").String())
+ require.Equal(t, "resp_bridge_second", gjson.GetBytes(secondTurnEvent, "response.id").String())
+
+ require.NoError(t, clientConn.Close(coderws.StatusNormalClosure, "done"))
+ select {
+ case proxyErr := <-errCh:
+ require.NoError(t, proxyErr)
+ case <-time.After(5 * time.Second):
+ t.Fatal("timed out waiting for websocket bridge proxy to finish")
+ }
+
+ require.Len(t, upstream.bodies, 2, "进入 HTTP bridge 后同一客户端 WS 连接内应保持 HTTP/SSE bridge")
+ require.False(t, gjson.GetBytes(upstream.bodies[0], "previous_response_id").Exists())
+ require.False(t, gjson.GetBytes(upstream.bodies[1], "previous_response_id").Exists())
+ secondInput := gjson.GetBytes(upstream.bodies[1], "input").Array()
+ require.Len(t, secondInput, 3)
+ require.Equal(t, "first", secondInput[0].String())
+ require.Equal(t, "function_call", secondInput[1].Get("type").String())
+ require.Equal(t, "call_bridge_1", secondInput[1].Get("call_id").String())
+ require.Equal(t, "function_call_output", secondInput[2].Get("type").String())
+ require.Equal(t, "call_bridge_1", secondInput[2].Get("call_id").String())
+ require.Equal(t, 0, captureDialer.DialCount())
+ require.Empty(t, captureConn.writes)
+}
diff --git a/backend/internal/service/openai_ws_ratelimit_signal_test.go b/backend/internal/service/openai_ws_ratelimit_signal_test.go
index 4ee85a3a..a3673d74 100644
--- a/backend/internal/service/openai_ws_ratelimit_signal_test.go
+++ b/backend/internal/service/openai_ws_ratelimit_signal_test.go
@@ -338,6 +338,9 @@ func TestOpenAIGatewayService_ProxyResponsesWebSocketFromClient_ErrorEventUsageL
select {
case serverErr := <-serverErrCh:
require.Error(t, serverErr)
+ var failoverErr *UpstreamFailoverError
+ require.ErrorAs(t, serverErr, &failoverErr)
+ require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode)
require.Len(t, repo.rateLimitCalls, 1)
require.WithinDuration(t, time.Unix(resetAt, 0), repo.rateLimitCalls[0], 2*time.Second)
case <-time.After(5 * time.Second):
diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay.go b/backend/internal/service/openai_ws_v2/passthrough_relay.go
index 2b7e2add..6aba3b7d 100644
--- a/backend/internal/service/openai_ws_v2/passthrough_relay.go
+++ b/backend/internal/service/openai_ws_v2/passthrough_relay.go
@@ -25,6 +25,7 @@ type Usage struct {
OutputTokens int
CacheCreationInputTokens int
CacheReadInputTokens int
+ ImageOutputTokens int
}
type RelayResult struct {
@@ -55,14 +56,18 @@ type RelayExit struct {
}
type RelayOptions struct {
- WriteTimeout time.Duration
- IdleTimeout time.Duration
- UpstreamDrainTimeout time.Duration
- FirstMessageType coderws.MessageType
- OnUsageParseFailure func(eventType string, usageRaw string)
- OnTurnComplete func(turn RelayTurnResult)
- OnTrace func(event RelayTraceEvent)
- Now func() time.Time
+ WriteTimeout time.Duration
+ IdleTimeout time.Duration
+ UpstreamDrainTimeout time.Duration
+ FirstMessageType coderws.MessageType
+ FirstMessageSent bool
+ StartClientAfterFirstDownstream bool
+ OnUsageParseFailure func(eventType string, usageRaw string)
+ OnTurnComplete func(turn RelayTurnResult)
+ BeforeWriteClient func(msgType coderws.MessageType, payload []byte, wroteDownstream bool) error
+ ReadClientFrame func(ctx context.Context, clientConn FrameConn) (coderws.MessageType, []byte, error)
+ OnTrace func(event RelayTraceEvent)
+ Now func() time.Time
}
type RelayTraceEvent struct {
@@ -170,29 +175,47 @@ func Relay(
MessageType: relayMessageTypeString(firstMessageType),
})
- if err := writeUpstream(firstMessageType, firstClientMessage); err != nil {
- result.Duration = nowFn().Sub(startAt)
+ if options.FirstMessageSent {
emitRelayTrace(onTrace, RelayTraceEvent{
- Stage: "write_first_message_failed",
+ Stage: "write_first_message_skipped",
+ Direction: "client_to_upstream",
+ MessageType: relayMessageTypeString(firstMessageType),
+ PayloadBytes: len(firstClientMessage),
+ })
+ } else {
+ if err := writeUpstream(firstMessageType, firstClientMessage); err != nil {
+ result.Duration = nowFn().Sub(startAt)
+ emitRelayTrace(onTrace, RelayTraceEvent{
+ Stage: "write_first_message_failed",
+ Direction: "client_to_upstream",
+ MessageType: relayMessageTypeString(firstMessageType),
+ PayloadBytes: len(firstClientMessage),
+ Error: err.Error(),
+ })
+ return result, &RelayExit{Stage: "write_upstream", Err: err}
+ }
+ emitRelayTrace(onTrace, RelayTraceEvent{
+ Stage: "write_first_message_ok",
Direction: "client_to_upstream",
MessageType: relayMessageTypeString(firstMessageType),
PayloadBytes: len(firstClientMessage),
- Error: err.Error(),
})
- return result, &RelayExit{Stage: "write_upstream", Err: err}
}
clientToUpstreamFrames.Add(1)
- emitRelayTrace(onTrace, RelayTraceEvent{
- Stage: "write_first_message_ok",
- Direction: "client_to_upstream",
- MessageType: relayMessageTypeString(firstMessageType),
- PayloadBytes: len(firstClientMessage),
- })
markActivity()
exitCh := make(chan relayExitSignal, 3)
dropDownstreamWrites := atomic.Bool{}
- go runClientToUpstream(relayCtx, clientConn, writeUpstream, markActivity, clientToUpstreamFrames, onTrace, exitCh)
+ clientReaderStarted := atomic.Bool{}
+ startClientReader := func() {
+ if !clientReaderStarted.CompareAndSwap(false, true) {
+ return
+ }
+ go runClientToUpstream(relayCtx, clientConn, options.ReadClientFrame, writeUpstream, markActivity, clientToUpstreamFrames, onTrace, exitCh)
+ }
+ if !options.StartClientAfterFirstDownstream {
+ startClientReader()
+ }
go runUpstreamToClient(
relayCtx,
upstreamConn,
@@ -202,6 +225,12 @@ func Relay(
state,
options.OnUsageParseFailure,
options.OnTurnComplete,
+ options.BeforeWriteClient,
+ func() {
+ if options.StartClientAfterFirstDownstream {
+ startClientReader()
+ }
+ },
&dropDownstreamWrites,
upstreamToClientFrames,
droppedDownstreamFrames,
@@ -230,7 +259,9 @@ func Relay(
} else {
relayCancel()
_ = upstreamConn.Close()
- secondExit, hasSecondExit = waitRelayExit(exitCh, 200*time.Millisecond)
+ if clientReaderStarted.Load() {
+ secondExit, hasSecondExit = waitRelayExit(exitCh, 200*time.Millisecond)
+ }
}
if hasSecondExit {
combinedWroteDownstream = combinedWroteDownstream || secondExit.wroteDownstream
@@ -250,6 +281,14 @@ func Relay(
result.ClientToUpstreamFrames = clientToUpstreamFrames.Load()
result.UpstreamToClientFrames = upstreamToClientFrames.Load()
result.DroppedDownstreamFrames = droppedDownstreamFrames.Load()
+ if options.FirstMessageSent && firstExit.stage == "read_client" && firstExit.graceful {
+ emitRelayTrace(onTrace, RelayTraceEvent{
+ Stage: "relay_client_closed",
+ Graceful: true,
+ WroteDownstream: combinedWroteDownstream,
+ })
+ return result, nil
+ }
if firstExit.stage == "read_client" && firstExit.graceful {
stage := "client_disconnected"
exitErr := firstExit.err
@@ -310,6 +349,14 @@ func Relay(
WroteDownstream: combinedWroteDownstream,
}
}
+ if options.FirstMessageSent {
+ emitRelayTrace(onTrace, RelayTraceEvent{
+ Stage: "relay_client_closed",
+ Graceful: true,
+ WroteDownstream: combinedWroteDownstream,
+ })
+ return result, nil
+ }
emitRelayTrace(onTrace, RelayTraceEvent{
Stage: "relay_complete",
Graceful: true,
@@ -322,14 +369,20 @@ func Relay(
func runClientToUpstream(
ctx context.Context,
clientConn FrameConn,
+ readClientFrame func(context.Context, FrameConn) (coderws.MessageType, []byte, error),
writeUpstream func(msgType coderws.MessageType, payload []byte) error,
markActivity func(),
forwardedFrames *atomic.Int64,
onTrace func(event RelayTraceEvent),
exitCh chan<- relayExitSignal,
) {
+ if readClientFrame == nil {
+ readClientFrame = func(ctx context.Context, conn FrameConn) (coderws.MessageType, []byte, error) {
+ return conn.ReadFrame(ctx)
+ }
+ }
for {
- msgType, payload, err := clientConn.ReadFrame(ctx)
+ msgType, payload, err := readClientFrame(ctx, clientConn)
if err != nil {
emitRelayTrace(onTrace, RelayTraceEvent{
Stage: "read_client_failed",
@@ -368,6 +421,8 @@ func runUpstreamToClient(
state *relayState,
onUsageParseFailure func(eventType string, usageRaw string),
onTurnComplete func(turn RelayTurnResult),
+ beforeWriteClient func(msgType coderws.MessageType, payload []byte, wroteDownstream bool) error,
+ afterWriteClient func(),
dropDownstreamWrites *atomic.Bool,
forwardedFrames *atomic.Int64,
droppedFrames *atomic.Int64,
@@ -395,6 +450,24 @@ func runUpstreamToClient(
return
}
markActivity()
+ if beforeWriteClient != nil {
+ if err := beforeWriteClient(msgType, payload, wroteDownstream); err != nil {
+ emitRelayTrace(onTrace, RelayTraceEvent{
+ Stage: "upstream_message_rejected",
+ Direction: "upstream_to_client",
+ MessageType: relayMessageTypeString(msgType),
+ PayloadBytes: len(payload),
+ WroteDownstream: wroteDownstream,
+ Error: err.Error(),
+ })
+ exitCh <- relayExitSignal{
+ stage: "upstream_message",
+ err: err,
+ wroteDownstream: wroteDownstream,
+ }
+ return
+ }
+ }
observedEvent := observedUpstreamEvent{}
switch msgType {
case coderws.MessageText:
@@ -438,6 +511,9 @@ func runUpstreamToClient(
return
}
wroteDownstream = true
+ if afterWriteClient != nil {
+ afterWriteClient()
+ }
if forwardedFrames != nil {
forwardedFrames.Add(1)
}
@@ -681,8 +757,21 @@ func parseUsageAndAccumulate(
}
inputResult := gjson.GetBytes(message, "response.usage.input_tokens")
+ if !inputResult.Exists() {
+ inputResult = gjson.GetBytes(message, "response.usage.prompt_tokens")
+ }
outputResult := gjson.GetBytes(message, "response.usage.output_tokens")
+ if !outputResult.Exists() {
+ outputResult = gjson.GetBytes(message, "response.usage.completion_tokens")
+ }
cachedResult := gjson.GetBytes(message, "response.usage.input_tokens_details.cached_tokens")
+ if !cachedResult.Exists() {
+ cachedResult = gjson.GetBytes(message, "response.usage.prompt_tokens_details.cached_tokens")
+ }
+ imageTokens := usageResult.Get("output_tokens_details.image_tokens").Int()
+ if imageTokens == 0 {
+ imageTokens = usageResult.Get("completion_tokens_details.image_tokens").Int()
+ }
inputTokens, inputOK := parseUsageIntField(inputResult, true)
outputTokens, outputOK := parseUsageIntField(outputResult, true)
@@ -696,14 +785,18 @@ func parseUsageAndAccumulate(
return Usage{}
}
parsedUsage := Usage{
- InputTokens: inputTokens,
- OutputTokens: outputTokens,
- CacheReadInputTokens: cachedTokens,
+ InputTokens: inputTokens,
+ OutputTokens: outputTokens,
+ CacheCreationInputTokens: int(usageResult.Get("cache_creation_input_tokens").Int()),
+ CacheReadInputTokens: cachedTokens,
+ ImageOutputTokens: int(imageTokens),
}
state.usage.InputTokens += parsedUsage.InputTokens
state.usage.OutputTokens += parsedUsage.OutputTokens
+ state.usage.CacheCreationInputTokens += parsedUsage.CacheCreationInputTokens
state.usage.CacheReadInputTokens += parsedUsage.CacheReadInputTokens
+ state.usage.ImageOutputTokens += parsedUsage.ImageOutputTokens
return parsedUsage
}
@@ -765,7 +858,7 @@ func isTerminalEvent(eventType string) bool {
func shouldParseUsage(eventType string) bool {
switch eventType {
- case "response.completed", "response.done", "response.failed":
+ case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled":
return true
default:
return false
diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go
index 123e10ce..13c51f66 100644
--- a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go
+++ b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go
@@ -45,6 +45,7 @@ func TestRunClientToUpstream_ErrorPaths(t *testing.T) {
runClientToUpstream(
context.Background(),
newPassthroughTestFrameConn(nil, true),
+ nil,
func(_ coderws.MessageType, _ []byte) error { return nil },
func() {},
nil,
@@ -65,6 +66,7 @@ func TestRunClientToUpstream_ErrorPaths(t *testing.T) {
newPassthroughTestFrameConn([]passthroughTestFrame{
{msgType: coderws.MessageText, payload: []byte(`{"x":1}`)},
}, true),
+ nil,
func(_ coderws.MessageType, _ []byte) error { return errors.New("boom") },
func() {},
nil,
@@ -87,6 +89,7 @@ func TestRunClientToUpstream_ErrorPaths(t *testing.T) {
newPassthroughTestFrameConn([]passthroughTestFrame{
{msgType: coderws.MessageText, payload: []byte(`{"x":1}`)},
}, true),
+ nil,
func(_ coderws.MessageType, _ []byte) error { return nil },
func() {},
forwarded,
@@ -120,6 +123,8 @@ func TestRunUpstreamToClient_ErrorAndDropPaths(t *testing.T) {
&relayState{},
nil,
nil,
+ nil,
+ nil,
drop,
nil,
nil,
@@ -149,6 +154,8 @@ func TestRunUpstreamToClient_ErrorAndDropPaths(t *testing.T) {
&relayState{},
nil,
nil,
+ nil,
+ nil,
drop,
nil,
nil,
@@ -181,6 +188,8 @@ func TestRunUpstreamToClient_ErrorAndDropPaths(t *testing.T) {
&relayState{},
nil,
nil,
+ nil,
+ nil,
drop,
nil,
dropped,
@@ -291,20 +300,41 @@ func TestParseUsageAndEnrichCoverage(t *testing.T) {
require.Equal(t, 0, state.usage.OutputTokens)
require.Equal(t, 0, state.usage.CacheReadInputTokens)
- parseUsageAndAccumulate(state, []byte(`{"type":"response.completed","response":{"usage":{"input_tokens":2,"output_tokens":1,"input_tokens_details":{"cached_tokens":1}}}}`), "response.completed", nil)
+ parseUsageAndAccumulate(state, []byte(`{"type":"response.completed","response":{"usage":{"input_tokens":2,"output_tokens":1,"input_tokens_details":{"cached_tokens":1},"cache_creation_input_tokens":4,"output_tokens_details":{"image_tokens":3}}}}`), "response.completed", nil)
require.Equal(t, 2, state.usage.InputTokens)
require.Equal(t, 1, state.usage.OutputTokens)
require.Equal(t, 1, state.usage.CacheReadInputTokens)
+ require.Equal(t, 4, state.usage.CacheCreationInputTokens)
+ require.Equal(t, 3, state.usage.ImageOutputTokens)
result := &RelayResult{}
enrichResult(result, state, 5*time.Millisecond)
require.Equal(t, state.usage.InputTokens, result.Usage.InputTokens)
+ require.Equal(t, state.usage.CacheCreationInputTokens, result.Usage.CacheCreationInputTokens)
+ require.Equal(t, state.usage.ImageOutputTokens, result.Usage.ImageOutputTokens)
require.Equal(t, 5*time.Millisecond, result.Duration)
parseUsageAndAccumulate(state, []byte(`{"type":"response.in_progress","response":{"usage":{"input_tokens":9}}}`), "response.in_progress", nil)
require.Equal(t, 2, state.usage.InputTokens)
enrichResult(nil, state, 0)
}
+func TestParseUsageAndAccumulateAcceptsChatUsageAliases(t *testing.T) {
+ t.Parallel()
+
+ state := &relayState{}
+ got := parseUsageAndAccumulate(
+ state,
+ []byte(`{"type":"response.done","response":{"usage":{"prompt_tokens":12,"completion_tokens":6,"prompt_tokens_details":{"cached_tokens":4},"completion_tokens_details":{"image_tokens":2}}}}`),
+ "response.done",
+ nil,
+ )
+ require.Equal(t, 12, got.InputTokens)
+ require.Equal(t, 6, got.OutputTokens)
+ require.Equal(t, 4, got.CacheReadInputTokens)
+ require.Equal(t, 2, got.ImageOutputTokens)
+ require.Equal(t, got, state.usage)
+}
+
func TestEmitTurnCompleteCoverage(t *testing.T) {
t.Parallel()
@@ -368,6 +398,23 @@ func TestIsTokenEventCoverageBranches(t *testing.T) {
require.True(t, isTokenEvent("response.done"))
}
+func TestShouldParseUsageTerminalEvents(t *testing.T) {
+ t.Parallel()
+
+ for _, eventType := range []string{
+ "response.completed",
+ "response.done",
+ "response.failed",
+ "response.incomplete",
+ "response.cancelled",
+ "response.canceled",
+ } {
+ require.True(t, shouldParseUsage(eventType), eventType)
+ }
+ require.False(t, shouldParseUsage("response.output_text.delta"))
+ require.False(t, shouldParseUsage(""))
+}
+
func TestRelayTurnTimingHelpersCoverage(t *testing.T) {
t.Parallel()
diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter.go b/backend/internal/service/openai_ws_v2_passthrough_adapter.go
index 347a3b44..c93d0981 100644
--- a/backend/internal/service/openai_ws_v2_passthrough_adapter.go
+++ b/backend/internal/service/openai_ws_v2_passthrough_adapter.go
@@ -312,6 +312,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
// goroutine)和 OnTurnComplete / final result(runUpstreamToClient
// goroutine)之间同步当前 turn 的 usage metadata。
usageMeta.initFromFirstFrame(firstClientMessage)
+ promptCacheKey := strings.TrimSpace(gjson.GetBytes(firstClientMessage, "prompt_cache_key").String())
wsURL, err := s.buildOpenAIResponsesWSURL(account)
if err != nil {
@@ -338,7 +339,13 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
if s.cfg != nil && s.cfg.Gateway.ForceCodexCLI {
isCodexCLI = true
}
- headers, _ := s.buildOpenAIWSHeaders(c, account, token, wsDecision, isCodexCLI, "", "", "")
+ turnState := ""
+ turnMetadata := ""
+ if c != nil {
+ turnState = strings.TrimSpace(c.GetHeader(openAIWSTurnStateHeader))
+ turnMetadata = strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader))
+ }
+ headers, _ := s.buildOpenAIWSHeaders(c, account, token, wsDecision, isCodexCLI, turnState, turnMetadata, promptCacheKey)
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
@@ -359,6 +366,13 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
statusCode,
truncateOpenAIWSLogValue(err.Error(), openAIWSLogValueMaxLen),
)
+ if statusCode == http.StatusTooManyRequests {
+ s.persistOpenAIWSRateLimitSignal(ctx, account, handshakeHeaders, nil, "rate_limit_exceeded", "rate_limit_error", strings.TrimSpace(err.Error()))
+ return &UpstreamFailoverError{
+ StatusCode: http.StatusTooManyRequests,
+ ResponseHeaders: cloneHeader(handshakeHeaders),
+ }
+ }
return s.mapOpenAIWSPassthroughDialError(err, statusCode, handshakeHeaders)
}
defer func() {
@@ -456,15 +470,46 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
cancel()
},
}
+ upstreamFirstMessageSent := false
+ firstWriteCtx, cancelFirstWrite := context.WithTimeout(ctx, s.openAIWSWriteTimeout())
+ firstWriteErr := upstreamFrameConn.WriteFrame(firstWriteCtx, coderws.MessageText, firstClientMessage)
+ cancelFirstWrite()
+ if firstWriteErr != nil {
+ return wrapOpenAIWSIngressTurnError(
+ "write_upstream",
+ fmt.Errorf("write first upstream websocket request: %w", firstWriteErr),
+ false,
+ )
+ }
+ upstreamFirstMessageSent = true
+
+ readNextClientFrame := func(readCtx context.Context, conn openaiwsv2.FrameConn) (coderws.MessageType, []byte, error) {
+ for {
+ msgType, payload, readErr := conn.ReadFrame(readCtx)
+ if readErr != nil {
+ return msgType, payload, readErr
+ }
+ if msgType == coderws.MessageText && strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" {
+ return msgType, payload, nil
+ }
+ if writeErr := upstreamFrameConn.WriteFrame(readCtx, msgType, payload); writeErr != nil {
+ return msgType, payload, writeErr
+ }
+ }
+ }
+
relayResult, relayExit := openaiwsv2.RunEntry(openaiwsv2.EntryInput{
Ctx: ctx,
ClientConn: policyClientConn,
UpstreamConn: upstreamFrameConn,
FirstClientMessage: firstClientMessage,
Options: openaiwsv2.RelayOptions{
- WriteTimeout: s.openAIWSWriteTimeout(),
- IdleTimeout: s.openAIWSPassthroughIdleTimeout(),
- FirstMessageType: coderws.MessageText,
+ WriteTimeout: s.openAIWSWriteTimeout(),
+ IdleTimeout: s.openAIWSPassthroughIdleTimeout(),
+ FirstMessageType: coderws.MessageText,
+ FirstMessageSent: upstreamFirstMessageSent,
+ StartClientAfterFirstDownstream: true,
+ ReadClientFrame: readNextClientFrame,
OnUsageParseFailure: func(eventType string, usageRaw string) {
logOpenAIWSV2Passthrough(
"usage_parse_failed event_type=%s usage_raw=%s",
@@ -481,6 +526,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
OutputTokens: turn.Usage.OutputTokens,
CacheCreationInputTokens: turn.Usage.CacheCreationInputTokens,
CacheReadInputTokens: turn.Usage.CacheReadInputTokens,
+ ImageOutputTokens: turn.Usage.ImageOutputTokens,
},
Model: turn.RequestModel,
ServiceTier: usageMeta.serviceTier.Load(),
@@ -507,6 +553,31 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
hooks.AfterTurn(turnNo, turnResult, nil)
}
},
+ BeforeWriteClient: func(msgType coderws.MessageType, payload []byte, wroteDownstream bool) error {
+ if msgType != coderws.MessageText || wroteDownstream {
+ return nil
+ }
+ if eventType, _, _ := parseOpenAIWSEventEnvelope(payload); eventType != "error" {
+ return nil
+ }
+ errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(payload)
+ if !isOpenAIWSRateLimitError(errCodeRaw, errTypeRaw, errMsgRaw) {
+ return nil
+ }
+ s.persistOpenAIWSRateLimitSignal(ctx, account, handshakeHeaders, payload, errCodeRaw, errTypeRaw, errMsgRaw)
+ logOpenAIWSV2Passthrough(
+ "relay_rate_limit_failover account_id=%d err_code=%s err_type=%s err_message=%s",
+ account.ID,
+ truncateOpenAIWSLogValue(errCodeRaw, openAIWSLogValueMaxLen),
+ truncateOpenAIWSLogValue(errTypeRaw, openAIWSLogValueMaxLen),
+ truncateOpenAIWSLogValue(errMsgRaw, openAIWSLogValueMaxLen),
+ )
+ return &UpstreamFailoverError{
+ StatusCode: http.StatusTooManyRequests,
+ ResponseBody: append([]byte(nil), payload...),
+ ResponseHeaders: cloneHeader(handshakeHeaders),
+ }
+ },
OnTrace: func(event openaiwsv2.RelayTraceEvent) {
logOpenAIWSV2Passthrough(
"relay_trace account_id=%d stage=%s direction=%s msg_type=%s bytes=%d graceful=%v wrote_downstream=%v err=%s",
@@ -530,6 +601,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough(
OutputTokens: relayResult.Usage.OutputTokens,
CacheCreationInputTokens: relayResult.Usage.CacheCreationInputTokens,
CacheReadInputTokens: relayResult.Usage.CacheReadInputTokens,
+ ImageOutputTokens: relayResult.Usage.ImageOutputTokens,
},
Model: relayResult.RequestModel,
ServiceTier: usageMeta.serviceTier.Load(),
diff --git a/backend/internal/service/ops_models.go b/backend/internal/service/ops_models.go
index ba735346..0bbe4220 100644
--- a/backend/internal/service/ops_models.go
+++ b/backend/internal/service/ops_models.go
@@ -64,6 +64,10 @@ type OpsErrorLog struct {
RequestedModel string `json:"requested_model"`
UpstreamModel string `json:"upstream_model"`
RequestType *int16 `json:"request_type"`
+
+ // 关联 api_key 名称(LEFT JOIN api_keys 取得;软删只覆盖 key 列,name 保留,故已删 key 仍有原名)。
+ APIKeyName string `json:"api_key_name,omitempty"`
+ APIKeyDeleted bool `json:"api_key_deleted,omitempty"`
}
type OpsErrorLogDetail struct {
@@ -87,6 +91,15 @@ type OpsErrorLogDetail struct {
// vNext metric semantics
IsBusinessLimited bool `json:"is_business_limited"`
+
+ // Deleted key owner info (populated when INVALID_API_KEY and key was previously deleted)
+ AttemptedKeyPrefix string `json:"attempted_key_prefix,omitempty"`
+ DeletedKeyOwnerUserID *int64 `json:"deleted_key_owner_user_id,omitempty"`
+ DeletedKeyOwnerEmail string `json:"deleted_key_owner_email,omitempty"`
+ DeletedKeyName string `json:"deleted_key_name,omitempty"`
+
+ // Bound (non-deleted) key prefix, snapshotted at error time; mutually exclusive with AttemptedKeyPrefix.
+ APIKeyPrefix string `json:"api_key_prefix,omitempty"`
}
type OpsErrorLogFilter struct {
@@ -99,7 +112,7 @@ type OpsErrorLogFilter struct {
StatusCodes []int
StatusCodesOther bool
- Phase string
+ Phase string // Special: Phase=="upstream" bypasses status>=400 clause; do not set together with ErrorPhasesAny.
Owner string
Source string
Resolved *bool
@@ -110,6 +123,32 @@ type OpsErrorLogFilter struct {
RequestID string
ClientRequestID string
+ // User-scoped filters (used by the user-facing error requests endpoint and
+ // by admin drill-down from the usage page).
+ UserID *int64
+ APIKeyID *int64
+
+ // MatchDeletedKeyOwner: 用户侧专用。UserID 设置且为 true 时,归属从 user_id=UserID
+ // 放宽为 (user_id=UserID OR deleted_key_owner_user_id=UserID),使原所有者能看到
+ // 自己「已删除 key 认证失败」的记录。admin 路径不设此开关 → 行为不变。
+ MatchDeletedKeyOwner bool
+
+ // Model matches against requested_model first, then model.
+ Model string
+ // ModelFuzzy 为 true 时 Model 走 ILIKE 模糊匹配(仅用户端启用);false(默认)保持精确 =,管理端语义不变。
+ ModelFuzzy bool
+
+ // ExcludeCountTokens drops count_tokens probe errors (is_count_tokens=true).
+ ExcludeCountTokens bool
+
+ // ErrorPhasesAny / ErrorTypesAny add plain ANY() filters WITHOUT touching the
+ // special-cased single `Phase` field (only Phase=="upstream" bypasses the status>=400 clause).
+ // NOTE: these ANY filters do NOT bypass status>=400; records with error_phase='upstream'
+ // but status_code<400 (recovered upstream errors) remain excluded.
+ // Used to map user-facing coarse categories to backend conditions.
+ ErrorPhasesAny []string
+ ErrorTypesAny []string
+
// View controls error categorization for list endpoints.
// - errors: show actionable errors (exclude business-limited / 429 / 529)
// - excluded: only show excluded errors
diff --git a/backend/internal/service/ops_port.go b/backend/internal/service/ops_port.go
index 30145ed3..0cba300d 100644
--- a/backend/internal/service/ops_port.go
+++ b/backend/internal/service/ops_port.go
@@ -10,6 +10,8 @@ type OpsRepository interface {
BatchInsertErrorLogs(ctx context.Context, inputs []*OpsInsertErrorLogInput) (int64, error)
ListErrorLogs(ctx context.Context, filter *OpsErrorLogFilter) (*OpsErrorLogList, error)
GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLogDetail, error)
+ // LookupDeletedKeyAudit 按明文 key 反查最近一条已删除 key 审计;未命中返回 (nil, nil)。
+ LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error)
ListRequestDetails(ctx context.Context, filter *OpsRequestDetailFilter) ([]*OpsRequestDetail, int64, error)
BatchInsertSystemLogs(ctx context.Context, inputs []*OpsInsertSystemLogInput) (int64, error)
ListSystemLogs(ctx context.Context, filter *OpsSystemLogFilter) (*OpsSystemLogList, error)
@@ -61,6 +63,12 @@ type OpsRepository interface {
GetLatestDailyBucketDate(ctx context.Context) (time.Time, bool, error)
}
+// DeletedKeyAuditResult 是按明文 key 反查 deleted_api_key_audits 的结果。
+type DeletedKeyAuditResult struct {
+ UserID int64
+ KeyName string
+}
+
type OpsInsertErrorLogInput struct {
RequestID string
ClientRequestID string
@@ -118,6 +126,15 @@ type OpsInsertErrorLogInput struct {
TimeToFirstTokenMs *int64
CreatedAt time.Time
+
+ // 已删除 key 归因(仅 INVALID_API_KEY 认证失败时可能非空)
+ AttemptedKeyPrefix string // 提交 key 的脱敏前缀(前 8 位)
+ DeletedKeyOwnerUserID *int64 // 反查命中的原所有者 user_id
+ DeletedKeyName string // 反查命中的 key 名称
+
+ // 有效(未删除)key 报错时快照的 key 脱敏前缀(前 8 位);与 AttemptedKeyPrefix 互斥。
+ // 落库快照而非读时 JOIN:key 之后被删(key 列被 tombstone 覆盖)仍保留当时前缀。
+ APIKeyPrefix string
}
type OpsInsertSystemMetricsInput struct {
diff --git a/backend/internal/service/ops_repo_mock_test.go b/backend/internal/service/ops_repo_mock_test.go
index 4138ea77..5e33bffb 100644
--- a/backend/internal/service/ops_repo_mock_test.go
+++ b/backend/internal/service/ops_repo_mock_test.go
@@ -13,6 +13,7 @@ type opsRepoMock struct {
ListSystemLogsFn func(ctx context.Context, filter *OpsSystemLogFilter) (*OpsSystemLogList, error)
DeleteSystemLogsFn func(ctx context.Context, filter *OpsSystemLogCleanupFilter) (int64, error)
InsertSystemLogCleanupAuditFn func(ctx context.Context, input *OpsSystemLogCleanupAudit) error
+ LookupDeletedKeyAuditFn func(ctx context.Context, key string) (*DeletedKeyAuditResult, error)
}
func (m *opsRepoMock) InsertErrorLog(ctx context.Context, input *OpsInsertErrorLogInput) (int64, error) {
@@ -189,4 +190,11 @@ func (m *opsRepoMock) GetLatestDailyBucketDate(ctx context.Context) (time.Time,
return time.Time{}, false, nil
}
+func (m *opsRepoMock) LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error) {
+ if m.LookupDeletedKeyAuditFn != nil {
+ return m.LookupDeletedKeyAuditFn(ctx, key)
+ }
+ return nil, nil
+}
+
var _ OpsRepository = (*opsRepoMock)(nil)
diff --git a/backend/internal/service/ops_service.go b/backend/internal/service/ops_service.go
index 1cea72fa..a8c8a4bb 100644
--- a/backend/internal/service/ops_service.go
+++ b/backend/internal/service/ops_service.go
@@ -41,6 +41,11 @@ type OpsService struct {
// cleanupReloader 由 wire 在 OpsCleanupService 构造完成后通过 SetCleanupReloader 注入。
// 解耦避免 OpsService -> OpsCleanupService 的硬依赖(cleanup 也读 settings,会循环)。
cleanupReloader CleanupReloader
+
+ // quotaAutoPauseSink 由 wire 注入(通常是 SettingService.SetOpenAIQuotaAutoPauseSettings)。
+ // UpdateOpsAdvancedSettings 写入新配置后调用,把最新的 quota auto-pause 全局默认阈值
+ // 立即同步到调度热路径读取的内存缓存,避免下次请求才能感知新值。
+ quotaAutoPauseSink func(OpsOpenAIAccountQuotaAutoPauseSettings)
}
// CleanupReloader 由 OpsCleanupService 实现。
@@ -57,6 +62,16 @@ func (s *OpsService) SetCleanupReloader(r CleanupReloader) {
s.cleanupReloader = r
}
+// SetOpenAIQuotaAutoPauseSettingsSink 由 wire 注入,把最新的 quota auto-pause 全局默认
+// 阈值 push 到调度热路径读取的内存缓存。同 SetCleanupReloader 的解耦目的:避免 OpsService
+// 持有 *SettingService 引入循环依赖。
+func (s *OpsService) SetOpenAIQuotaAutoPauseSettingsSink(sink func(OpsOpenAIAccountQuotaAutoPauseSettings)) {
+ if s == nil {
+ return
+ }
+ s.quotaAutoPauseSink = sink
+}
+
func NewOpsService(
opsRepo OpsRepository,
settingRepo SettingRepository,
@@ -323,6 +338,50 @@ func (s *OpsService) GetErrorLogs(ctx context.Context, filter *OpsErrorLogFilter
return result, nil
}
+// ListUserErrorRequests 返回某个用户自己的错误请求(精简脱敏)。
+// 强制:仅当前用户、View=all(含业务限流/余额类)、排除 count_tokens 噪声。
+func (s *OpsService) ListUserErrorRequests(ctx context.Context, userID int64, filter *OpsErrorLogFilter) (*UserErrorRequestList, error) {
+ if filter == nil {
+ filter = &OpsErrorLogFilter{}
+ }
+ f := *filter // 拷贝快照,避免原地篡改调用方的 filter(slice 字段只读,浅拷贝足够)
+ filter = &f
+ uid := userID
+ filter.UserID = &uid
+ // 用户侧放宽归属:纳入「删 key 后认证失败」(user_id=NULL,靠 deleted_key_owner 归因)的记录。
+ filter.MatchDeletedKeyOwner = true
+ // APIKeyID 透传:保留 handler 传入的值。安全由 buildOpsErrorLogsWhere 的
+ // "user_id = 自己 AND api_key_id = X" 双重约束保证——传入他人 key 只会得到空集,无泄露。
+ filter.View = "all"
+ filter.ExcludeCountTokens = true
+ filter.ModelFuzzy = true // 用户端模型过滤走 ILIKE 模糊;管理端不设此字段,保持精确
+ // 防御:用户端不接受这些 admin-only / 特殊维度
+ filter.UserQuery = ""
+ filter.Owner = ""
+ filter.Source = ""
+ // 清空 Phase 是防御:Phase 是单值特殊字段,仅当其 == "upstream" 时 buildOpsErrorLogsWhere 才跳过 status>=400 子句。
+ // 用户端一律改走 category→ErrorPhasesAny/ErrorTypesAny(纯 ANY 过滤,不影响 status>=400 子句),
+ // 因此 recovered upstream(error_phase='upstream' 但 status<400,最终成功返回)记录对用户不可见——符合预期。
+ filter.Phase = ""
+
+ list, err := s.opsRepo.ListErrorLogs(ctx, filter)
+ if err != nil {
+ return nil, err
+ }
+ items := make([]*UserErrorRequest, 0, len(list.Errors))
+ for _, e := range list.Errors {
+ if r := ToUserErrorRequest(e); r != nil {
+ items = append(items, r)
+ }
+ }
+ return &UserErrorRequestList{
+ Items: items,
+ Total: list.Total,
+ Page: list.Page,
+ PageSize: list.PageSize,
+ }, nil
+}
+
func (s *OpsService) GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLogDetail, error) {
if err := s.RequireMonitoringEnabled(ctx); err != nil {
return nil, err
@@ -340,6 +399,39 @@ func (s *OpsService) GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLo
return detail, nil
}
+// GetUserErrorRequestDetail 返回某用户自己某条错误请求的脱敏详情(含 error_body)。
+// 安全:强制按用户归属校验;非本人记录一律返回 NotFound(不泄露存在性)。
+func (s *OpsService) GetUserErrorRequestDetail(ctx context.Context, userID, id int64) (*UserErrorRequestDetail, error) {
+ if s.opsRepo == nil {
+ return nil, infraerrors.NotFound("OPS_ERROR_NOT_FOUND", "ops error log not found")
+ }
+ if id <= 0 {
+ return nil, infraerrors.BadRequest("OPS_ERROR_INVALID_ID", "invalid error id")
+ }
+ detail, err := s.opsRepo.GetErrorLogByID(ctx, id)
+ if err != nil {
+ if errors.Is(err, sql.ErrNoRows) {
+ return nil, infraerrors.NotFound("OPS_ERROR_NOT_FOUND", "ops error log not found")
+ }
+ return nil, infraerrors.InternalServer("OPS_ERROR_LOAD_FAILED", "Failed to load ops error log").WithCause(err)
+ }
+ // 归属:直接归属(user_id)或经「已删除 key 归因」(deleted_key_owner_user_id)二者之一即可。
+ ownedDirectly := detail.UserID != nil && *detail.UserID == userID
+ ownedViaDeletedKey := detail.DeletedKeyOwnerUserID != nil && *detail.DeletedKeyOwnerUserID == userID
+ if !ownedDirectly && !ownedViaDeletedKey {
+ return nil, infraerrors.NotFound("OPS_ERROR_NOT_FOUND", "ops error log not found")
+ }
+ return ToUserErrorRequestDetail(detail), nil
+}
+
+// LookupDeletedKeyAudit 按明文 key 反查已删除 key 的原所有者;未命中或未启用返回 (nil, nil)。
+func (s *OpsService) LookupDeletedKeyAudit(ctx context.Context, key string) (*DeletedKeyAuditResult, error) {
+ if s.opsRepo == nil {
+ return nil, nil
+ }
+ return s.opsRepo.LookupDeletedKeyAudit(ctx, key)
+}
+
func (s *OpsService) UpdateErrorResolution(ctx context.Context, errorID int64, resolved bool, resolvedByUserID *int64) error {
if err := s.RequireMonitoringEnabled(ctx); err != nil {
return err
diff --git a/backend/internal/service/ops_service_user_error_test.go b/backend/internal/service/ops_service_user_error_test.go
new file mode 100644
index 00000000..9027ff07
--- /dev/null
+++ b/backend/internal/service/ops_service_user_error_test.go
@@ -0,0 +1,222 @@
+package service
+
+import (
+ "context"
+ "database/sql"
+ "testing"
+
+ infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+)
+
+type stubOpsRepoForUserErr struct {
+ OpsRepository // 嵌入接口,未实现的方法 panic,仅覆盖 ListErrorLogs
+ gotFilter *OpsErrorLogFilter
+
+ // GetErrorLogByID 控制字段
+ detailToReturn *OpsErrorLogDetail
+ detailErrToReturn error
+}
+
+func (s *stubOpsRepoForUserErr) ListErrorLogs(ctx context.Context, f *OpsErrorLogFilter) (*OpsErrorLogList, error) {
+ s.gotFilter = f
+ return &OpsErrorLogList{
+ Errors: []*OpsErrorLog{{
+ Phase: "request", Type: "rate_limit_error",
+ Model: "m", RequestedModel: "rm", StatusCode: 429,
+ Message: "secret", UserEmail: "a@b.c",
+ }},
+ Total: 1, Page: 1, PageSize: 20,
+ }, nil
+}
+
+func (s *stubOpsRepoForUserErr) GetErrorLogByID(ctx context.Context, id int64) (*OpsErrorLogDetail, error) {
+ if s.detailErrToReturn != nil {
+ return nil, s.detailErrToReturn
+ }
+ return s.detailToReturn, nil
+}
+
+func TestListUserErrorRequests_ForcesScopeAndRedacts(t *testing.T) {
+ stub := &stubOpsRepoForUserErr{}
+ svc := &OpsService{opsRepo: stub}
+ uid := int64(42)
+ kid := int64(7)
+ in := &OpsErrorLogFilter{UserID: nil, View: "errors", Phase: "upstream", APIKeyID: &kid}
+ out, err := svc.ListUserErrorRequests(context.Background(), uid, in)
+ if err != nil {
+ t.Fatal(err)
+ }
+ // 强制按用户
+ if stub.gotFilter.UserID == nil || *stub.gotFilter.UserID != uid {
+ t.Fatalf("UserID not forced: %+v", stub.gotFilter.UserID)
+ }
+ // 强制 View=all(含业务限流/余额)
+ if stub.gotFilter.View != "all" {
+ t.Fatalf("View not forced to all: %q", stub.gotFilter.View)
+ }
+ // 强制排除 count_tokens
+ if !stub.gotFilter.ExcludeCountTokens {
+ t.Fatal("ExcludeCountTokens not forced")
+ }
+ // 强制清空 Phase(防止 "upstream" 绕过 status>=400 子句 + 与 ErrorPhasesAny 双重约束)
+ if stub.gotFilter.Phase != "" {
+ t.Fatalf("Phase not cleared: %q", stub.gotFilter.Phase)
+ }
+ // APIKeyID 透传保留(用户可按自己 key 过滤;越权由 user_id AND api_key_id 双重防护)
+ if stub.gotFilter.APIKeyID == nil || *stub.gotFilter.APIKeyID != kid {
+ t.Fatalf("APIKeyID should be preserved, got %v", stub.gotFilter.APIKeyID)
+ }
+ // 调用方传入的 filter 不应被原地篡改(验证 shallow copy 隔离生效)
+ if in.View != "errors" || in.UserID != nil || in.Phase != "upstream" {
+ t.Fatalf("caller filter was mutated: View=%q UserID=%v Phase=%q", in.View, in.UserID, in.Phase)
+ }
+ // 脱敏:返回条目含 message 字段
+ if len(out.Items) != 1 || out.Items[0].Category != "rate_limit" || out.Items[0].Model != "rm" {
+ t.Fatalf("bad item: %+v", out.Items)
+ }
+}
+
+func TestGetUserErrorRequestDetail_OwnershipEnforced(t *testing.T) {
+ ownerUID := int64(999)
+ callerUID := int64(1)
+ upstreamStatus := 503
+
+ detail := &OpsErrorLogDetail{
+ OpsErrorLog: OpsErrorLog{
+ ID: 42,
+ Phase: "upstream",
+ Type: "api_error",
+ Model: "gpt-4",
+ RequestedModel: "gpt-4-turbo",
+ InboundEndpoint: "/v1/chat/completions",
+ StatusCode: 502,
+ Platform: "openai",
+ Message: "upstream failed",
+ UserID: &ownerUID,
+ },
+ ErrorBody: `{"error":"upstream"}`,
+ UpstreamStatusCode: &upstreamStatus,
+ }
+
+ stub := &stubOpsRepoForUserErr{detailToReturn: detail}
+ svc := &OpsService{opsRepo: stub}
+
+ // 越权调用(callerUID=1,但记录属于 ownerUID=999)→ 应返回 NotFound,detail 为 nil
+ got, err := svc.GetUserErrorRequestDetail(context.Background(), callerUID, 42)
+ if err == nil {
+ t.Fatal("expected error for unauthorized access, got nil")
+ }
+ if got != nil {
+ t.Fatalf("expected nil detail for unauthorized access, got %+v", got)
+ }
+ // 验证错误为 NotFound(不暴露存在性)
+ if !infraerrors.IsNotFound(err) {
+ t.Fatalf("expected NotFound error, got: %v", err)
+ }
+
+ // 合法调用(callerUID=999 = ownerUID)→ 应返回 non-nil detail
+ got2, err2 := svc.GetUserErrorRequestDetail(context.Background(), ownerUID, 42)
+ if err2 != nil {
+ t.Fatalf("expected no error for legitimate access, got %v", err2)
+ }
+ if got2 == nil {
+ t.Fatal("expected non-nil detail for legitimate access")
+ }
+ if got2.ID != 42 {
+ t.Errorf("want ID=42, got %d", got2.ID)
+ }
+ if got2.ErrorBody != `{"error":"upstream"}` {
+ t.Errorf("want ErrorBody=%q, got %q", `{"error":"upstream"}`, got2.ErrorBody)
+ }
+ if got2.UpstreamStatusCode == nil || *got2.UpstreamStatusCode != 503 {
+ t.Errorf("want UpstreamStatusCode=503, got %v", got2.UpstreamStatusCode)
+ }
+ if got2.Message != "upstream failed" {
+ t.Errorf("want Message=%q, got %q", "upstream failed", got2.Message)
+ }
+}
+
+func TestGetUserErrorRequestDetail_NotFound(t *testing.T) {
+ stub := &stubOpsRepoForUserErr{detailErrToReturn: sql.ErrNoRows}
+ svc := &OpsService{opsRepo: stub}
+
+ got, err := svc.GetUserErrorRequestDetail(context.Background(), 1, 999)
+ if err == nil {
+ t.Fatal("expected error for not found, got nil")
+ }
+ if got != nil {
+ t.Fatalf("expected nil detail, got %+v", got)
+ }
+}
+
+func TestGetUserErrorRequestDetail_InvalidID(t *testing.T) {
+ stub := &stubOpsRepoForUserErr{}
+ svc := &OpsService{opsRepo: stub}
+
+ _, err := svc.GetUserErrorRequestDetail(context.Background(), 1, 0)
+ if err == nil {
+ t.Fatal("expected error for id=0")
+ }
+ _, err = svc.GetUserErrorRequestDetail(context.Background(), 1, -5)
+ if err == nil {
+ t.Fatal("expected error for id=-5")
+ }
+}
+
+func TestListUserErrorRequests_EnablesMatchDeletedKeyOwner(t *testing.T) {
+ stub := &stubOpsRepoForUserErr{}
+ svc := &OpsService{opsRepo: stub}
+ uid := int64(42)
+
+ if _, err := svc.ListUserErrorRequests(context.Background(), uid, &OpsErrorLogFilter{}); err != nil {
+ t.Fatal(err)
+ }
+ if stub.gotFilter == nil || !stub.gotFilter.MatchDeletedKeyOwner {
+ t.Fatal("ListUserErrorRequests should enable MatchDeletedKeyOwner for the user scope")
+ }
+}
+
+func TestGetUserErrorRequestDetail_DeletedKeyOwnerAccess(t *testing.T) {
+ ownerUID := int64(777)
+ otherUID := int64(2)
+
+ // 情况2:user_id=NULL,靠 deleted_key_owner_user_id 归因到 ownerUID
+ mk := func() *OpsErrorLogDetail {
+ return &OpsErrorLogDetail{
+ OpsErrorLog: OpsErrorLog{
+ ID: 55,
+ Phase: "auth",
+ Type: "api_error",
+ StatusCode: 401,
+ Message: "Invalid API key",
+ UserID: nil,
+ APIKeyName: "my-old-key",
+ APIKeyDeleted: true,
+ },
+ DeletedKeyOwnerUserID: &ownerUID,
+ }
+ }
+
+ // 原所有者(经 deleted_key 归因)→ 放行
+ svcOwner := &OpsService{opsRepo: &stubOpsRepoForUserErr{detailToReturn: mk()}}
+ got, err := svcOwner.GetUserErrorRequestDetail(context.Background(), ownerUID, 55)
+ if err != nil {
+ t.Fatalf("owner via deleted_key should be allowed, got err: %v", err)
+ }
+ if got == nil || got.ID != 55 {
+ t.Fatalf("expected detail ID=55, got %+v", got)
+ }
+ if !got.KeyDeleted || got.KeyName != "my-old-key" {
+ t.Fatalf("expected KeyDeleted=true KeyName=my-old-key, got %+v", got)
+ }
+
+ // 他人 → NotFound,不泄露存在性
+ svcOther := &OpsService{opsRepo: &stubOpsRepoForUserErr{detailToReturn: mk()}}
+ got2, err2 := svcOther.GetUserErrorRequestDetail(context.Background(), otherUID, 55)
+ if err2 == nil || got2 != nil {
+ t.Fatalf("non-owner should get (nil, NotFound), got detail=%+v err=%v", got2, err2)
+ }
+ if !infraerrors.IsNotFound(err2) {
+ t.Fatalf("expected NotFound, got %v", err2)
+ }
+}
diff --git a/backend/internal/service/ops_settings.go b/backend/internal/service/ops_settings.go
index 68c1d9dd..472f4e32 100644
--- a/backend/internal/service/ops_settings.go
+++ b/backend/internal/service/ops_settings.go
@@ -369,6 +369,7 @@ func defaultOpsAdvancedSettings() *OpsAdvancedSettings {
Aggregation: OpsAggregationSettings{
AggregationEnabled: false,
},
+ OpenAIAccountQuotaAutoPause: OpsOpenAIAccountQuotaAutoPauseSettings{},
IgnoreCountTokensErrors: true, // count_tokens 404 是预期行为,默认忽略
IgnoreContextCanceled: true, // Default to true - client disconnects are not errors
IgnoreNoAvailableAccounts: false, // Default to false - this is a real routing issue
@@ -384,6 +385,8 @@ func normalizeOpsAdvancedSettings(cfg *OpsAdvancedSettings) {
if cfg == nil {
return
}
+ cfg.OpenAIAccountQuotaAutoPause.DefaultThreshold5h = clampOpsQuotaAutoPauseThreshold(cfg.OpenAIAccountQuotaAutoPause.DefaultThreshold5h)
+ cfg.OpenAIAccountQuotaAutoPause.DefaultThreshold7d = clampOpsQuotaAutoPauseThreshold(cfg.OpenAIAccountQuotaAutoPause.DefaultThreshold7d)
cfg.DataRetention.CleanupSchedule = strings.TrimSpace(cfg.DataRetention.CleanupSchedule)
if cfg.DataRetention.CleanupSchedule == "" {
cfg.DataRetention.CleanupSchedule = opsCleanupDefaultSchedule
@@ -405,6 +408,16 @@ func normalizeOpsAdvancedSettings(cfg *OpsAdvancedSettings) {
}
}
+func clampOpsQuotaAutoPauseThreshold(value float64) float64 {
+ if value <= 0 {
+ return 0
+ }
+ if value > 1 {
+ return 1
+ }
+ return value
+}
+
func validateOpsAdvancedSettings(cfg *OpsAdvancedSettings) error {
if cfg == nil {
return errors.New("invalid config")
@@ -477,6 +490,12 @@ func (s *OpsService) UpdateOpsAdvancedSettings(ctx context.Context, cfg *OpsAdva
if err := s.settingRepo.Set(ctx, SettingKeyOpsAdvancedSettings, string(raw)); err != nil {
return nil, err
}
+ // Push the new quota auto-pause settings straight into the in-memory cache that
+ // the OpenAI scheduling hot path reads, so the next request observes the new value
+ // without waiting for the background refresher's TTL.
+ if s.quotaAutoPauseSink != nil {
+ s.quotaAutoPauseSink(cfg.OpenAIAccountQuotaAutoPause)
+ }
// notify cleanup service to reload schedule/enabled.
if s.cleanupReloader != nil {
diff --git a/backend/internal/service/ops_settings_advanced_test.go b/backend/internal/service/ops_settings_advanced_test.go
index 06cc545b..62803f94 100644
--- a/backend/internal/service/ops_settings_advanced_test.go
+++ b/backend/internal/service/ops_settings_advanced_test.go
@@ -4,6 +4,9 @@ import (
"context"
"encoding/json"
"testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
)
func TestGetOpsAdvancedSettings_DefaultHidesOpenAITokenStats(t *testing.T) {
@@ -95,3 +98,64 @@ func TestGetOpsAdvancedSettings_BackfillsNewDisplayFlagsFromDefaults(t *testing.
t.Fatalf("DisplayAlertEvents = false, want true default backfill")
}
}
+
+func TestGetOpenAIQuotaAutoPauseSettings_ReadsDefaultsFromOpsAdvancedSettings(t *testing.T) {
+ repo := newRuntimeSettingRepoStub()
+ repo.values[SettingKeyOpsAdvancedSettings] = `{"openai_account_quota_auto_pause":{"default_threshold_5h":0.95,"default_threshold_7d":0.9}}`
+ svc := NewSettingService(repo, &config.Config{})
+
+ // Warm the in-memory cache synchronously so the assertion below is deterministic.
+ // GetOpenAIQuotaAutoPauseSettings is non-blocking on the hot path (returns the
+ // cached value, refreshes asynchronously); for tests and startup, Warm is the
+ // synchronous entry point that guarantees a populated cache.
+ settings := svc.WarmOpenAIQuotaAutoPauseSettings(context.Background())
+ if settings.DefaultThreshold5h != 0.95 {
+ t.Fatalf("DefaultThreshold5h = %v, want 0.95", settings.DefaultThreshold5h)
+ }
+ if settings.DefaultThreshold7d != 0.9 {
+ t.Fatalf("DefaultThreshold7d = %v, want 0.9", settings.DefaultThreshold7d)
+ }
+
+ // Subsequent Get must hit the warm cache and return the same value without any DB
+ // access — that's the hot-path invariant.
+ cached := svc.GetOpenAIQuotaAutoPauseSettings(context.Background())
+ if cached.DefaultThreshold5h != 0.95 || cached.DefaultThreshold7d != 0.9 {
+ t.Fatalf("cached read = %+v, want {0.95, 0.9}", cached)
+ }
+}
+
+// Hot-path invariant: a Get with cold cache must return immediately (zero defaults)
+// rather than blocking on the DB. The async refresher will populate the cache for
+// subsequent calls.
+func TestGetOpenAIQuotaAutoPauseSettings_ColdCacheNonBlocking(t *testing.T) {
+ repo := newRuntimeSettingRepoStub()
+ repo.values[SettingKeyOpsAdvancedSettings] = `{"openai_account_quota_auto_pause":{"default_threshold_5h":0.7}}`
+ svc := NewSettingService(repo, &config.Config{})
+
+ start := time.Now()
+ settings := svc.GetOpenAIQuotaAutoPauseSettings(context.Background())
+ elapsed := time.Since(start)
+ if elapsed > 50*time.Millisecond {
+ t.Fatalf("cold-cache Get must be non-blocking, took %v", elapsed)
+ }
+ // Cold cache means we get zero defaults (the async refresh hasn't completed yet).
+ if settings.DefaultThreshold5h != 0 || settings.DefaultThreshold7d != 0 {
+ t.Fatalf("cold-cache Get = %+v, want zeroes", settings)
+ }
+}
+
+// Explicit cache write (e.g. from UpdateOpsAdvancedSettings) must be visible on the
+// very next read without any DB roundtrip.
+func TestSetOpenAIQuotaAutoPauseSettings_VisibleImmediately(t *testing.T) {
+ svc := NewSettingService(newRuntimeSettingRepoStub(), &config.Config{})
+
+ svc.SetOpenAIQuotaAutoPauseSettings(OpsOpenAIAccountQuotaAutoPauseSettings{
+ DefaultThreshold5h: 0.88,
+ DefaultThreshold7d: 0.77,
+ })
+
+ got := svc.GetOpenAIQuotaAutoPauseSettings(context.Background())
+ if got.DefaultThreshold5h != 0.88 || got.DefaultThreshold7d != 0.77 {
+ t.Fatalf("after Set, Get = %+v, want {0.88, 0.77}", got)
+ }
+}
diff --git a/backend/internal/service/ops_settings_models.go b/backend/internal/service/ops_settings_models.go
index fa18b05f..4d459e2b 100644
--- a/backend/internal/service/ops_settings_models.go
+++ b/backend/internal/service/ops_settings_models.go
@@ -92,17 +92,23 @@ type OpsAlertRuntimeSettings struct {
// OpsAdvancedSettings stores advanced ops configuration (data retention, aggregation).
type OpsAdvancedSettings struct {
- DataRetention OpsDataRetentionSettings `json:"data_retention"`
- Aggregation OpsAggregationSettings `json:"aggregation"`
- IgnoreCountTokensErrors bool `json:"ignore_count_tokens_errors"`
- IgnoreContextCanceled bool `json:"ignore_context_canceled"`
- IgnoreNoAvailableAccounts bool `json:"ignore_no_available_accounts"`
- IgnoreInvalidApiKeyErrors bool `json:"ignore_invalid_api_key_errors"`
- IgnoreInsufficientBalanceErrors bool `json:"ignore_insufficient_balance_errors"`
- DisplayOpenAITokenStats bool `json:"display_openai_token_stats"`
- DisplayAlertEvents bool `json:"display_alert_events"`
- AutoRefreshEnabled bool `json:"auto_refresh_enabled"`
- AutoRefreshIntervalSec int `json:"auto_refresh_interval_seconds"`
+ DataRetention OpsDataRetentionSettings `json:"data_retention"`
+ Aggregation OpsAggregationSettings `json:"aggregation"`
+ OpenAIAccountQuotaAutoPause OpsOpenAIAccountQuotaAutoPauseSettings `json:"openai_account_quota_auto_pause"`
+ IgnoreCountTokensErrors bool `json:"ignore_count_tokens_errors"`
+ IgnoreContextCanceled bool `json:"ignore_context_canceled"`
+ IgnoreNoAvailableAccounts bool `json:"ignore_no_available_accounts"`
+ IgnoreInvalidApiKeyErrors bool `json:"ignore_invalid_api_key_errors"`
+ IgnoreInsufficientBalanceErrors bool `json:"ignore_insufficient_balance_errors"`
+ DisplayOpenAITokenStats bool `json:"display_openai_token_stats"`
+ DisplayAlertEvents bool `json:"display_alert_events"`
+ AutoRefreshEnabled bool `json:"auto_refresh_enabled"`
+ AutoRefreshIntervalSec int `json:"auto_refresh_interval_seconds"`
+}
+
+type OpsOpenAIAccountQuotaAutoPauseSettings struct {
+ DefaultThreshold5h float64 `json:"default_threshold_5h"`
+ DefaultThreshold7d float64 `json:"default_threshold_7d"`
}
type OpsDataRetentionSettings struct {
diff --git a/backend/internal/service/ops_user_error.go b/backend/internal/service/ops_user_error.go
new file mode 100644
index 00000000..817c139a
--- /dev/null
+++ b/backend/internal/service/ops_user_error.go
@@ -0,0 +1,123 @@
+package service
+
+import "time"
+
+// UserErrorRequest 是面向终端用户的"错误请求"精简脱敏视图(白名单)。
+// 严禁包含 client_ip / user_agent / account / api_key_prefix / upstream_endpoint /
+// user_email 等敏感或内部字段。注:message(网关标准化错误描述)与 key_name
+// (用户自有 API Key 名称,KeysView 中本就可见)经产品决策对该用户开放;
+// error_body 仅在详情接口(GetUserErrorRequestDetail)按归属校验后返回。
+type UserErrorRequest struct {
+ ID int64 `json:"id"`
+ CreatedAt time.Time `json:"created_at"`
+ Model string `json:"model"`
+ InboundEndpoint string `json:"inbound_endpoint"`
+ StatusCode int `json:"status_code"`
+ Category string `json:"category"`
+ Platform string `json:"platform"`
+ Message string `json:"message"`
+ KeyName string `json:"key_name"`
+ KeyDeleted bool `json:"key_deleted"`
+}
+
+// UserErrorRequestList 是用户错误请求分页结果。
+type UserErrorRequestList struct {
+ Items []*UserErrorRequest `json:"items"`
+ Total int `json:"total"`
+ Page int `json:"page"`
+ PageSize int `json:"page_size"`
+}
+
+// MapUserErrorCategory 把后端 error_phase + error_type 映射为用户侧粗分类码。
+// 返回的是稳定的分类 code(前端做 i18n),不是展示文案。
+func MapUserErrorCategory(phase, errType string) string {
+ switch phase {
+ case "auth":
+ return "auth"
+ case "routing":
+ return "service_unavailable"
+ case "upstream", "network":
+ return "upstream"
+ case "internal":
+ return "internal"
+ case "request":
+ switch errType {
+ case "rate_limit_error":
+ return "rate_limit"
+ case "billing_error", "subscription_error":
+ return "quota"
+ case "invalid_request_error":
+ return "invalid_request"
+ }
+ }
+ return "other"
+}
+
+// CategoryToFilter 把用户侧分类码反向映射为后端过滤条件(plain ANY)。
+// 未知分类返回两个空切片(即不施加分类过滤)。
+// 注意:"other" 与未知分类都走 default 返回空切片——"other" 无对应的 phase/type 组合,无法精确反查,因此等价于不过滤。
+func CategoryToFilter(category string) (phases []string, errorTypes []string) {
+ switch category {
+ case "auth":
+ return []string{"auth"}, nil
+ case "service_unavailable":
+ return []string{"routing"}, nil
+ case "upstream":
+ return []string{"upstream", "network"}, nil
+ case "internal":
+ return []string{"internal"}, nil
+ case "rate_limit":
+ return nil, []string{"rate_limit_error"}
+ case "quota":
+ return nil, []string{"billing_error", "subscription_error"}
+ case "invalid_request":
+ return nil, []string{"invalid_request_error"}
+ default:
+ return nil, nil
+ }
+}
+
+// ToUserErrorRequest 把内部 OpsErrorLog 裁剪为用户安全视图。
+func ToUserErrorRequest(e *OpsErrorLog) *UserErrorRequest {
+ if e == nil {
+ return nil
+ }
+ model := e.RequestedModel
+ if model == "" {
+ model = e.Model
+ }
+ return &UserErrorRequest{
+ ID: e.ID,
+ CreatedAt: e.CreatedAt,
+ Model: model,
+ InboundEndpoint: e.InboundEndpoint,
+ StatusCode: e.StatusCode,
+ Category: MapUserErrorCategory(e.Phase, e.Type),
+ Platform: e.Platform,
+ Message: e.Message,
+ KeyName: e.APIKeyName,
+ KeyDeleted: e.APIKeyDeleted,
+ }
+}
+
+// UserErrorRequestDetail 是错误请求详情的脱敏视图(点击单行查看)。
+// 在 UserErrorRequest 基础上额外暴露 error_body(上游错误响应正文)与 upstream_status_code;
+// 仍严禁任何内部/敏感字段。
+type UserErrorRequestDetail struct {
+ UserErrorRequest
+ ErrorBody string `json:"error_body"`
+ UpstreamStatusCode *int `json:"upstream_status_code,omitempty"`
+}
+
+// ToUserErrorRequestDetail 把内部 OpsErrorLogDetail 裁剪为用户安全详情视图。
+func ToUserErrorRequestDetail(e *OpsErrorLogDetail) *UserErrorRequestDetail {
+ if e == nil {
+ return nil
+ }
+ base := ToUserErrorRequest(&e.OpsErrorLog)
+ return &UserErrorRequestDetail{
+ UserErrorRequest: *base,
+ ErrorBody: e.ErrorBody,
+ UpstreamStatusCode: e.UpstreamStatusCode,
+ }
+}
diff --git a/backend/internal/service/ops_user_error_test.go b/backend/internal/service/ops_user_error_test.go
new file mode 100644
index 00000000..31b0c269
--- /dev/null
+++ b/backend/internal/service/ops_user_error_test.go
@@ -0,0 +1,167 @@
+package service
+
+import (
+ "encoding/json"
+ "strings"
+ "testing"
+ "time"
+)
+
+func TestMapUserErrorCategory(t *testing.T) {
+ cases := []struct {
+ phase, etype, want string
+ }{
+ {"auth", "authentication_error", "auth"},
+ {"request", "rate_limit_error", "rate_limit"},
+ {"request", "billing_error", "quota"},
+ {"request", "subscription_error", "quota"},
+ {"request", "invalid_request_error", "invalid_request"},
+ {"routing", "api_error", "service_unavailable"},
+ {"upstream", "upstream_error", "upstream"},
+ {"network", "api_error", "upstream"},
+ {"internal", "api_error", "internal"},
+ {"weird", "weird", "other"},
+ }
+ for _, c := range cases {
+ if got := MapUserErrorCategory(c.phase, c.etype); got != c.want {
+ t.Errorf("MapUserErrorCategory(%q,%q)=%q want %q", c.phase, c.etype, got, c.want)
+ }
+ }
+}
+
+func TestCategoryToFilter(t *testing.T) {
+ phases, types := CategoryToFilter("rate_limit")
+ if len(types) != 1 || types[0] != "rate_limit_error" || len(phases) != 0 {
+ t.Fatalf("rate_limit => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("auth")
+ if len(phases) != 1 || phases[0] != "auth" || len(types) != 0 {
+ t.Fatalf("auth => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("service_unavailable")
+ if len(phases) != 1 || phases[0] != "routing" || len(types) != 0 {
+ t.Fatalf("service_unavailable => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("upstream")
+ if len(phases) != 2 || phases[0] != "upstream" || phases[1] != "network" || len(types) != 0 {
+ t.Fatalf("upstream => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("internal")
+ if len(phases) != 1 || phases[0] != "internal" || len(types) != 0 {
+ t.Fatalf("internal => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("quota")
+ if len(types) != 2 || types[0] != "billing_error" || types[1] != "subscription_error" || len(phases) != 0 {
+ t.Fatalf("quota => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("invalid_request")
+ if len(types) != 1 || types[0] != "invalid_request_error" || len(phases) != 0 {
+ t.Fatalf("invalid_request => phases=%v types=%v", phases, types)
+ }
+ phases, types = CategoryToFilter("other")
+ if len(phases) != 0 || len(types) != 0 {
+ t.Fatalf("other => phases=%v types=%v", phases, types)
+ }
+}
+
+func TestToUserErrorRequest_RedactsSensitiveFields(t *testing.T) {
+ src := &OpsErrorLog{
+ ID: 123,
+ CreatedAt: time.Unix(0, 0).UTC(),
+ Model: "m",
+ RequestedModel: "rm",
+ InboundEndpoint: "/v1/chat/completions",
+ StatusCode: 429,
+ Platform: "openai",
+ Phase: "request",
+ Type: "rate_limit_error",
+ Message: "rate limit exceeded",
+ APIKeyName: "my-key",
+ APIKeyDeleted: true,
+ }
+ out := ToUserErrorRequest(src)
+ if out.ID != 123 {
+ t.Errorf("want ID=123, got %d", out.ID)
+ }
+ if out.Model != "rm" {
+ t.Errorf("want requested_model preferred, got %q", out.Model)
+ }
+ if out.Category != "rate_limit" {
+ t.Errorf("category=%q", out.Category)
+ }
+ if out.StatusCode != 429 || out.InboundEndpoint != "/v1/chat/completions" || out.Platform != "openai" {
+ t.Errorf("basic fields wrong: %+v", out)
+ }
+ if out.Message != "rate limit exceeded" {
+ t.Errorf("want message=%q, got %q", "rate limit exceeded", out.Message)
+ }
+ if out.KeyName != "my-key" {
+ t.Errorf("want key_name=my-key, got %q", out.KeyName)
+ }
+ if !out.KeyDeleted {
+ t.Error("want key_deleted=true")
+ }
+}
+
+func TestToUserErrorRequestDetail_WhitelistAndRedacts(t *testing.T) {
+ uid := int64(42)
+ upstreamStatus := 503
+ src := &OpsErrorLogDetail{
+ OpsErrorLog: OpsErrorLog{
+ ID: 999,
+ CreatedAt: time.Unix(1000, 0).UTC(),
+ Model: "gpt-4",
+ RequestedModel: "gpt-4-turbo",
+ InboundEndpoint: "/v1/chat/completions",
+ StatusCode: 502,
+ Platform: "openai",
+ Phase: "upstream",
+ Type: "api_error",
+ Message: "upstream error",
+ UserID: &uid,
+ UserEmail: "secret@example.com",
+ ClientIP: func() *string { s := "1.2.3.4"; return &s }(),
+ UpstreamEndpoint: "https://api.openai.com/v1/chat/completions",
+ },
+ ErrorBody: `{"error":{"message":"upstream failed","type":"server_error"}}`,
+ UserAgent: "Mozilla/5.0 secret-agent",
+ UpstreamStatusCode: &upstreamStatus,
+ }
+
+ out := ToUserErrorRequestDetail(src)
+ if out == nil {
+ t.Fatal("expected non-nil detail")
+ }
+
+ // 基础字段正确映射
+ if out.ID != 999 {
+ t.Errorf("want ID=999, got %d", out.ID)
+ }
+ if out.Message != "upstream error" {
+ t.Errorf("want message=%q, got %q", "upstream error", out.Message)
+ }
+ if out.ErrorBody != src.ErrorBody {
+ t.Errorf("ErrorBody mismatch")
+ }
+ if out.UpstreamStatusCode == nil || *out.UpstreamStatusCode != 503 {
+ t.Errorf("UpstreamStatusCode mismatch")
+ }
+
+ // 序列化后不含敏感字段
+ b, err := json.Marshal(out)
+ if err != nil {
+ t.Fatalf("json.Marshal failed: %v", err)
+ }
+ raw := string(b)
+ for _, forbidden := range []string{"user_email", "client_ip", "upstream_endpoint", "user_agent"} {
+ if strings.Contains(raw, forbidden) {
+ t.Errorf("sensitive field %q leaked in JSON output: %s", forbidden, raw)
+ }
+ }
+}
+
+func TestToUserErrorRequestDetail_Nil(t *testing.T) {
+ if out := ToUserErrorRequestDetail(nil); out != nil {
+ t.Errorf("expected nil for nil input, got %+v", out)
+ }
+}
diff --git a/backend/internal/service/pricing_service_test.go b/backend/internal/service/pricing_service_test.go
index cc8b120a..f4252f95 100644
--- a/backend/internal/service/pricing_service_test.go
+++ b/backend/internal/service/pricing_service_test.go
@@ -124,9 +124,9 @@ func TestDefaultPricingIncludesCodexAutoReview(t *testing.T) {
got := svc.GetModelPricing("codex-auto-review")
require.NotNil(t, got)
- require.InDelta(t, 2.5e-6, got.InputCostPerToken, 1e-12)
- require.InDelta(t, 1.5e-5, got.OutputCostPerToken, 1e-12)
- require.InDelta(t, 2.5e-7, got.CacheReadInputTokenCost, 1e-12)
+ require.InDelta(t, 5e-6, got.InputCostPerToken, 1e-12)
+ require.InDelta(t, 3e-5, got.OutputCostPerToken, 1e-12)
+ require.InDelta(t, 5e-7, got.CacheReadInputTokenCost, 1e-12)
}
func TestGetModelPricing_Gpt54MiniUsesDedicatedStaticFallbackWhenRemoteMissing(t *testing.T) {
diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go
index c3b160e7..ecbd86d1 100644
--- a/backend/internal/service/ratelimit_service.go
+++ b/backend/internal/service/ratelimit_service.go
@@ -153,7 +153,7 @@ func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Accoun
// HandleUpstreamError 处理上游错误响应,标记账号状态
// 返回是否应该停止该账号的调度
-func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte) (shouldDisable bool) {
+func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte, requestedModel ...string) (shouldDisable bool) {
customErrorCodesEnabled := account.IsCustomErrorCodesEnabled()
// 池模式默认不标记本地账号状态;仅当用户显式配置自定义错误码时按本地策略处理。
@@ -169,6 +169,10 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc
return false
}
+ if len(requestedModel) > 0 && s.HandleUpstreamModelNotFound(ctx, account, requestedModel[0], statusCode, responseBody) {
+ return true
+ }
+
// 先尝试临时不可调度规则(401除外)
// 如果匹配成功,直接返回,不执行后续禁用逻辑
if statusCode != 401 {
@@ -244,17 +248,15 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc
shouldDisable = true
break
}
- // 2. 设置 expires_at 为当前时间,强制下次请求刷新 token
- if account.Credentials == nil {
- account.Credentials = make(map[string]any)
- }
- account.Credentials["expires_at"] = time.Now().Format(time.RFC3339)
- if err := persistAccountCredentials(ctx, s.accountRepo, account, account.Credentials); err != nil {
- slog.Warn("oauth_401_force_refresh_update_failed", "account_id", account.ID, "error", err)
- } else {
- slog.Info("oauth_401_force_refresh_set", "account_id", account.ID, "platform", account.Platform)
- }
- // 3. 临时不可调度,替代 SetError(保持 status=active 让刷新服务能拾取)
+ // 2. 临时不可调度,替代 SetError(保持 status=active 让刷新服务能拾取)
+ // 注意:此处不再写回 account.Credentials/expires_at。
+ // 原实现使用请求开始时的 account 快照整列覆盖 credentials JSONB(见
+ // persistAccountCredentials → accountRepository.UpdateCredentials → SetCredentials),
+ // 在另一个 worker 刚刷新完 refresh_token 的窄窗口内会把新 refresh_token 回滚为旧值,
+ // 导致下一周期用旧 refresh_token 调上游拿到 invalid_grant 后,
+ // tryRecoverFromRefreshRace 重读 DB 发现 currentRT == usedRT 也救不回来,账号被错误 disable。
+ // 这里仅依赖 InvalidateToken + SetTempUnschedulable 让账号在冷却期内不被调度,
+ // 冷却结束后由 token_provider 的 NeedsRefresh / token_refresh_service 走带分布式锁的正路刷新。
msg := "Authentication failed (401): invalid or expired credentials"
if upstreamMsg != "" {
msg = "OAuth 401: " + upstreamMsg
@@ -1616,9 +1618,51 @@ func (s *RateLimitService) HandleTempUnschedulable(ctx context.Context, account
return s.tryTempUnschedulable(ctx, account, statusCode, responseBody)
}
+const upstreamModelNotFoundCooldown = 30 * time.Minute
+const upstreamModelNotFoundReason = "upstream_404_model_not_found"
const tempUnschedBodyMaxBytes = 64 << 10
const tempUnschedMessageMaxBytes = 2048
+func (s *RateLimitService) HandleUpstreamModelNotFound(ctx context.Context, account *Account, requestedModel string, statusCode int, responseBody []byte) bool {
+ if s == nil || account == nil || s.accountRepo == nil {
+ return false
+ }
+ if !account.ShouldHandleErrorCode(statusCode) {
+ return false
+ }
+ if !isUpstreamModelNotFoundError(statusCode, responseBody) {
+ return false
+ }
+ modelKey := modelRateLimitKeyForUpstreamModelNotFound(ctx, account, requestedModel)
+ if modelKey == "" {
+ return false
+ }
+ resetAt := time.Now().Add(upstreamModelNotFoundCooldown)
+ if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, modelKey, resetAt, upstreamModelNotFoundReason); err != nil {
+ slog.Warn("upstream_model_not_found_set_model_rate_limit_failed", "account_id", account.ID, "model", modelKey, "error", err)
+ return true
+ }
+ slog.Info("upstream_model_not_found_model_rate_limited", "account_id", account.ID, "model", modelKey, "reset_at", resetAt)
+ return true
+}
+
+func modelRateLimitKeyForUpstreamModelNotFound(ctx context.Context, account *Account, requestedModel string) string {
+ modelKey := strings.TrimSpace(requestedModel)
+ if account == nil || modelKey == "" {
+ return modelKey
+ }
+ if account.Platform == PlatformAntigravity {
+ if resolved := strings.TrimSpace(resolveFinalAntigravityModelKey(ctx, account, modelKey)); resolved != "" {
+ return resolved
+ }
+ return modelKey
+ }
+ if mapped := strings.TrimSpace(account.GetMappedModel(modelKey)); mapped != "" {
+ return mapped
+ }
+ return modelKey
+}
+
func (s *RateLimitService) tryTempUnschedulable(ctx context.Context, account *Account, statusCode int, responseBody []byte) bool {
if account == nil {
return false
diff --git a/backend/internal/service/ratelimit_service_401_test.go b/backend/internal/service/ratelimit_service_401_test.go
index a964775e..873aaf33 100644
--- a/backend/internal/service/ratelimit_service_401_test.go
+++ b/backend/internal/service/ratelimit_service_401_test.go
@@ -129,7 +129,10 @@ func TestRateLimitService_HandleUpstreamError_OAuth401SetsTempUnschedulable(t *t
}
// TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError
-// OpenAI OAuth 401 缓存失效出错时仍走 temp_unschedulable
+// OpenAI OAuth 401 缓存失效出错时仍走 temp_unschedulable。
+// 注意:401 handler 不再回写 credentials(避免请求开始时的快照整列覆盖 DB
+// 把另一个 worker 刚刷新出来的新 refresh_token 回滚为旧值),
+// 因此 updateCredentialsCalls 应当为 0。
func TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError(t *testing.T) {
repo := &rateLimitAccountRepoStub{}
invalidator := &tokenCacheInvalidatorRecorder{err: errors.New("boom")}
@@ -149,7 +152,7 @@ func TestRateLimitService_HandleUpstreamError_OAuth401InvalidatorError(t *testin
require.True(t, shouldDisable)
require.Equal(t, 0, repo.setErrorCalls)
require.Equal(t, 1, repo.tempCalls)
- require.Equal(t, 1, repo.updateCredentialsCalls)
+ require.Equal(t, 0, repo.updateCredentialsCalls)
require.Len(t, invalidator.accounts, 1)
}
@@ -171,7 +174,12 @@ func TestRateLimitService_HandleUpstreamError_NonOAuth401(t *testing.T) {
require.Empty(t, invalidator.accounts)
}
-func TestRateLimitService_HandleUpstreamError_OAuth401UsesCredentialsUpdater(t *testing.T) {
+// TestRateLimitService_HandleUpstreamError_OAuth401DoesNotOverwriteCredentials
+// 回归测试:确保 401 handler 不再使用请求开始时的 account 快照写回 credentials。
+// 原实现会通过 persistAccountCredentials → UpdateCredentials → SetCredentials
+// 整列覆盖 credentials JSONB,在另一个 worker 刚刷新完 refresh_token 的窄窗口内
+// 会把新 refresh_token 回滚为快照中的旧值,导致下一周期拿 invalid_grant 被错误 disable。
+func TestRateLimitService_HandleUpstreamError_OAuth401DoesNotOverwriteCredentials(t *testing.T) {
repo := &rateLimitAccountRepoStub{}
service := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
account := &Account{
@@ -187,8 +195,9 @@ func TestRateLimitService_HandleUpstreamError_OAuth401UsesCredentialsUpdater(t *
shouldDisable := service.HandleUpstreamError(context.Background(), account, 401, http.Header{}, []byte("unauthorized"))
require.True(t, shouldDisable)
- require.Equal(t, 1, repo.updateCredentialsCalls)
- require.NotEmpty(t, repo.lastCredentials["expires_at"])
+ require.Equal(t, 0, repo.updateCredentialsCalls, "401 handler must not write credentials back from the request-start snapshot")
+ require.Equal(t, 1, repo.tempCalls, "401 handler should still set temp-unschedulable cooldown")
+ require.Nil(t, repo.lastCredentials, "no credentials should have been persisted")
}
// 缺少 refresh_token 的 OAuth 账号 401 应直接 SetError 永久禁用,
diff --git a/backend/internal/service/ratelimit_service_model_not_found_test.go b/backend/internal/service/ratelimit_service_model_not_found_test.go
new file mode 100644
index 00000000..dfd18c5f
--- /dev/null
+++ b/backend/internal/service/ratelimit_service_model_not_found_test.go
@@ -0,0 +1,127 @@
+//go:build unit
+
+package service
+
+import (
+ "context"
+ "errors"
+ "net/http"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/require"
+)
+
+type modelNotFoundRateLimitCall struct {
+ accountID int64
+ scope string
+ resetAt time.Time
+ reason string
+}
+
+type modelNotFoundAccountRepoStub struct {
+ mockAccountRepoForGemini
+ tempCalls int
+ modelRateLimitCalls []modelNotFoundRateLimitCall
+ modelRateLimitErr error
+}
+
+func (r *modelNotFoundAccountRepoStub) SetTempUnschedulable(ctx context.Context, id int64, until time.Time, reason string) error {
+ r.tempCalls++
+ return nil
+}
+
+func (r *modelNotFoundAccountRepoStub) SetModelRateLimit(ctx context.Context, id int64, scope string, resetAt time.Time, reason ...string) error {
+ call := modelNotFoundRateLimitCall{
+ accountID: id,
+ scope: scope,
+ resetAt: resetAt,
+ }
+ if len(reason) > 0 {
+ call.reason = reason[0]
+ }
+ r.modelRateLimitCalls = append(r.modelRateLimitCalls, call)
+ return r.modelRateLimitErr
+}
+
+func TestRateLimitService_HandleUpstreamError_ModelNotFoundUsesModelRateLimit(t *testing.T) {
+ repo := &modelNotFoundAccountRepoStub{}
+ svc := &RateLimitService{accountRepo: repo}
+ account := openAIModelNotFoundTempAccount()
+
+ handled := svc.HandleUpstreamError(
+ context.Background(),
+ account,
+ http.StatusNotFound,
+ http.Header{},
+ []byte(`{"error":{"code":"model_not_found","message":"model not found"}}`),
+ "gpt-5.4",
+ )
+
+ require.True(t, handled)
+ require.Zero(t, repo.tempCalls)
+ require.Len(t, repo.modelRateLimitCalls, 1)
+ call := repo.modelRateLimitCalls[0]
+ require.Equal(t, account.ID, call.accountID)
+ require.Equal(t, "gpt-5.4", call.scope)
+ require.Equal(t, upstreamModelNotFoundReason, call.reason)
+ require.WithinDuration(t, time.Now().Add(upstreamModelNotFoundCooldown), call.resetAt, 5*time.Second)
+}
+
+func TestRateLimitService_HandleUpstreamError_ModelNotFoundWriteFailureDoesNotTempUnschedule(t *testing.T) {
+ repo := &modelNotFoundAccountRepoStub{modelRateLimitErr: errors.New("write failed")}
+ svc := &RateLimitService{accountRepo: repo}
+ account := openAIModelNotFoundTempAccount()
+
+ handled := svc.HandleUpstreamError(
+ context.Background(),
+ account,
+ http.StatusNotFound,
+ http.Header{},
+ []byte(`{"error":{"code":"model_not_found","message":"model not found"}}`),
+ "gpt-5.4",
+ )
+
+ require.True(t, handled)
+ require.Zero(t, repo.tempCalls)
+ require.Len(t, repo.modelRateLimitCalls, 1)
+}
+
+func TestRateLimitService_HandleUpstreamError_Bare404KeepsTempUnschedulablePath(t *testing.T) {
+ repo := &modelNotFoundAccountRepoStub{}
+ svc := &RateLimitService{accountRepo: repo}
+ account := openAIModelNotFoundTempAccount()
+
+ handled := svc.HandleUpstreamError(
+ context.Background(),
+ account,
+ http.StatusNotFound,
+ http.Header{},
+ []byte(`{"error":{"message":"endpoint not found"}}`),
+ "gpt-5.4",
+ )
+
+ require.True(t, handled)
+ require.Equal(t, 1, repo.tempCalls)
+ require.Empty(t, repo.modelRateLimitCalls)
+}
+
+func openAIModelNotFoundTempAccount() *Account {
+ return &Account{
+ ID: 101,
+ Platform: PlatformOpenAI,
+ Type: AccountTypeAPIKey,
+ Status: StatusActive,
+ Schedulable: true,
+ Credentials: map[string]any{
+ "temp_unschedulable_enabled": true,
+ "temp_unschedulable_rules": []any{
+ map[string]any{
+ "error_code": float64(http.StatusNotFound),
+ "keywords": []any{"not found"},
+ "duration_minutes": float64(10),
+ },
+ },
+ },
+ }
+}
diff --git a/backend/internal/service/ratelimit_session_window_test.go b/backend/internal/service/ratelimit_session_window_test.go
index 7796a85e..be6cb309 100644
--- a/backend/internal/service/ratelimit_session_window_test.go
+++ b/backend/internal/service/ratelimit_session_window_test.go
@@ -137,7 +137,7 @@ func (m *sessionWindowMockRepo) ListSchedulableUngroupedByPlatforms(context.Cont
func (m *sessionWindowMockRepo) SetRateLimited(context.Context, int64, time.Time) error {
panic("unexpected")
}
-func (m *sessionWindowMockRepo) SetModelRateLimit(context.Context, int64, string, time.Time) error {
+func (m *sessionWindowMockRepo) SetModelRateLimit(context.Context, int64, string, time.Time, ...string) error {
panic("unexpected")
}
func (m *sessionWindowMockRepo) SetOverloaded(context.Context, int64, time.Time) error {
diff --git a/backend/internal/service/scheduler_snapshot_hydration_test.go b/backend/internal/service/scheduler_snapshot_hydration_test.go
index 778cab23..0a1d0a0a 100644
--- a/backend/internal/service/scheduler_snapshot_hydration_test.go
+++ b/backend/internal/service/scheduler_snapshot_hydration_test.go
@@ -6,6 +6,8 @@ import (
"context"
"testing"
"time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
)
type snapshotHydrationCache struct {
@@ -186,3 +188,89 @@ func TestGatewaySelectAccountWithLoadAwareness_HydratesSelectedAccountFromSchedu
t.Fatalf("expected hydrated api key, got %q", got)
}
}
+
+func TestGatewaySelectAccountWithLoadAwareness_SkipsAntigravityGeminiFamilyRateLimitedSnapshot(t *testing.T) {
+ resetAt := time.Now().Add(10 * time.Minute).Format(time.RFC3339)
+ cache := &snapshotHydrationCache{
+ snapshot: []*Account{
+ {
+ ID: 1,
+ Platform: PlatformAntigravity,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 1,
+ AccountGroups: []AccountGroup{
+ {AccountID: 1, GroupID: 22},
+ },
+ GroupIDs: []int64{22},
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ modelRateLimitsKey: map[string]any{
+ antigravityGeminiModelRateLimitKey: map[string]any{
+ "rate_limit_reset_at": resetAt,
+ },
+ },
+ },
+ },
+ {
+ ID: 2,
+ Platform: PlatformAntigravity,
+ Type: AccountTypeOAuth,
+ Status: StatusActive,
+ Schedulable: true,
+ Concurrency: 1,
+ Priority: 2,
+ AccountGroups: []AccountGroup{
+ {AccountID: 2, GroupID: 22},
+ },
+ GroupIDs: []int64{22},
+ Extra: map[string]any{
+ "mixed_scheduling": true,
+ },
+ },
+ },
+ accounts: map[int64]*Account{
+ 1: {ID: 1, Platform: PlatformAntigravity, Type: AccountTypeOAuth},
+ 2: {ID: 2, Platform: PlatformAntigravity, Type: AccountTypeOAuth},
+ },
+ }
+ groupID := int64(22)
+ svc := &GatewayService{
+ schedulerSnapshot: NewSchedulerSnapshotService(cache, nil, nil, nil, nil),
+ groupRepo: &mockGroupRepoForGateway{
+ groups: map[int64]*Group{
+ groupID: {
+ ID: groupID,
+ Platform: PlatformGemini,
+ Status: StatusActive,
+ Hydrated: true,
+ },
+ },
+ },
+ concurrencyService: NewConcurrencyService(&mockConcurrencyCache{}),
+ cfg: &config.Config{
+ Gateway: config.GatewayConfig{
+ Scheduling: config.GatewaySchedulingConfig{
+ LoadBatchEnabled: true,
+ StickySessionMaxWaiting: 3,
+ StickySessionWaitTimeout: time.Second,
+ FallbackWaitTimeout: time.Second,
+ FallbackMaxWaiting: 10,
+ },
+ },
+ },
+ }
+
+ result, err := svc.SelectAccountWithLoadAwareness(context.Background(), &groupID, "", "gemini-3-flash-preview", nil, "", 0)
+ if err != nil {
+ t.Fatalf("SelectAccountWithLoadAwareness error: %v", err)
+ }
+ if result == nil || result.Account == nil {
+ t.Fatalf("expected selected account")
+ }
+ if result.Account.ID != 2 {
+ t.Fatalf("expected scheduler to skip Gemini-family limited antigravity account 1, got %d", result.Account.ID)
+ }
+}
diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go
index f69499b6..4daf7f5f 100644
--- a/backend/internal/service/setting_service.go
+++ b/backend/internal/service/setting_service.go
@@ -141,10 +141,32 @@ type cachedOpenAICodexUserAgent struct {
expiresAt int64 // unix nano
}
+type cachedOpenAIQuotaAutoPauseSettings struct {
+ settings OpsOpenAIAccountQuotaAutoPauseSettings
+ expiresAt int64
+}
+
const openAICodexUserAgentCacheTTL = 60 * time.Second
const openAICodexUserAgentErrorTTL = 5 * time.Second
const openAICodexUserAgentDBTimeout = 5 * time.Second
+// cachedOpenAIAllowCodexPlugin Codex 插件放行开关缓存(进程内缓存,60s TTL)。
+// IsOpenAIAllowClaudeCodeCodexPluginEnabled 在每个 codex_cli_only 账号的网关请求热路径上被调用,避免每次访问 DB。
+type cachedOpenAIAllowCodexPlugin struct {
+ value bool
+ expiresAt int64 // unix nano
+}
+
+const openAIAllowCodexPluginCacheTTL = 60 * time.Second
+const openAIAllowCodexPluginErrorTTL = 5 * time.Second
+const openAIAllowCodexPluginDBTimeout = 5 * time.Second
+
+const openAIQuotaAutoPauseSettingsCacheTTL = 60 * time.Second
+const openAIQuotaAutoPauseSettingsErrorTTL = 5 * time.Second
+const openAIQuotaAutoPauseSettingsDBTimeout = 5 * time.Second
+
+const openAIQuotaAutoPauseSettingsRefreshKey = "openai_quota_auto_pause_settings"
+
// DefaultSubscriptionGroupReader validates group references used by default subscriptions.
type DefaultSubscriptionGroupReader interface {
GetByID(ctx context.Context, id int64) (*Group, error)
@@ -156,17 +178,28 @@ type WebSearchManagerBuilder func(cfg *WebSearchEmulationConfig, proxyURLs map[i
// SettingService 系统设置服务
type SettingService struct {
- settingRepo SettingRepository
- defaultSubGroupReader DefaultSubscriptionGroupReader
- proxyRepo ProxyRepository // for resolving websearch provider proxy URLs
- cfg *config.Config
- onUpdate func() // Callback when settings are updated (for cache invalidation)
- version string // Application version
- webSearchManagerBuilder WebSearchManagerBuilder
- antigravityUAVersionCache atomic.Value // *cachedAntigravityUserAgentVersion
- antigravityUAVersionSF singleflight.Group
- openAICodexUACache atomic.Value // *cachedOpenAICodexUserAgent
- openAICodexUASF singleflight.Group
+ settingRepo SettingRepository
+ defaultSubGroupReader DefaultSubscriptionGroupReader
+ proxyRepo ProxyRepository // for resolving websearch provider proxy URLs
+ cfg *config.Config
+ onUpdate func() // Callback when settings are updated (for cache invalidation)
+ version string // Application version
+ webSearchManagerBuilder WebSearchManagerBuilder
+ antigravityUAVersionCache atomic.Value // *cachedAntigravityUserAgentVersion
+ antigravityUAVersionSF singleflight.Group
+ openAICodexUACache atomic.Value // *cachedOpenAICodexUserAgent
+ openAICodexUASF singleflight.Group
+ openAIAllowCodexPluginCache atomic.Value // *cachedOpenAIAllowCodexPlugin
+ openAIAllowCodexPluginSF singleflight.Group
+
+ // openAIQuotaAutoPauseSettingsCache holds the most recently observed quota auto-pause
+ // settings. GetOpenAIQuotaAutoPauseSettings reads this atomic.Value on the request hot
+ // path without ever blocking on the DB; when the cached entry expires, a background
+ // goroutine refreshes it via openAIQuotaAutoPauseSettingsSF (stale-while-revalidate).
+ // This per-service field also gives tests natural isolation — each SettingService
+ // instance owns its own cache, no shared package-level state.
+ openAIQuotaAutoPauseSettingsCache atomic.Value // *cachedOpenAIQuotaAutoPauseSettings
+ openAIQuotaAutoPauseSettingsSF singleflight.Group
}
// DefaultPlatformQuotaSetting 单 platform 三档限额(nil = 沿用上层;0 = 显式禁用;>0 = 上限)
@@ -741,6 +774,7 @@ func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings
SettingKeyAvailableChannelsEnabled,
SettingKeyAffiliateEnabled,
SettingKeyRiskControlEnabled,
+ SettingKeyAllowUserViewErrorRequests,
}
settings, err := s.settingRepo.GetMultiple(ctx, keys)
@@ -854,6 +888,8 @@ func (s *SettingService) GetPublicSettings(ctx context.Context) (*PublicSettings
AffiliateEnabled: settings[SettingKeyAffiliateEnabled] == "true",
RiskControlEnabled: settings[SettingKeyRiskControlEnabled] == "true",
+
+ AllowUserViewErrorRequests: settings[SettingKeyAllowUserViewErrorRequests] == "true",
}, nil
}
@@ -931,6 +967,17 @@ func (s *SettingService) GetAvailableChannelsRuntime(ctx context.Context) Availa
}
}
+// IsUserErrorViewAllowed reads the user-facing error-requests visibility switch
+// directly from the settings store. Fail-closed: on error returns false (opt-in default).
+func (s *SettingService) IsUserErrorViewAllowed(ctx context.Context) bool {
+ vals, err := s.settingRepo.GetMultiple(ctx, []string{SettingKeyAllowUserViewErrorRequests})
+ if err != nil {
+ slog.Warn("failed to get allow_user_view_error_requests setting, defaulting to false", "error", err)
+ return false
+ }
+ return vals[SettingKeyAllowUserViewErrorRequests] == "true"
+}
+
// GetAntigravityUserAgentVersion 返回 Antigravity 上游请求使用的版本号。
// 后台设置优先;为空、缺失或非法时回退到 ANTIGRAVITY_USER_AGENT_VERSION / 内置默认值。
func (s *SettingService) GetAntigravityUserAgentVersion(ctx context.Context) string {
@@ -1029,6 +1076,54 @@ func (s *SettingService) GetOpenAICodexUserAgent(ctx context.Context) string {
return fallback
}
+// IsOpenAIAllowClaudeCodeCodexPluginEnabled 全局开关:是否额外放行 Claude Code 的 Codex 插件(默认关闭)。
+// 仅在调用方已确认账号 codex_cli_only 开启时读取,避免对非受限账号产生无谓查询。
+// 使用进程内 atomic.Value 缓存(60s TTL),避免在每个网关请求热路径上访问 DB。
+func (s *SettingService) IsOpenAIAllowClaudeCodeCodexPluginEnabled(ctx context.Context) bool {
+ if cached, ok := s.openAIAllowCodexPluginCache.Load().(*cachedOpenAIAllowCodexPlugin); ok && cached != nil {
+ if time.Now().UnixNano() < cached.expiresAt {
+ return cached.value
+ }
+ }
+ result, _, _ := s.openAIAllowCodexPluginSF.Do("openai_allow_codex_plugin_enabled", func() (any, error) {
+ if cached, ok := s.openAIAllowCodexPluginCache.Load().(*cachedOpenAIAllowCodexPlugin); ok && cached != nil {
+ if time.Now().UnixNano() < cached.expiresAt {
+ return cached.value, nil
+ }
+ }
+ dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), openAIAllowCodexPluginDBTimeout)
+ defer cancel()
+ value, err := s.settingRepo.GetValue(dbCtx, SettingKeyOpenAIAllowClaudeCodeCodexPlugin)
+ if err != nil {
+ if errors.Is(err, ErrSettingNotFound) {
+ // 设置不存在 → 默认关闭,正常 TTL 缓存
+ s.openAIAllowCodexPluginCache.Store(&cachedOpenAIAllowCodexPlugin{
+ value: false,
+ expiresAt: time.Now().Add(openAIAllowCodexPluginCacheTTL).UnixNano(),
+ })
+ return false, nil
+ }
+ slog.Warn("failed to get openai_allow_claude_code_codex_plugin setting", "error", err)
+ // DB 错误 → 安全默认关闭,短 TTL 快速重试
+ s.openAIAllowCodexPluginCache.Store(&cachedOpenAIAllowCodexPlugin{
+ value: false,
+ expiresAt: time.Now().Add(openAIAllowCodexPluginErrorTTL).UnixNano(),
+ })
+ return false, nil
+ }
+ enabled := value == "true"
+ s.openAIAllowCodexPluginCache.Store(&cachedOpenAIAllowCodexPlugin{
+ value: enabled,
+ expiresAt: time.Now().Add(openAIAllowCodexPluginCacheTTL).UnixNano(),
+ })
+ return enabled, nil
+ })
+ if val, ok := result.(bool); ok {
+ return val
+ }
+ return false
+}
+
// SetOnUpdateCallback sets a callback function to be called when settings are updated
// This is used for cache invalidation (e.g., HTML cache in frontend server)
func (s *SettingService) SetOnUpdateCallback(callback func()) {
@@ -1108,6 +1203,7 @@ type PublicSettingsInjectionPayload struct {
AvailableChannelsEnabled bool `json:"available_channels_enabled"`
AffiliateEnabled bool `json:"affiliate_enabled"`
RiskControlEnabled bool `json:"risk_control_enabled"`
+ AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"`
}
// GetPublicSettingsForInjection returns public settings in a format suitable for HTML injection.
@@ -1170,6 +1266,7 @@ func (s *SettingService) GetPublicSettingsForInjection(ctx context.Context) (any
AvailableChannelsEnabled: settings.AvailableChannelsEnabled,
AffiliateEnabled: settings.AffiliateEnabled,
RiskControlEnabled: settings.RiskControlEnabled,
+ AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests,
}, nil
}
@@ -1830,6 +1927,7 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting
updates[SettingKeyRewriteMessageCacheControl] = strconv.FormatBool(settings.RewriteMessageCacheControl)
updates[SettingKeyAntigravityUserAgentVersion] = antigravity.NormalizeUserAgentVersion(settings.AntigravityUserAgentVersion)
updates[SettingKeyOpenAICodexUserAgent] = strings.TrimSpace(settings.OpenAICodexUserAgent)
+ updates[SettingKeyOpenAIAllowClaudeCodeCodexPlugin] = strconv.FormatBool(settings.OpenAIAllowClaudeCodeCodexPlugin)
updates[SettingPaymentVisibleMethodAlipaySource] = settings.PaymentVisibleMethodAlipaySource
updates[SettingPaymentVisibleMethodWxpaySource] = settings.PaymentVisibleMethodWxpaySource
updates[SettingPaymentVisibleMethodAlipayEnabled] = strconv.FormatBool(settings.PaymentVisibleMethodAlipayEnabled)
@@ -1856,6 +1954,8 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting
updates[SettingKeyDefaultPlatformQuotas] = string(blob)
}
+ updates[SettingKeyAllowUserViewErrorRequests] = strconv.FormatBool(settings.AllowUserViewErrorRequests)
+
return updates, nil
}
@@ -1979,9 +2079,25 @@ func (s *SettingService) refreshCachedSettings(settings *SystemSettings) {
enabled: settings.OpenAIAdvancedSchedulerEnabled,
expiresAt: time.Now().Add(openAIAdvancedSchedulerSettingCacheTTL).UnixNano(),
})
+ // Invalidate the quota auto-pause cache and let the next read trigger a fresh load.
+ // We can't know from here whether ops_advanced_settings was also touched, so be
+ // defensive: store an expired entry — GetOpenAIQuotaAutoPauseSettings will serve
+ // stale and kick off an async refresh, never blocking the request that follows.
+ s.openAIQuotaAutoPauseSettingsSF.Forget(openAIQuotaAutoPauseSettingsRefreshKey)
+ if cached, _ := s.openAIQuotaAutoPauseSettingsCache.Load().(*cachedOpenAIQuotaAutoPauseSettings); cached != nil {
+ s.openAIQuotaAutoPauseSettingsCache.Store(&cachedOpenAIQuotaAutoPauseSettings{
+ settings: cached.settings,
+ expiresAt: 0,
+ })
+ }
if s.cfg != nil {
s.cfg.SetTrustForwardedIPForAPIKeyACL(settings.APIKeyACLTrustForwardedIP)
}
+ s.openAIAllowCodexPluginSF.Forget("openai_allow_codex_plugin_enabled")
+ s.openAIAllowCodexPluginCache.Store(&cachedOpenAIAllowCodexPlugin{
+ value: settings.OpenAIAllowClaudeCodeCodexPlugin,
+ expiresAt: time.Now().Add(openAIAllowCodexPluginCacheTTL).UnixNano(),
+ })
if s.onUpdate != nil {
s.onUpdate() // Invalidate cache after settings update
}
@@ -2732,6 +2848,8 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error {
SettingPaymentVisibleMethodAlipayEnabled: "false",
SettingPaymentVisibleMethodWxpayEnabled: "false",
openAIAdvancedSchedulerSettingKey: "false",
+
+ SettingKeyAllowUserViewErrorRequests: "false",
}
return s.settingRepo.SetMultiple(ctx, defaults)
@@ -3247,6 +3365,7 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin
}
result.AntigravityUserAgentVersion = antigravity.NormalizeUserAgentVersion(settings[SettingKeyAntigravityUserAgentVersion])
result.OpenAICodexUserAgent = strings.TrimSpace(settings[SettingKeyOpenAICodexUserAgent])
+ result.OpenAIAllowClaudeCodeCodexPlugin = settings[SettingKeyOpenAIAllowClaudeCodeCodexPlugin] == "true"
// Web search emulation: quick enabled check from the JSON config
if raw := settings[SettingKeyWebSearchEmulationConfig]; raw != "" {
@@ -3288,6 +3407,8 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin
}
}
+ result.AllowUserViewErrorRequests = settings[SettingKeyAllowUserViewErrorRequests] == "true" // default false
+
return result
}
@@ -4394,6 +4515,106 @@ func (s *SettingService) GetClaudeCodeVersionBounds(ctx context.Context) (min, m
return b.min, b.max
}
+// GetOpenAIQuotaAutoPauseSettings returns the current global default quota auto-pause
+// settings. It is invoked on the OpenAI scheduling hot path (once per request) and is
+// therefore designed to never block on the DB:
+//
+// - Fresh cached value → returned immediately.
+// - Stale or empty cache → the last known value is returned, and a background
+// goroutine refreshes the cache via singleflight (stale-while-revalidate).
+// - First call with no cache yet → zero defaults are returned and the same async
+// refresh is kicked off; the next call gets the freshly populated value.
+//
+// Callers that need the freshly persisted value synchronously (tests, post-update
+// confirmation, optional startup warm-up) should call WarmOpenAIQuotaAutoPauseSettings.
+func (s *SettingService) GetOpenAIQuotaAutoPauseSettings(ctx context.Context) OpsOpenAIAccountQuotaAutoPauseSettings {
+ if s == nil {
+ return OpsOpenAIAccountQuotaAutoPauseSettings{}
+ }
+ cached, _ := s.openAIQuotaAutoPauseSettingsCache.Load().(*cachedOpenAIQuotaAutoPauseSettings)
+ now := time.Now().UnixNano()
+ if cached != nil && now < cached.expiresAt {
+ return cached.settings
+ }
+ // Stale or unset: trigger background refresh without blocking this request.
+ // singleflight.DoChan dedupes concurrent refreshes; we deliberately ignore the
+ // returned channel — the result is observable via the atomic cache.
+ s.openAIQuotaAutoPauseSettingsSF.DoChan(openAIQuotaAutoPauseSettingsRefreshKey, func() (any, error) {
+ s.refreshOpenAIQuotaAutoPauseSettings(context.Background())
+ return nil, nil
+ })
+ if cached != nil {
+ return cached.settings // serve stale value while revalidating
+ }
+ return OpsOpenAIAccountQuotaAutoPauseSettings{}
+}
+
+// WarmOpenAIQuotaAutoPauseSettings synchronously loads the quota auto-pause settings
+// into the in-memory cache. Useful for application startup (so the first request hits
+// a warm cache) and for tests that need deterministic reads immediately after
+// constructing the service.
+func (s *SettingService) WarmOpenAIQuotaAutoPauseSettings(ctx context.Context) OpsOpenAIAccountQuotaAutoPauseSettings {
+ if s == nil {
+ return OpsOpenAIAccountQuotaAutoPauseSettings{}
+ }
+ s.refreshOpenAIQuotaAutoPauseSettings(ctx)
+ cached, _ := s.openAIQuotaAutoPauseSettingsCache.Load().(*cachedOpenAIQuotaAutoPauseSettings)
+ if cached == nil {
+ return OpsOpenAIAccountQuotaAutoPauseSettings{}
+ }
+ return cached.settings
+}
+
+// refreshOpenAIQuotaAutoPauseSettings reads the latest settings from the DB and stores
+// them into the in-memory cache. On error it stores the prior value (or zero defaults
+// if nothing is cached yet) with the shorter error TTL so the next refresh comes
+// sooner. Always uses its own timeout-bounded context to keep refresh latency
+// predictable regardless of the caller.
+func (s *SettingService) refreshOpenAIQuotaAutoPauseSettings(ctx context.Context) {
+ if s == nil || s.settingRepo == nil {
+ return
+ }
+ dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), openAIQuotaAutoPauseSettingsDBTimeout)
+ defer cancel()
+
+ settings := OpsOpenAIAccountQuotaAutoPauseSettings{}
+ ttl := openAIQuotaAutoPauseSettingsCacheTTL
+ raw, err := s.settingRepo.GetValue(dbCtx, SettingKeyOpsAdvancedSettings)
+ if err == nil {
+ cfg := defaultOpsAdvancedSettings()
+ if strings.TrimSpace(raw) != "" {
+ if jsonErr := json.Unmarshal([]byte(raw), cfg); jsonErr == nil {
+ normalizeOpsAdvancedSettings(cfg)
+ }
+ }
+ settings = cfg.OpenAIAccountQuotaAutoPause
+ } else if !errors.Is(err, ErrSettingNotFound) {
+ // Real error: keep serving prior value but refresh sooner.
+ if prior, _ := s.openAIQuotaAutoPauseSettingsCache.Load().(*cachedOpenAIQuotaAutoPauseSettings); prior != nil {
+ settings = prior.settings
+ }
+ ttl = openAIQuotaAutoPauseSettingsErrorTTL
+ }
+
+ s.openAIQuotaAutoPauseSettingsCache.Store(&cachedOpenAIQuotaAutoPauseSettings{
+ settings: settings,
+ expiresAt: time.Now().Add(ttl).UnixNano(),
+ })
+}
+
+// SetOpenAIQuotaAutoPauseSettings writes the given settings directly into the in-memory
+// cache. Called from settings-write code paths so that the next read reflects the new
+// value immediately, without waiting for the background refresh.
+func (s *SettingService) SetOpenAIQuotaAutoPauseSettings(settings OpsOpenAIAccountQuotaAutoPauseSettings) {
+ if s == nil {
+ return
+ }
+ s.openAIQuotaAutoPauseSettingsCache.Store(&cachedOpenAIQuotaAutoPauseSettings{
+ settings: settings,
+ expiresAt: time.Now().Add(openAIQuotaAutoPauseSettingsCacheTTL).UnixNano(),
+ })
+}
+
// GetRectifierSettings 获取请求整流器配置
func (s *SettingService) GetRectifierSettings(ctx context.Context) (*RectifierSettings, error) {
value, err := s.settingRepo.GetValue(ctx, SettingKeyRectifierSettings)
diff --git a/backend/internal/service/setting_service_openai_allow_claude_code_test.go b/backend/internal/service/setting_service_openai_allow_claude_code_test.go
new file mode 100644
index 00000000..22059f07
--- /dev/null
+++ b/backend/internal/service/setting_service_openai_allow_claude_code_test.go
@@ -0,0 +1,55 @@
+package service
+
+import (
+ "context"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/stretchr/testify/require"
+)
+
+type allowClaudeCodeSettingRepoStub struct{ values map[string]string }
+
+func (s *allowClaudeCodeSettingRepoStub) Get(ctx context.Context, key string) (*Setting, error) {
+ panic("unused")
+}
+func (s *allowClaudeCodeSettingRepoStub) GetValue(ctx context.Context, key string) (string, error) {
+ if v, ok := s.values[key]; ok {
+ return v, nil
+ }
+ return "", ErrSettingNotFound
+}
+func (s *allowClaudeCodeSettingRepoStub) Set(ctx context.Context, key, value string) error {
+ panic("unused")
+}
+func (s *allowClaudeCodeSettingRepoStub) GetMultiple(ctx context.Context, keys []string) (map[string]string, error) {
+ panic("unused")
+}
+func (s *allowClaudeCodeSettingRepoStub) SetMultiple(ctx context.Context, settings map[string]string) error {
+ panic("unused")
+}
+func (s *allowClaudeCodeSettingRepoStub) GetAll(ctx context.Context) (map[string]string, error) {
+ panic("unused")
+}
+func (s *allowClaudeCodeSettingRepoStub) Delete(ctx context.Context, key string) error {
+ panic("unused")
+}
+
+func TestSettingService_IsOpenAIAllowClaudeCodeCodexPluginEnabled(t *testing.T) {
+ t.Run("默认关闭(设置缺失)", func(t *testing.T) {
+ svc := NewSettingService(&allowClaudeCodeSettingRepoStub{values: map[string]string{}}, &config.Config{})
+ require.False(t, svc.IsOpenAIAllowClaudeCodeCodexPluginEnabled(context.Background()))
+ })
+ t.Run("值为 true 时开启", func(t *testing.T) {
+ svc := NewSettingService(&allowClaudeCodeSettingRepoStub{values: map[string]string{
+ SettingKeyOpenAIAllowClaudeCodeCodexPlugin: "true",
+ }}, &config.Config{})
+ require.True(t, svc.IsOpenAIAllowClaudeCodeCodexPluginEnabled(context.Background()))
+ })
+ t.Run("值非 true 时关闭", func(t *testing.T) {
+ svc := NewSettingService(&allowClaudeCodeSettingRepoStub{values: map[string]string{
+ SettingKeyOpenAIAllowClaudeCodeCodexPlugin: "false",
+ }}, &config.Config{})
+ require.False(t, svc.IsOpenAIAllowClaudeCodeCodexPluginEnabled(context.Background()))
+ })
+}
diff --git a/backend/internal/service/setting_service_public_test.go b/backend/internal/service/setting_service_public_test.go
index cafe3051..3dedb1e4 100644
--- a/backend/internal/service/setting_service_public_test.go
+++ b/backend/internal/service/setting_service_public_test.go
@@ -135,6 +135,19 @@ func TestSettingService_GetPublicSettingsForInjection_PreservesTrustedDeployment
require.Nil(t, items[0].OmitAuthContext)
}
+func TestSettingService_GetPublicSettings_ExposesAllowUserViewErrorRequests(t *testing.T) {
+ repo := &settingPublicRepoStub{
+ values: map[string]string{
+ SettingKeyAllowUserViewErrorRequests: "true",
+ },
+ }
+ svc := NewSettingService(repo, &config.Config{})
+
+ settings, err := svc.GetPublicSettings(context.Background())
+ require.NoError(t, err)
+ require.True(t, settings.AllowUserViewErrorRequests)
+}
+
func TestSettingService_GetPublicSettings_ExposesWeChatOAuthModeCapabilities(t *testing.T) {
svc := NewSettingService(&settingPublicRepoStub{
values: map[string]string{
diff --git a/backend/internal/service/setting_service_user_error_persist_test.go b/backend/internal/service/setting_service_user_error_persist_test.go
new file mode 100644
index 00000000..1ec76d46
--- /dev/null
+++ b/backend/internal/service/setting_service_user_error_persist_test.go
@@ -0,0 +1,31 @@
+//go:build unit
+
+package service
+
+import (
+ "context"
+ "testing"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/stretchr/testify/require"
+)
+
+// TestAllowUserViewErrorRequests_PersistsToDB 验证 buildSystemSettingsUpdates 会将
+// AllowUserViewErrorRequests 写入 updates map(即最终落库),这是对 bug 的回归测试:
+// 该字段曾因漏写而永远无法持久化。
+func TestAllowUserViewErrorRequests_PersistsToDB(t *testing.T) {
+ // bmUpdateRepoStub 已在 setting_service_backend_mode_test.go 中定义(同 package)。
+ // 本测试不触及需要 GetValue 的设置项,getValueFn 设为 nil 即可,无需 stub。
+ repo := &bmUpdateRepoStub{}
+ svc := NewSettingService(repo, &config.Config{})
+
+ err := svc.UpdateSettings(context.Background(), &SystemSettings{
+ AllowUserViewErrorRequests: true,
+ })
+ require.NoError(t, err)
+
+ // 断言 updates 中含有该 key,且值为 "true"
+ val, ok := repo.updates[SettingKeyAllowUserViewErrorRequests]
+ require.True(t, ok, "updates map 中应包含 SettingKeyAllowUserViewErrorRequests,但未找到(bug:buildSystemSettingsUpdates 漏写)")
+ require.Equal(t, "true", val)
+}
diff --git a/backend/internal/service/setting_user_error_view_test.go b/backend/internal/service/setting_user_error_view_test.go
new file mode 100644
index 00000000..264a5a79
--- /dev/null
+++ b/backend/internal/service/setting_user_error_view_test.go
@@ -0,0 +1,9 @@
+package service
+
+import "testing"
+
+func TestSettingKeyAllowUserViewErrorRequests_Constant(t *testing.T) {
+ if SettingKeyAllowUserViewErrorRequests != "allow_user_view_error_requests" {
+ t.Fatalf("unexpected key: %s", SettingKeyAllowUserViewErrorRequests)
+ }
+}
diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go
index fb4bf485..a7146ddc 100644
--- a/backend/internal/service/settings_view.go
+++ b/backend/internal/service/settings_view.go
@@ -195,6 +195,7 @@ type SystemSettings struct {
RewriteMessageCacheControl bool // 是否改写 messages[*].content[*].cache_control(默认 false)
AntigravityUserAgentVersion string // Antigravity 上游 User-Agent 版本号;空值使用配置/默认值
OpenAICodexUserAgent string // OpenAI Codex 上游完整 User-Agent;空值使用内置默认
+ OpenAIAllowClaudeCodeCodexPlugin bool // 全局开关:是否额外放行 Claude Code 的 Codex 插件(默认 false)
// Web Search Emulation
WebSearchEmulationEnabled bool // 是否启用 web search 模拟
@@ -222,6 +223,9 @@ type SystemSettings struct {
// 系统全局默认平台配额(key = platform,nil/缺省 = 不限制)
DefaultPlatformQuotas map[string]*DefaultPlatformQuotaSetting `json:"default_platform_quotas"`
+
+ // 允许终端用户在用量页查看自己的失败请求
+ AllowUserViewErrorRequests bool
}
type DefaultSubscriptionSetting struct {
@@ -293,6 +297,9 @@ type PublicSettings struct {
// 风控中心功能开关
RiskControlEnabled bool `json:"risk_control_enabled"`
+
+ // 允许终端用户在用量页查看自己的失败请求
+ AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"`
}
type LoginAgreementDocument struct {
diff --git a/backend/internal/service/update_service.go b/backend/internal/service/update_service.go
index 34ad4610..de8c5e16 100644
--- a/backend/internal/service/update_service.go
+++ b/backend/internal/service/update_service.go
@@ -17,6 +17,12 @@ import (
"strconv"
"strings"
"time"
+
+ infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
+)
+
+var (
+ ErrNoUpdateAvailable = infraerrors.Conflict("ALREADY_UP_TO_DATE", "no update available; current version is latest")
)
const (
@@ -146,7 +152,7 @@ func (s *UpdateService) PerformUpdate(ctx context.Context) error {
}
if !info.HasUpdate {
- return fmt.Errorf("no update available")
+ return ErrNoUpdateAvailable
}
// Find matching archive and checksum for current platform
diff --git a/backend/internal/service/update_service_test.go b/backend/internal/service/update_service_test.go
new file mode 100644
index 00000000..8d8310d4
--- /dev/null
+++ b/backend/internal/service/update_service_test.go
@@ -0,0 +1,64 @@
+//go:build unit
+
+package service
+
+import (
+ "context"
+ "errors"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/require"
+)
+
+type updateServiceCacheStub struct {
+ data string
+}
+
+func (s *updateServiceCacheStub) GetUpdateInfo(context.Context) (string, error) {
+ if s.data == "" {
+ return "", errors.New("cache miss")
+ }
+ return s.data, nil
+}
+
+func (s *updateServiceCacheStub) SetUpdateInfo(_ context.Context, data string, _ time.Duration) error {
+ s.data = data
+ return nil
+}
+
+type updateServiceGitHubClientStub struct {
+ release *GitHubRelease
+}
+
+func (s *updateServiceGitHubClientStub) FetchLatestRelease(context.Context, string) (*GitHubRelease, error) {
+ return s.release, nil
+}
+
+func (s *updateServiceGitHubClientStub) DownloadFile(context.Context, string, string, int64) error {
+ panic("DownloadFile should not be called when no update is available")
+}
+
+func (s *updateServiceGitHubClientStub) FetchChecksumFile(context.Context, string) ([]byte, error) {
+ panic("FetchChecksumFile should not be called when no update is available")
+}
+
+func TestUpdateServicePerformUpdateNoUpdateReturnsSentinel(t *testing.T) {
+ svc := NewUpdateService(
+ &updateServiceCacheStub{},
+ &updateServiceGitHubClientStub{
+ release: &GitHubRelease{
+ TagName: "v0.1.132",
+ Name: "v0.1.132",
+ },
+ },
+ "0.1.132",
+ "release",
+ )
+
+ err := svc.PerformUpdate(context.Background())
+
+ require.Error(t, err)
+ require.True(t, errors.Is(err, ErrNoUpdateAvailable))
+ require.ErrorIs(t, err, ErrNoUpdateAvailable)
+}
diff --git a/backend/internal/service/user.go b/backend/internal/service/user.go
index f9833611..edb944ee 100644
--- a/backend/internal/service/user.go
+++ b/backend/internal/service/user.go
@@ -32,6 +32,7 @@ type User struct {
LastUsedAt *time.Time
CreatedAt time.Time
UpdatedAt time.Time
+ DeletedAt *time.Time // 非 nil 表示用户已软删除
// GroupRates 用户专属分组倍率配置
// map[groupID]rateMultiplier
diff --git a/backend/internal/service/user_msg_queue_service.go b/backend/internal/service/user_msg_queue_service.go
index a0ce95a8..f3f105ac 100644
--- a/backend/internal/service/user_msg_queue_service.go
+++ b/backend/internal/service/user_msg_queue_service.go
@@ -12,6 +12,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
+ "github.com/tidwall/gjson"
)
// UserMsgQueueCache 用户消息串行队列 Redis 缓存接口
@@ -62,43 +63,48 @@ func NewUserMessageQueueService(cache UserMsgQueueCache, rpmCache RPMCache, cfg
// 2. 最后一条消息 role == "user"
// 3. 最后一条消息 content(如果是数组)中不含 type:"tool_result" / "tool_use_result"
func IsRealUserMessage(parsed *ParsedRequest) bool {
- if parsed == nil || len(parsed.Messages) == 0 {
+ if parsed == nil {
+ return false
+ }
+ messagesRaw := parsed.MessagesRaw()
+ if len(messagesRaw) == 0 {
return false
}
- lastMsg := parsed.Messages[len(parsed.Messages)-1]
- msgMap, ok := lastMsg.(map[string]any)
- if !ok {
+ messages := gjson.ParseBytes(messagesRaw)
+ if !messages.IsArray() {
+ return false
+ }
+ lastMsg := gjson.Result{}
+ messages.ForEach(func(_, msg gjson.Result) bool {
+ lastMsg = msg
+ return true
+ })
+ if !lastMsg.Exists() || !lastMsg.IsObject() {
+ return false
+ }
+ if lastMsg.Get("role").String() != "user" {
return false
}
- role, _ := msgMap["role"].(string)
- if role != "user" {
- return false
+ content := lastMsg.Get("content")
+ if !content.Exists() {
+ return true
+ }
+ if !content.IsArray() {
+ return true
}
- // 检查 content 是否包含 tool_result 类型
- content, ok := msgMap["content"]
- if !ok {
- return true // 没有 content 字段,视为普通用户消息
- }
-
- contentArr, ok := content.([]any)
- if !ok {
- return true // content 不是数组(可能是 string),视为普通用户消息
- }
-
- for _, item := range contentArr {
- itemMap, ok := item.(map[string]any)
- if !ok {
- continue
- }
- itemType, _ := itemMap["type"].(string)
+ isReal := true
+ content.ForEach(func(_, item gjson.Result) bool {
+ itemType := item.Get("type").String()
if itemType == "tool_result" || itemType == "tool_use_result" {
+ isReal = false
return false
}
- }
- return true
+ return true
+ })
+ return isReal
}
// TryAcquire 尝试立即获取串行锁
diff --git a/backend/internal/service/user_platform_quota_flusher.go b/backend/internal/service/user_platform_quota_flusher.go
new file mode 100644
index 00000000..3ee23d2c
--- /dev/null
+++ b/backend/internal/service/user_platform_quota_flusher.go
@@ -0,0 +1,267 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "sync/atomic"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+ "github.com/Wei-Shaw/sub2api/internal/pkg/logger"
+)
+
+// quotaDirtyCache 是 flusher 依赖的窄接口(来自 BillingCache)。
+type quotaDirtyCache interface {
+ PopDirtyUserPlatformQuotaKeys(ctx context.Context, n int) ([]UserPlatformQuotaKey, error)
+ ReaddDirtyUserPlatformQuotaKeys(ctx context.Context, keys []UserPlatformQuotaKey) error
+ BatchGetUserPlatformQuotaCache(ctx context.Context, keys []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error)
+}
+
+// quotaSnapshotWriter 是 flusher 依赖的 DB 写入窄接口。
+// 使用 service 层的 UserPlatformQuotaSnapshot,避免与 repository 包形成循环依赖;
+// 实际实现由 repository adapter 在 B7 注入。
+type quotaSnapshotWriter interface {
+ BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error
+}
+
+// FlusherMetrics 记录 flusher 运行时指标(原子量,零值可用)。
+type FlusherMetrics struct {
+ FlushSuccessTotal atomic.Int64
+ FlushErrorTotal atomic.Int64
+ FlushBatchSizeTotal atomic.Int64
+ FlushLatencyMsMax atomic.Int64
+ DirtyReaddTotal atomic.Int64
+ // DirtyLostTotal:Readd 失败导致脏 key 丢失——已 SPOP+主操作失败+Readd 也失败;
+ // Redis 仍权威,活跃 key 下次 SADD 自愈。
+ DirtyLostTotal atomic.Int64
+ FlushFKViolationTotal atomic.Int64
+}
+
+// flusherMaxBatchesPerTick 单次 tick 最多消费的批数,防止 tick 执行时间过长。
+const flusherMaxBatchesPerTick = 16
+
+// maxFlushBatchSize 限制单批行数,必须 ≤ repository.BatchSnapshotUsage 的 batchRows(6000),
+// 以保证单次 flush 的 snapshots 仅生成一条 UPSERT(单事务原子)。两处需手动保持一致。
+const maxFlushBatchSize = 6000
+
+// defaultFlushBatchSize 是配置 flush_batch_size 非法(≤0)时的回退值。
+const defaultFlushBatchSize = 1000
+
+// UserPlatformQuotaUsageFlusher 将 Redis 脏集快照定期批量写入 DB。
+// 不维护任何 delta/in-process 状态;每批读取 Redis 当前绝对值覆盖写入。
+type UserPlatformQuotaUsageFlusher struct {
+ cache quotaDirtyCache
+ quotaRepo quotaSnapshotWriter
+ timingWheel *TimingWheelService
+ // enabled 对应 flusher_enabled 配置;false 时 Start() 不注册定时器。
+ enabled bool
+ interval time.Duration
+ batchSize int
+ flushTimeout time.Duration
+ metrics *FlusherMetrics
+ stopped atomic.Bool
+}
+
+// NewUserPlatformQuotaUsageFlusher 创建 UserPlatformQuotaUsageFlusher。
+// cache(BillingCache) 隐式满足 quotaDirtyCache;quotaRepo(UserPlatformQuotaRepository) 隐式满足 quotaSnapshotWriter。
+func NewUserPlatformQuotaUsageFlusher(cfg *config.Config, cache BillingCache, quotaRepo UserPlatformQuotaRepository, tw *TimingWheelService) *UserPlatformQuotaUsageFlusher {
+ batchSize := cfg.Database.UserPlatformQuotaFlushBatchSize
+ if batchSize <= 0 {
+ batchSize = defaultFlushBatchSize
+ }
+ if batchSize > maxFlushBatchSize {
+ logger.LegacyPrintf("quota_flusher",
+ "[QuotaFlusher] flush_batch_size %d 超过上限 %d,已 clamp(避免 BatchSnapshotUsage 多子批非原子)",
+ cfg.Database.UserPlatformQuotaFlushBatchSize, maxFlushBatchSize)
+ batchSize = maxFlushBatchSize
+ }
+ interval := time.Duration(cfg.Database.UserPlatformQuotaFlushIntervalMs) * time.Millisecond
+ if interval <= 0 {
+ logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] flush_interval_ms %d 非法,回退 2000ms", cfg.Database.UserPlatformQuotaFlushIntervalMs)
+ interval = 2 * time.Second
+ }
+ return &UserPlatformQuotaUsageFlusher{
+ cache: cache,
+ quotaRepo: quotaRepo,
+ timingWheel: tw,
+ enabled: cfg.Database.UserPlatformQuotaFlusherEnabled,
+ interval: interval,
+ batchSize: batchSize,
+ flushTimeout: 3 * time.Second,
+ metrics: &FlusherMetrics{},
+ }
+}
+
+// updateLatencyMax 用 CAS 单调更新最大延迟。
+func (s *UserPlatformQuotaUsageFlusher) updateLatencyMax(ms int64) {
+ for {
+ old := s.metrics.FlushLatencyMsMax.Load()
+ if ms <= old {
+ return
+ }
+ if s.metrics.FlushLatencyMsMax.CompareAndSwap(old, ms) {
+ return
+ }
+ }
+}
+
+// readdOrCountLost 尝试把 keys 回填脏集:成功计 DirtyReaddTotal,失败计 DirtyLostTotal 并 ALERT。
+func (s *UserPlatformQuotaUsageFlusher) readdOrCountLost(ctx context.Context, keys []UserPlatformQuotaKey, stage string) {
+ if err := s.cache.ReaddDirtyUserPlatformQuotaKeys(ctx, keys); err != nil {
+ s.metrics.DirtyLostTotal.Add(int64(len(keys)))
+ logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] ALERT: Readd after %s failed, %d keys 丢出脏集(DB 镜像缺这批,Redis 仍权威,活跃 key 下次 SADD 自愈): %v", stage, len(keys), err)
+ return
+ }
+ s.metrics.DirtyReaddTotal.Add(int64(len(keys)))
+}
+
+// flushOneBatch 处理单批:Pop → BatchGet → 组装 snaps → BatchSnapshotUsage。
+// 返回 (shouldContinue bool):false 表示本轮循环应停止(空集/错误/最后一批)。
+// 每次调用独立创建带 timeout 的 ctx 并 defer cancel,不会在循环中累积泄漏。
+func (s *UserPlatformQuotaUsageFlusher) flushOneBatch(parentCtx context.Context) bool {
+ ctx, cancel := context.WithTimeout(parentCtx, s.flushTimeout)
+ defer cancel()
+
+ // 1. Pop 脏集
+ keys, err := s.cache.PopDirtyUserPlatformQuotaKeys(ctx, s.batchSize)
+ if err != nil {
+ s.metrics.FlushErrorTotal.Add(1)
+ logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] PopDirty error: %v", err)
+ return false
+ }
+ if len(keys) == 0 {
+ // 脏集已空
+ return false
+ }
+
+ // 2. 批量读 Redis 快照
+ entries, err := s.cache.BatchGetUserPlatformQuotaCache(ctx, keys)
+ if err != nil {
+ s.metrics.FlushErrorTotal.Add(1)
+ s.readdOrCountLost(ctx, keys, "BatchGet")
+ logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] BatchGet error: %v", err)
+ return false
+ }
+
+ // 3. 组装 snapshots(MISS 或任一 WindowStart==nil → 跳过)
+ snaps := make([]UserPlatformQuotaSnapshot, 0, len(keys))
+ for i, key := range keys {
+ e := entries[i]
+ if e == nil {
+ continue
+ }
+ if e.DailyWindowStart == nil || e.WeeklyWindowStart == nil || e.MonthlyWindowStart == nil {
+ continue
+ }
+ snaps = append(snaps, UserPlatformQuotaSnapshot{
+ UserID: key.UserID,
+ Platform: key.Platform,
+ DailyUsageUSD: e.DailyUsageUSD,
+ WeeklyUsageUSD: e.WeeklyUsageUSD,
+ MonthlyUsageUSD: e.MonthlyUsageUSD,
+ DailyWindowStart: *e.DailyWindowStart,
+ WeeklyWindowStart: *e.WeeklyWindowStart,
+ MonthlyWindowStart: *e.MonthlyWindowStart,
+ })
+ }
+
+ // 4. 全部 MISS/异常跳过时
+ if len(snaps) == 0 {
+ // 若 Pop 数量已不满一批,表示脏集将空,停止
+ if len(keys) < s.batchSize {
+ return false
+ }
+ // 否则继续下一批(可能还有更多脏 key)
+ return true
+ }
+
+ // 已知竞态(admin 写 × flusher 刷,仅 flusher_enabled=true 时存在):
+ // admin ResetExpiredWindow/UpsertForUser 是"先写 DB 再 DeleteCache"。若本批已 SPOP + BatchGet
+ // 读到旧 usage 快照(此刻 member 已离开脏集),而 admin 随后写 DB、本行 UPSERT 又在 admin 写之后落库,
+ // 则旧快照会覆盖 admin 刚写入的值;DeleteCache 后 Redis MISS,下次 preflight 从 DB 重载被覆盖的旧值。
+ // 因 member 已被 SPOP,admin 侧 SREM/清脏标记无法拦截本批(故未做)。影响有限,暂列为已知取舍:
+ // - UpsertForUser 改 limit,而本 UPSERT 不写 limit 列 → limit 配置不受影响;
+ // - ResetExpiredWindow 改 usage,但 preflight windowExpired 会在窗口真正过期时自愈重置,
+ // 仅"强制重置未过期窗口"且与本批精确交错时短暂失效;
+ // - 低频 admin 操作 + 默认 flusher_enabled=false。彻底消除需 version OCC(DB 加 version 列条件 UPSERT),
+ // 成本高;启用 flusher 后如需强一致再评估。
+
+ // 5. 写入 DB
+ start := time.Now()
+ writeErr := s.quotaRepo.BatchSnapshotUsage(ctx, snaps, time.Now().UTC())
+ s.updateLatencyMax(time.Since(start).Milliseconds())
+
+ if writeErr != nil {
+ if errors.Is(writeErr, ErrUserPlatformQuotaFKViolation) {
+ // 注意:PG FK violation 是整条 INSERT 回滚 → 整批(含同批正常用户)均未写入 DB,
+ // 且这些 key 已被 SPOP 出脏集、此处不 Readd。活跃 key 会在下次请求重新 SADD,
+ // flusher 读 Redis 当前累计绝对值刷库即自愈;低活跃 key 这轮 DB usage 偏低
+ // (Redis 仍是 enforcement 权威,不受影响;DB 仅展示)。已删用户边角的接受取舍,不做逐行重试。
+ // FK 违反:用户已被删除,直接丢弃不 Readd
+ s.metrics.FlushFKViolationTotal.Add(1)
+ s.metrics.FlushErrorTotal.Add(1)
+ logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] FK violation (dropped %d snaps): %v", len(snaps), writeErr)
+ } else {
+ // 其他错误:回填脏集,保留下次重试
+ s.metrics.FlushErrorTotal.Add(1)
+ s.readdOrCountLost(ctx, keys, "BatchSnapshotUsage")
+ logger.LegacyPrintf("quota_flusher", "[QuotaFlusher] BatchSnapshotUsage error: %v", writeErr)
+ }
+ return false
+ }
+
+ // 6. 成功
+ s.metrics.FlushSuccessTotal.Add(1)
+ s.metrics.FlushBatchSizeTotal.Add(int64(len(snaps)))
+
+ // 若 Pop 数量不满一批,脏集已空,停止
+ if len(keys) < s.batchSize {
+ return false
+ }
+ return true
+}
+
+// flush 执行一次完整的 flush,循环消费至脏集空或达到 maxBatchesPerTick。
+func (s *UserPlatformQuotaUsageFlusher) flush() {
+ if s == nil {
+ return
+ }
+ parentCtx := context.Background()
+ for b := 0; b < flusherMaxBatchesPerTick; b++ {
+ if !s.flushOneBatch(parentCtx) {
+ return
+ }
+ }
+ // 连续消费满 flusherMaxBatchesPerTick 批仍未取空脏集:本 tick 主动让出,剩余积压留待下一 tick。
+ // 记一条 log 便于 oncall 发现 distinct 活跃 key 远超 maxBatchesPerTick×batchSize(DB 镜像延迟上升);
+ // 可配合 Redis SCARD billing:upq:dirty 观察脏集存量。
+ logger.LegacyPrintf("quota_flusher",
+ "[QuotaFlusher] 单 tick 达到 max batches 上限(%d × batchSize=%d),脏集仍非空,积压顺延至下一 tick",
+ flusherMaxBatchesPerTick, s.batchSize)
+}
+
+// tick 是 TimingWheel 回调。若 flusher 已停止则直接返回。
+func (s *UserPlatformQuotaUsageFlusher) tick() {
+ if s == nil || s.stopped.Load() {
+ return
+ }
+ s.flush()
+}
+
+// Start 注册定时 tick。flusher_enabled=false 时直接返回,不注册定时器。
+func (s *UserPlatformQuotaUsageFlusher) Start() {
+ if s == nil || !s.enabled {
+ return
+ }
+ s.timingWheel.ScheduleRecurring("deferred:platform_quota", s.interval, s.tick)
+}
+
+// Stop 停止 flusher:标记 stopped → Cancel 定时器 → 执行最后一次 flush。
+func (s *UserPlatformQuotaUsageFlusher) Stop() {
+ if s == nil {
+ return
+ }
+ s.stopped.Store(true)
+ s.timingWheel.Cancel("deferred:platform_quota")
+ s.flush()
+}
diff --git a/backend/internal/service/user_platform_quota_flusher_test.go b/backend/internal/service/user_platform_quota_flusher_test.go
new file mode 100644
index 00000000..4f734481
--- /dev/null
+++ b/backend/internal/service/user_platform_quota_flusher_test.go
@@ -0,0 +1,511 @@
+package service
+
+import (
+ "context"
+ "errors"
+ "testing"
+ "time"
+
+ "github.com/Wei-Shaw/sub2api/internal/config"
+)
+
+// ---------------------------------------------------------------------------
+// Mock: quotaDirtyCache
+// ---------------------------------------------------------------------------
+
+type mockQuotaDirtyCache struct {
+ // popSequence: 第 0 次 Pop 返回 popSequence[0],之后返回 nil(空集)
+ popSequence [][]UserPlatformQuotaKey
+ popCallIdx int
+
+ // getEntries: BatchGetUserPlatformQuotaCache 返回的 entries(与 keys 对齐)
+ getEntries []*UserPlatformQuotaCacheEntry
+ getErr error
+
+ // readdCalled: 记录 Readd 收到的 keys(累积所有次调用)
+ readdCalled [][]UserPlatformQuotaKey
+ readdErr error
+}
+
+func (m *mockQuotaDirtyCache) PopDirtyUserPlatformQuotaKeys(_ context.Context, _ int) ([]UserPlatformQuotaKey, error) {
+ if m.popCallIdx < len(m.popSequence) {
+ keys := m.popSequence[m.popCallIdx]
+ m.popCallIdx++
+ return keys, nil
+ }
+ // 超出序列 → 空集(模拟脏集已清空)
+ return nil, nil
+}
+
+func (m *mockQuotaDirtyCache) ReaddDirtyUserPlatformQuotaKeys(_ context.Context, keys []UserPlatformQuotaKey) error {
+ m.readdCalled = append(m.readdCalled, keys)
+ return m.readdErr
+}
+
+func (m *mockQuotaDirtyCache) BatchGetUserPlatformQuotaCache(_ context.Context, _ []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) {
+ if m.getErr != nil {
+ return nil, m.getErr
+ }
+ return m.getEntries, nil
+}
+
+// ---------------------------------------------------------------------------
+// Mock: quotaSnapshotWriter
+// ---------------------------------------------------------------------------
+
+type mockQuotaSnapshotWriter struct {
+ receivedSnaps []UserPlatformQuotaSnapshot
+ returnErr error
+}
+
+func (m *mockQuotaSnapshotWriter) BatchSnapshotUsage(_ context.Context, snaps []UserPlatformQuotaSnapshot, _ time.Time) error {
+ m.receivedSnaps = append(m.receivedSnaps, snaps...)
+ return m.returnErr
+}
+
+// ---------------------------------------------------------------------------
+// Helper: 构造窗口起始时间(非 nil)
+// ---------------------------------------------------------------------------
+
+func flusherPtrTime(t time.Time) *time.Time { return &t }
+
+func makeEntry(daily, weekly, monthly float64) *UserPlatformQuotaCacheEntry {
+ now := time.Now().UTC()
+ return &UserPlatformQuotaCacheEntry{
+ DailyUsageUSD: daily,
+ WeeklyUsageUSD: weekly,
+ MonthlyUsageUSD: monthly,
+ DailyWindowStart: flusherPtrTime(now),
+ WeeklyWindowStart: flusherPtrTime(now),
+ MonthlyWindowStart: flusherPtrTime(now),
+ }
+}
+
+// ---------------------------------------------------------------------------
+// newTestFlusher: 直接构造 struct(跳过构造函数,B7 才注入)
+// ---------------------------------------------------------------------------
+
+func newTestFlusher(cache quotaDirtyCache, writer quotaSnapshotWriter) *UserPlatformQuotaUsageFlusher {
+ return &UserPlatformQuotaUsageFlusher{
+ cache: cache,
+ quotaRepo: writer,
+ timingWheel: nil, // 单测不启动 TimingWheel
+ interval: 5 * time.Second,
+ batchSize: 100,
+ flushTimeout: 5 * time.Second,
+ metrics: &FlusherMetrics{},
+ }
+}
+
+// ---------------------------------------------------------------------------
+// 场景 1: PopSnapshotUpsert — 2 key + 2 个含 window 的 entry → writer 收 2 行
+// ---------------------------------------------------------------------------
+
+func TestFlusher_PopSnapshotUpsert(t *testing.T) {
+ keys := []UserPlatformQuotaKey{
+ {UserID: 1, Platform: "anthropic"},
+ {UserID: 2, Platform: "openai"},
+ }
+ cache := &mockQuotaDirtyCache{
+ popSequence: [][]UserPlatformQuotaKey{keys}, // 第 1 次返回 keys,之后空
+ getEntries: []*UserPlatformQuotaCacheEntry{
+ makeEntry(1.0, 2.0, 3.0),
+ makeEntry(4.0, 5.0, 6.0),
+ },
+ }
+ writer := &mockQuotaSnapshotWriter{}
+ f := newTestFlusher(cache, writer)
+
+ f.flush()
+
+ if len(writer.receivedSnaps) != 2 {
+ t.Fatalf("expected 2 snaps, got %d", len(writer.receivedSnaps))
+ }
+ if f.metrics.FlushBatchSizeTotal.Load() != 2 {
+ t.Errorf("FlushBatchSizeTotal = %d, want 2", f.metrics.FlushBatchSizeTotal.Load())
+ }
+ if f.metrics.FlushSuccessTotal.Load() != 1 {
+ t.Errorf("FlushSuccessTotal = %d, want 1", f.metrics.FlushSuccessTotal.Load())
+ }
+ if f.metrics.FlushErrorTotal.Load() != 0 {
+ t.Errorf("FlushErrorTotal = %d, want 0", f.metrics.FlushErrorTotal.Load())
+ }
+}
+
+// ---------------------------------------------------------------------------
+// 场景 2: MissKeySkipped — 2 key,BatchGet 返回 [entry, nil] → 只刷 1 行,nil 跳过,不 Readd
+// ---------------------------------------------------------------------------
+
+func TestFlusher_MissKeySkipped(t *testing.T) {
+ keys := []UserPlatformQuotaKey{
+ {UserID: 1, Platform: "anthropic"},
+ {UserID: 2, Platform: "openai"},
+ }
+ cache := &mockQuotaDirtyCache{
+ popSequence: [][]UserPlatformQuotaKey{keys},
+ getEntries: []*UserPlatformQuotaCacheEntry{
+ makeEntry(1.0, 2.0, 3.0),
+ nil, // MISS
+ },
+ }
+ writer := &mockQuotaSnapshotWriter{}
+ f := newTestFlusher(cache, writer)
+
+ f.flush()
+
+ if len(writer.receivedSnaps) != 1 {
+ t.Fatalf("expected 1 snap, got %d", len(writer.receivedSnaps))
+ }
+ if writer.receivedSnaps[0].UserID != 1 {
+ t.Errorf("expected snap for UserID=1, got %d", writer.receivedSnaps[0].UserID)
+ }
+ if len(cache.readdCalled) != 0 {
+ t.Errorf("Readd should NOT be called on MISS, got %d calls", len(cache.readdCalled))
+ }
+ if f.metrics.FlushSuccessTotal.Load() != 1 {
+ t.Errorf("FlushSuccessTotal = %d, want 1", f.metrics.FlushSuccessTotal.Load())
+ }
+}
+
+// ---------------------------------------------------------------------------
+// 场景 3: UpsertFailReadds — writer 返普通 error → keys 被 Readd,FlushErrorTotal=1,DirtyReaddTotal=len
+// ---------------------------------------------------------------------------
+
+func TestFlusher_UpsertFailReadds(t *testing.T) {
+ keys := []UserPlatformQuotaKey{
+ {UserID: 1, Platform: "anthropic"},
+ {UserID: 2, Platform: "openai"},
+ }
+ cache := &mockQuotaDirtyCache{
+ popSequence: [][]UserPlatformQuotaKey{keys},
+ getEntries: []*UserPlatformQuotaCacheEntry{
+ makeEntry(1.0, 2.0, 3.0),
+ makeEntry(4.0, 5.0, 6.0),
+ },
+ }
+ writeErr := errors.New("db connection timeout")
+ writer := &mockQuotaSnapshotWriter{returnErr: writeErr}
+ f := newTestFlusher(cache, writer)
+
+ f.flush()
+
+ if f.metrics.FlushErrorTotal.Load() != 1 {
+ t.Errorf("FlushErrorTotal = %d, want 1", f.metrics.FlushErrorTotal.Load())
+ }
+ if len(cache.readdCalled) == 0 {
+ t.Fatal("Readd should be called after write error")
+ }
+ totalReadd := 0
+ for _, rk := range cache.readdCalled {
+ totalReadd += len(rk)
+ }
+ if totalReadd != len(keys) {
+ t.Errorf("DirtyReaddTotal (from Readd calls) = %d, want %d", totalReadd, len(keys))
+ }
+ if f.metrics.DirtyReaddTotal.Load() != int64(len(keys)) {
+ t.Errorf("DirtyReaddTotal metric = %d, want %d", f.metrics.DirtyReaddTotal.Load(), len(keys))
+ }
+ if f.metrics.FlushSuccessTotal.Load() != 0 {
+ t.Errorf("FlushSuccessTotal = %d, want 0", f.metrics.FlushSuccessTotal.Load())
+ }
+}
+
+// ---------------------------------------------------------------------------
+// 场景 4: FKViolationDropsNoReadd — writer 返 ErrUserPlatformQuotaFKViolation → 不 Readd,FlushFKViolationTotal=1
+// ---------------------------------------------------------------------------
+
+func TestFlusher_FKViolationDropsNoReadd(t *testing.T) {
+ keys := []UserPlatformQuotaKey{
+ {UserID: 999, Platform: "anthropic"},
+ }
+ cache := &mockQuotaDirtyCache{
+ popSequence: [][]UserPlatformQuotaKey{keys},
+ getEntries: []*UserPlatformQuotaCacheEntry{
+ makeEntry(1.0, 2.0, 3.0),
+ },
+ }
+ writer := &mockQuotaSnapshotWriter{returnErr: ErrUserPlatformQuotaFKViolation}
+ f := newTestFlusher(cache, writer)
+
+ f.flush()
+
+ if f.metrics.FlushFKViolationTotal.Load() != 1 {
+ t.Errorf("FlushFKViolationTotal = %d, want 1", f.metrics.FlushFKViolationTotal.Load())
+ }
+ if f.metrics.FlushErrorTotal.Load() != 1 {
+ t.Errorf("FlushErrorTotal = %d, want 1", f.metrics.FlushErrorTotal.Load())
+ }
+ if len(cache.readdCalled) != 0 {
+ t.Errorf("Readd should NOT be called for FK violation (drop), got %d calls", len(cache.readdCalled))
+ }
+ if f.metrics.DirtyReaddTotal.Load() != 0 {
+ t.Errorf("DirtyReaddTotal = %d, want 0 (FK violation drops)", f.metrics.DirtyReaddTotal.Load())
+ }
+}
+
+// ---------------------------------------------------------------------------
+// 场景 5: NilSafe — var f *UserPlatformQuotaUsageFlusher; f.flush(); f.Stop() 不 panic
+// ---------------------------------------------------------------------------
+
+func TestFlusher_NilSafe(t *testing.T) {
+ var f *UserPlatformQuotaUsageFlusher
+ // 下面两行不应 panic
+ f.flush()
+ f.Stop()
+}
+
+// ---------------------------------------------------------------------------
+// 场景 6: StopPreventsFlush — stopped=true 后 tick() 不调 flush(writer 没收到 snaps)
+// ---------------------------------------------------------------------------
+
+func TestFlusher_StopPreventsFlush(t *testing.T) {
+ keys := []UserPlatformQuotaKey{
+ {UserID: 1, Platform: "anthropic"},
+ }
+ cache := &mockQuotaDirtyCache{
+ popSequence: [][]UserPlatformQuotaKey{keys},
+ getEntries: []*UserPlatformQuotaCacheEntry{
+ makeEntry(1.0, 2.0, 3.0),
+ },
+ }
+ writer := &mockQuotaSnapshotWriter{}
+ f := newTestFlusher(cache, writer)
+
+ // 标记为已停止
+ f.stopped.Store(true)
+
+ // tick 应该直接返回,不触发 flush
+ f.tick()
+
+ if len(writer.receivedSnaps) != 0 {
+ t.Errorf("expected 0 snaps after stop, got %d", len(writer.receivedSnaps))
+ }
+ if cache.popCallIdx != 0 {
+ t.Errorf("Pop should not be called after stop, popCallIdx = %d", cache.popCallIdx)
+ }
+}
+
+// ---------------------------------------------------------------------------
+// 场景 B13-1: ZeroPercentCompany — 0% 公司脏集恒空,flusher 空跑无 DB 写
+//
+// 模拟几乎没有用户配置 quota limit 的公司:脏集始终为空(popSequence 为空切片),
+// Pop 每次返回空集。flush() 应早退,不写 DB、不计成功、不 Readd。
+// ---------------------------------------------------------------------------
+
+func TestScenario_ZeroPercentCompany(t *testing.T) {
+ cache := &mockQuotaDirtyCache{
+ // popSequence 为空 → Pop 超出序列 → 始终返回 nil(空集)
+ popSequence: [][]UserPlatformQuotaKey{},
+ }
+ writer := &mockQuotaSnapshotWriter{}
+ f := newTestFlusher(cache, writer)
+
+ f.flush()
+
+ if len(writer.receivedSnaps) != 0 {
+ t.Errorf("0%% company: expected 0 snaps, got %d", len(writer.receivedSnaps))
+ }
+ if f.metrics.FlushBatchSizeTotal.Load() != 0 {
+ t.Errorf("0%% company: FlushBatchSizeTotal = %d, want 0", f.metrics.FlushBatchSizeTotal.Load())
+ }
+ if f.metrics.FlushSuccessTotal.Load() != 0 {
+ t.Errorf("0%% company: FlushSuccessTotal = %d, want 0 (empty-set early return)", f.metrics.FlushSuccessTotal.Load())
+ }
+ if f.metrics.FlushErrorTotal.Load() != 0 {
+ t.Errorf("0%% company: FlushErrorTotal = %d, want 0", f.metrics.FlushErrorTotal.Load())
+ }
+ if len(cache.readdCalled) != 0 {
+ t.Errorf("0%% company: Readd should never be called, got %d calls", len(cache.readdCalled))
+ }
+}
+
+// ---------------------------------------------------------------------------
+// P1: IntervalFallback — flush_interval_ms ≤0 时回退 2s;正常值保留
+// ---------------------------------------------------------------------------
+
+func TestNewUserPlatformQuotaUsageFlusher_IntervalFallback(t *testing.T) {
+ cases := []struct {
+ name string
+ inMs int
+ wantDu time.Duration
+ }{
+ {"零值回退 2s", 0, 2 * time.Second},
+ {"负数回退 2s", -100, 2 * time.Second},
+ {"正常 2000ms 保留", 2000, 2 * time.Second},
+ {"正常 500ms 保留", 500, 500 * time.Millisecond},
+ }
+ for _, tc := range cases {
+ t.Run(tc.name, func(t *testing.T) {
+ cfg := &config.Config{}
+ cfg.Database.UserPlatformQuotaFlushIntervalMs = tc.inMs
+ f := NewUserPlatformQuotaUsageFlusher(cfg, nil, nil, nil)
+ if f.interval != tc.wantDu {
+ t.Fatalf("interval = %v, want %v", f.interval, tc.wantDu)
+ }
+ })
+ }
+}
+
+// ---------------------------------------------------------------------------
+// P1: EnabledField — flusher_enabled 配置正确写入 f.enabled
+// ---------------------------------------------------------------------------
+
+func TestNewUserPlatformQuotaUsageFlusher_EnabledField(t *testing.T) {
+ for _, enabled := range []bool{true, false} {
+ cfg := &config.Config{}
+ cfg.Database.UserPlatformQuotaFlusherEnabled = enabled
+ cfg.Database.UserPlatformQuotaFlushIntervalMs = 500
+ f := NewUserPlatformQuotaUsageFlusher(cfg, nil, nil, nil)
+ if f.enabled != enabled {
+ t.Errorf("enabled = %v, want %v", f.enabled, enabled)
+ }
+ }
+}
+
+// ---------------------------------------------------------------------------
+// P2: ReaddFailCounts — BatchGet 失败 + Readd 失败 → DirtyLostTotal 增、DirtyReaddTotal 不变
+// BatchGet 失败 + Readd 成功 → DirtyReaddTotal 增、DirtyLostTotal 不变
+// ---------------------------------------------------------------------------
+
+func TestFlusher_ReaddFailCounts(t *testing.T) {
+ keys := []UserPlatformQuotaKey{
+ {UserID: 10, Platform: "anthropic"},
+ {UserID: 11, Platform: "openai"},
+ }
+
+ t.Run("Readd 失败计 DirtyLostTotal", func(t *testing.T) {
+ cache := &mockQuotaDirtyCache{
+ popSequence: [][]UserPlatformQuotaKey{keys},
+ getErr: errors.New("redis timeout"), // 触发 BatchGet 失败路径
+ readdErr: errors.New("redis connection refused"), // Readd 也失败
+ }
+ f := newTestFlusher(cache, &mockQuotaSnapshotWriter{})
+
+ f.flush()
+
+ if f.metrics.DirtyLostTotal.Load() != int64(len(keys)) {
+ t.Errorf("DirtyLostTotal = %d, want %d", f.metrics.DirtyLostTotal.Load(), len(keys))
+ }
+ if f.metrics.DirtyReaddTotal.Load() != 0 {
+ t.Errorf("DirtyReaddTotal = %d, want 0 (Readd 失败不应计入)", f.metrics.DirtyReaddTotal.Load())
+ }
+ })
+
+ t.Run("Readd 成功计 DirtyReaddTotal", func(t *testing.T) {
+ cache := &mockQuotaDirtyCache{
+ popSequence: [][]UserPlatformQuotaKey{keys},
+ getErr: errors.New("redis timeout"), // 触发 BatchGet 失败路径
+ readdErr: nil, // Readd 成功
+ }
+ f := newTestFlusher(cache, &mockQuotaSnapshotWriter{})
+
+ f.flush()
+
+ if f.metrics.DirtyReaddTotal.Load() != int64(len(keys)) {
+ t.Errorf("DirtyReaddTotal = %d, want %d", f.metrics.DirtyReaddTotal.Load(), len(keys))
+ }
+ if f.metrics.DirtyLostTotal.Load() != 0 {
+ t.Errorf("DirtyLostTotal = %d, want 0 (Readd 成功不应计 lost)", f.metrics.DirtyLostTotal.Load())
+ }
+ })
+}
+
+// ---------------------------------------------------------------------------
+// ClampsBatchSize — NewUserPlatformQuotaUsageFlusher 构造时按
+// [defaultFlushBatchSize, maxFlushBatchSize] 区间 clamp batchSize
+// ---------------------------------------------------------------------------
+
+func TestNewUserPlatformQuotaUsageFlusher_ClampsBatchSize(t *testing.T) {
+ cases := []struct {
+ name string
+ in int
+ want int
+ }{
+ {"超上限被 clamp", 7000, maxFlushBatchSize},
+ {"恰好上限保留", maxFlushBatchSize, maxFlushBatchSize},
+ {"零回退默认", 0, defaultFlushBatchSize},
+ {"负数回退默认", -5, defaultFlushBatchSize},
+ {"正常值保留", 500, 500},
+ }
+ for _, tc := range cases {
+ t.Run(tc.name, func(t *testing.T) {
+ cfg := &config.Config{}
+ cfg.Database.UserPlatformQuotaFlushBatchSize = tc.in
+ f := NewUserPlatformQuotaUsageFlusher(cfg, nil, nil, nil)
+ if f.batchSize != tc.want {
+ t.Fatalf("batchSize = %d, want %d", f.batchSize, tc.want)
+ }
+ })
+ }
+}
+
+// ---------------------------------------------------------------------------
+// 场景 B13-2: NinetyPercentCompany — 90% 公司大量用户配 limit,一批 5 key 批量刷库
+//
+// 模拟大量用户配置了 quota limit 的公司:脏集第一次 Pop 返回 5 个不同用户的 key,
+// 之后返回空集(避免 flush 循环)。flush() 应构造 5 条 snapshot 写入 DB,
+// 断言绝对值语义(snap 的 DailyUsageUSD 等于 entry 的值)、metrics 正确、不 Readd。
+// ---------------------------------------------------------------------------
+
+func TestScenario_NinetyPercentCompany(t *testing.T) {
+ keys := []UserPlatformQuotaKey{
+ {UserID: 101, Platform: "anthropic"},
+ {UserID: 102, Platform: "anthropic"},
+ {UserID: 103, Platform: "openai"},
+ {UserID: 104, Platform: "openai"},
+ {UserID: 105, Platform: "anthropic"},
+ }
+ entries := []*UserPlatformQuotaCacheEntry{
+ makeEntry(1.1, 2.2, 3.3),
+ makeEntry(4.4, 5.5, 6.6),
+ makeEntry(7.7, 8.8, 9.9),
+ makeEntry(0.5, 1.0, 1.5),
+ makeEntry(10.0, 20.0, 30.0),
+ }
+ cache := &mockQuotaDirtyCache{
+ // 第 1 次 Pop 返回 5 keys,之后返回空集(防止 flush 无限循环)
+ popSequence: [][]UserPlatformQuotaKey{keys},
+ getEntries: entries,
+ }
+ writer := &mockQuotaSnapshotWriter{}
+ f := newTestFlusher(cache, writer)
+
+ f.flush()
+
+ // 应收到 5 条 snapshot
+ if len(writer.receivedSnaps) != 5 {
+ t.Fatalf("90%% company: expected 5 snaps, got %d", len(writer.receivedSnaps))
+ }
+
+ // 验证绝对值语义:第 1 条 snap 的各窗口 usage 应等于 entries[0] 的值
+ snap0 := writer.receivedSnaps[0]
+ entry0 := entries[0]
+ if snap0.DailyUsageUSD != entry0.DailyUsageUSD {
+ t.Errorf("snap[0].DailyUsageUSD = %v, want %v", snap0.DailyUsageUSD, entry0.DailyUsageUSD)
+ }
+ if snap0.WeeklyUsageUSD != entry0.WeeklyUsageUSD {
+ t.Errorf("snap[0].WeeklyUsageUSD = %v, want %v", snap0.WeeklyUsageUSD, entry0.WeeklyUsageUSD)
+ }
+ if snap0.MonthlyUsageUSD != entry0.MonthlyUsageUSD {
+ t.Errorf("snap[0].MonthlyUsageUSD = %v, want %v", snap0.MonthlyUsageUSD, entry0.MonthlyUsageUSD)
+ }
+
+ // FlushBatchSizeTotal 应为 5(本批 keys 数量)
+ if f.metrics.FlushBatchSizeTotal.Load() != 5 {
+ t.Errorf("90%% company: FlushBatchSizeTotal = %d, want 5", f.metrics.FlushBatchSizeTotal.Load())
+ }
+ // FlushSuccessTotal 应为 1(1 个批次写成功)
+ if f.metrics.FlushSuccessTotal.Load() != 1 {
+ t.Errorf("90%% company: FlushSuccessTotal = %d, want 1", f.metrics.FlushSuccessTotal.Load())
+ }
+ // 无错误、无 Readd
+ if f.metrics.FlushErrorTotal.Load() != 0 {
+ t.Errorf("90%% company: FlushErrorTotal = %d, want 0", f.metrics.FlushErrorTotal.Load())
+ }
+ if f.metrics.DirtyReaddTotal.Load() != 0 {
+ t.Errorf("90%% company: DirtyReaddTotal = %d, want 0", f.metrics.DirtyReaddTotal.Load())
+ }
+ if len(cache.readdCalled) != 0 {
+ t.Errorf("90%% company: Readd should not be called, got %d calls", len(cache.readdCalled))
+ }
+}
diff --git a/backend/internal/service/user_platform_quota_port.go b/backend/internal/service/user_platform_quota_port.go
index cb09542a..0f88eda4 100644
--- a/backend/internal/service/user_platform_quota_port.go
+++ b/backend/internal/service/user_platform_quota_port.go
@@ -11,6 +11,23 @@ import (
// handler 只需引用 service 包,无需直接依赖 repository 包。
var ErrUserPlatformQuotaNotFound = errors.New("user platform quota not found")
+// ErrUserPlatformQuotaFKViolation service 层 sentinel:批量 snapshot UPSERT 时存在
+// user_id 不在 users 表的记录(外键违反)。adapter 负责将 repository 层同名 sentinel 包装为此错误。
+var ErrUserPlatformQuotaFKViolation = errors.New("user platform quota snapshot FK violation")
+
+// UserPlatformQuotaSnapshot 是 service 层 flusher 向 DB 写入快照时使用的传输结构。
+// 字段语义与 repository.UserPlatformQuotaSnapshot 完全对应,由 adapter 负责转换。
+type UserPlatformQuotaSnapshot struct {
+ UserID int64
+ Platform string
+ DailyUsageUSD float64
+ WeeklyUsageUSD float64
+ MonthlyUsageUSD float64
+ DailyWindowStart time.Time
+ WeeklyWindowStart time.Time
+ MonthlyWindowStart time.Time
+}
+
// UserPlatformQuotaRecord service 层传输结构体(与 repository 层解耦)。
type UserPlatformQuotaRecord struct {
UserID int64
@@ -47,4 +64,6 @@ type UserPlatformQuotaRepository interface {
// ResetExpiredWindow 重置指定窗口("daily"|"weekly"|"monthly")的用量与起始时间。
// 未命中活跃记录时返回(service-side wrapper of repository.ErrUserPlatformQuotaNotFound)。
ResetExpiredWindow(ctx context.Context, userID int64, platform string, window string, newStart time.Time) error
+ // BatchSnapshotUsage 绝对值覆盖写入整批 usage 快照。FK 违反返回 ErrUserPlatformQuotaFKViolation。
+ BatchSnapshotUsage(ctx context.Context, snapshots []UserPlatformQuotaSnapshot, now time.Time) error
}
diff --git a/backend/internal/service/user_service.go b/backend/internal/service/user_service.go
index 36bcf1c8..f801e2c4 100644
--- a/backend/internal/service/user_service.go
+++ b/backend/internal/service/user_service.go
@@ -74,11 +74,16 @@ type UserListFilters struct {
// For large datasets this can be expensive; admin list pages should enable it on demand.
// nil means not specified (default: load subscriptions for backward compatibility).
IncludeSubscriptions *bool
+ // IncludeDeleted 为 true 时绕过软删除过滤,返回含已删除(deleted_at 非空)的用户。
+ // 仅供 /admin/usage 的 SearchUsers 端点使用,其他列表调用方不要设置。
+ IncludeDeleted bool
}
type UserRepository interface {
Create(ctx context.Context, user *User) error
GetByID(ctx context.Context, id int64) (*User, error)
+ // GetByIDIncludeDeleted 绕过软删除过滤按 ID 取用户(含已删)。仅供管理员审计/usage 点击使用。
+ GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error)
GetByEmail(ctx context.Context, email string) (*User, error)
GetFirstAdmin(ctx context.Context) (*User, error)
Update(ctx context.Context, user *User) error
diff --git a/backend/internal/service/user_service_test.go b/backend/internal/service/user_service_test.go
index 19aec5d3..417140ad 100644
--- a/backend/internal/service/user_service_test.go
+++ b/backend/internal/service/user_service_test.go
@@ -236,6 +236,10 @@ func (m *mockUserRepo) UnbindUserAuthProvider(_ context.Context, _ int64, provid
return nil
}
+func (m *mockUserRepo) GetByIDIncludeDeleted(ctx context.Context, id int64) (*User, error) {
+ return m.GetByID(ctx, id)
+}
+
func (m *mockUserRepo) WithUserProfileIdentityTx(ctx context.Context, fn func(txCtx context.Context) error) error {
m.txCalls++
txState := &mockUserRepoTxState{
@@ -327,10 +331,22 @@ func (m *mockBillingCache) DeleteUserPlatformQuotaCache(context.Context, int64,
return nil
}
-func (m *mockBillingCache) IncrUserPlatformQuotaUsageCache(context.Context, int64, string, float64, time.Duration) error {
+func (m *mockBillingCache) IncrUserPlatformQuotaUsageCache(context.Context, int64, string, float64, time.Duration, bool) error {
return nil
}
+func (m *mockBillingCache) PopDirtyUserPlatformQuotaKeys(context.Context, int) ([]UserPlatformQuotaKey, error) {
+ return nil, nil
+}
+
+func (m *mockBillingCache) ReaddDirtyUserPlatformQuotaKeys(context.Context, []UserPlatformQuotaKey) error {
+ return nil
+}
+
+func (m *mockBillingCache) BatchGetUserPlatformQuotaCache(context.Context, []UserPlatformQuotaKey) ([]*UserPlatformQuotaCacheEntry, error) {
+ return nil, nil
+}
+
// --- 测试 ---
func TestUpdateBalance_Success(t *testing.T) {
diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go
index b22e10ae..fbee1b05 100644
--- a/backend/internal/service/wire.go
+++ b/backend/internal/service/wire.go
@@ -45,6 +45,17 @@ func ProvideOAuthRefreshAPI(accountRepo AccountRepository, tokenCache GeminiToke
return NewOAuthRefreshAPI(accountRepo, tokenCache)
}
+// ProvideOpenAIOAuthService creates OpenAIOAuthService with privacy/account enrichment support.
+func ProvideOpenAIOAuthService(
+ proxyRepo ProxyRepository,
+ oauthClient OpenAIOAuthClient,
+ privacyClientFactory PrivacyClientFactory,
+) *OpenAIOAuthService {
+ svc := NewOpenAIOAuthService(proxyRepo, oauthClient)
+ svc.SetPrivacyClientFactory(privacyClientFactory)
+ return svc
+}
+
// ProvideTokenRefreshService creates and starts TokenRefreshService
func ProvideTokenRefreshService(
accountRepo AccountRepository,
@@ -396,6 +407,46 @@ func ProvideBackupService(
return svc
}
+// ProvideOpsService constructs OpsService and wires the SettingService-backed quota
+// auto-pause cache sink. Mirrors the SetCleanupReloader pattern: OpsService doesn't
+// hold a *SettingService reference, but wire injects a tiny callback so writes to
+// ops_advanced_settings immediately propagate into the scheduler hot-path cache.
+func ProvideOpsService(
+ opsRepo OpsRepository,
+ settingRepo SettingRepository,
+ cfg *config.Config,
+ accountRepo AccountRepository,
+ userRepo UserRepository,
+ concurrencyService *ConcurrencyService,
+ gatewayService *GatewayService,
+ openAIGatewayService *OpenAIGatewayService,
+ geminiCompatService *GeminiMessagesCompatService,
+ antigravityGatewayService *AntigravityGatewayService,
+ systemLogSink *OpsSystemLogSink,
+ settingService *SettingService,
+) *OpsService {
+ svc := NewOpsService(
+ opsRepo,
+ settingRepo,
+ cfg,
+ accountRepo,
+ userRepo,
+ concurrencyService,
+ gatewayService,
+ openAIGatewayService,
+ geminiCompatService,
+ antigravityGatewayService,
+ systemLogSink,
+ )
+ if settingService != nil {
+ svc.SetOpenAIQuotaAutoPauseSettingsSink(settingService.SetOpenAIQuotaAutoPauseSettings)
+ // Optional warm-up so the first scheduled request after process start observes
+ // a populated cache rather than zero defaults. Best-effort, sync-bounded.
+ settingService.WarmOpenAIQuotaAutoPauseSettings(context.Background())
+ }
+ return svc
+}
+
// ProvideSettingService wires SettingService with group reader and proxy repo.
func ProvideSettingService(settingRepo SettingRepository, groupRepo GroupRepository, proxyRepo ProxyRepository, cfg *config.Config) *SettingService {
svc := NewSettingService(settingRepo, cfg)
@@ -461,7 +512,7 @@ var ProviderSet = wire.NewSet(
NewOpenAIGatewayService,
wire.Bind(new(AccountRuntimeBlocker), new(*OpenAIGatewayService)),
NewOAuthService,
- NewOpenAIOAuthService,
+ ProvideOpenAIOAuthService,
NewGeminiOAuthService,
NewGeminiQuotaService,
NewCompositeTokenCacheInvalidator,
@@ -481,7 +532,7 @@ var ProviderSet = wire.NewSet(
NewDataManagementService,
ProvideBackupService,
ProvideOpsSystemLogSink,
- NewOpsService,
+ ProvideOpsService,
ProvideOpsMetricsCollector,
ProvideOpsAggregationService,
ProvideOpsAlertEvaluatorService,
@@ -531,8 +582,16 @@ var ProviderSet = wire.NewSet(
ProvideChannelMonitorService,
ProvideChannelMonitorRunner,
NewChannelMonitorRequestTemplateService,
+ ProvideUserPlatformQuotaUsageFlusher,
)
+// ProvideUserPlatformQuotaUsageFlusher 创建并启动 UserPlatformQuotaUsageFlusher。
+func ProvideUserPlatformQuotaUsageFlusher(cfg *config.Config, cache BillingCache, quotaRepo UserPlatformQuotaRepository, tw *TimingWheelService) *UserPlatformQuotaUsageFlusher {
+ svc := NewUserPlatformQuotaUsageFlusher(cfg, cache, quotaRepo, tw)
+ svc.Start()
+ return svc
+}
+
// ProvidePaymentConfigService wraps NewPaymentConfigService to accept the named
// payment.EncryptionKey type instead of raw []byte, avoiding Wire ambiguity.
func ProvidePaymentConfigService(entClient *dbent.Client, settingRepo SettingRepository, key payment.EncryptionKey) *PaymentConfigService {
diff --git a/backend/migrations/143_group_models_list_config.sql b/backend/migrations/143_group_models_list_config.sql
new file mode 100644
index 00000000..67f27623
--- /dev/null
+++ b/backend/migrations/143_group_models_list_config.sql
@@ -0,0 +1,5 @@
+-- 分组级自定义 /v1/models 展示列表配置。
+-- 仅用于控制 GET /v1/models 的展示结果,不参与账号白名单、模型映射或网关调度。
+
+ALTER TABLE groups
+ ADD COLUMN IF NOT EXISTS models_list_config JSONB NOT NULL DEFAULT '{}'::jsonb;
diff --git a/backend/migrations/144_add_opus48_to_model_mapping.sql b/backend/migrations/144_add_opus48_to_model_mapping.sql
new file mode 100644
index 00000000..18fbff93
--- /dev/null
+++ b/backend/migrations/144_add_opus48_to_model_mapping.sql
@@ -0,0 +1,16 @@
+-- 为已持久化的 Antigravity model_mapping 添加 claude-opus-4-8。
+--
+-- 未持久化 model_mapping 的账号会直接使用 DefaultAntigravityModelMapping,
+-- 因此这里只需要回填已有映射对象。
+
+UPDATE accounts
+SET credentials = jsonb_set(
+ credentials,
+ '{model_mapping,claude-opus-4-8}',
+ '"claude-opus-4-8"'::jsonb,
+ true
+)
+WHERE platform = 'antigravity'
+ AND deleted_at IS NULL
+ AND jsonb_typeof(credentials->'model_mapping') = 'object'
+ AND credentials->'model_mapping'->>'claude-opus-4-8' IS NULL;
diff --git a/backend/migrations/145_deleted_api_key_audit.sql b/backend/migrations/145_deleted_api_key_audit.sql
new file mode 100644
index 00000000..1364c094
--- /dev/null
+++ b/backend/migrations/145_deleted_api_key_audit.sql
@@ -0,0 +1,22 @@
+-- 已删除 API key 审计表:删除 key 时同步留存(明文 key、所有者、key 信息),
+-- 供认证失败(INVALID_API_KEY)反查"这个失效 key 曾属于谁"。
+-- 仅对本表上线后删除的 key 生效;此前已删的 key 原值已被 tombstone 覆盖,无法补录。
+SET LOCAL lock_timeout = '5s';
+SET LOCAL statement_timeout = '10min';
+
+CREATE TABLE IF NOT EXISTS deleted_api_key_audits (
+ id BIGSERIAL PRIMARY KEY,
+ key VARCHAR(128) NOT NULL, -- 原 key 明文(复用 api_keys 策略),非唯一
+ api_key_id BIGINT NOT NULL, -- 原 api_keys.id
+ user_id BIGINT NOT NULL, -- 原所有者(不加外键,与 ops 表设计哲学一致)
+ key_name VARCHAR(100) NOT NULL DEFAULT '', -- 原 key 名称,便于展示
+ deleted_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
+ created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
+);
+CREATE INDEX IF NOT EXISTS deletedapikeyaudit_key ON deleted_api_key_audits (key);
+CREATE INDEX IF NOT EXISTS deletedapikeyaudit_user_id ON deleted_api_key_audits (user_id);
+
+ALTER TABLE ops_error_logs
+ ADD COLUMN IF NOT EXISTS attempted_key_prefix VARCHAR(32),
+ ADD COLUMN IF NOT EXISTS deleted_key_owner_user_id BIGINT,
+ ADD COLUMN IF NOT EXISTS deleted_key_name VARCHAR(100);
diff --git a/backend/migrations/147_ops_error_log_api_key_prefix.sql b/backend/migrations/147_ops_error_log_api_key_prefix.sql
new file mode 100644
index 00000000..bfc07489
--- /dev/null
+++ b/backend/migrations/147_ops_error_log_api_key_prefix.sql
@@ -0,0 +1,12 @@
+-- 有效(未删除)key 报错时,在 ops 落库层快照该 key 的脱敏前缀(前 8 位),
+-- 便于在 /admin/ops 错误详情识别是用户的哪一个 key 出的错。
+-- 与 attempted_key_prefix 互补且互斥:
+-- api_key_id 非空(有效 key 报错) => api_key_prefix
+-- api_key_id 为空(INVALID_API_KEY 无效) => attempted_key_prefix
+-- 落库快照(而非读时 JOIN api_keys):key 之后被删时 api_keys.key 会被 tombstone
+-- 覆盖,快照可保留报错当时的真实前缀。
+SET LOCAL lock_timeout = '5s';
+SET LOCAL statement_timeout = '10min';
+
+ALTER TABLE ops_error_logs
+ ADD COLUMN IF NOT EXISTS api_key_prefix VARCHAR(32);
diff --git a/backend/migrations/148_add_ops_error_logs_user_time_index_notx.sql b/backend/migrations/148_add_ops_error_logs_user_time_index_notx.sql
new file mode 100644
index 00000000..54d73ba5
--- /dev/null
+++ b/backend/migrations/148_add_ops_error_logs_user_time_index_notx.sql
@@ -0,0 +1,6 @@
+-- 148_add_ops_error_logs_user_time_index_notx.sql
+-- 用户侧"错误请求"按 user_id 时间倒序分页所需的部分索引。
+-- 非事务迁移(_notx):CREATE INDEX CONCURRENTLY 不可在事务内执行。
+CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_ops_error_logs_user_time
+ ON ops_error_logs (user_id, created_at DESC)
+ WHERE user_id IS NOT NULL;
diff --git a/backend/resources/model-pricing/model_prices_and_context_window.json b/backend/resources/model-pricing/model_prices_and_context_window.json
index 3cae8c8b..e88ed2da 100644
--- a/backend/resources/model-pricing/model_prices_and_context_window.json
+++ b/backend/resources/model-pricing/model_prices_and_context_window.json
@@ -1,135 +1,4 @@
{
- "claude-3-5-haiku-20241022": {
- "cache_creation_input_token_cost": 1e-06,
- "cache_creation_input_token_cost_above_1hr": 6e-06,
- "cache_read_input_token_cost": 8e-08,
- "deprecation_date": "2025-10-01",
- "input_cost_per_token": 8e-07,
- "litellm_provider": "anthropic",
- "max_input_tokens": 200000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_token": 4e-06,
- "search_context_cost_per_query": {
- "search_context_size_high": 0.01,
- "search_context_size_low": 0.01,
- "search_context_size_medium": 0.01
- },
- "supports_assistant_prefill": true,
- "supports_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true,
- "tool_use_system_prompt_tokens": 264
- },
- "claude-3-5-haiku-latest": {
- "cache_creation_input_token_cost": 1.25e-06,
- "cache_creation_input_token_cost_above_1hr": 6e-06,
- "cache_read_input_token_cost": 1e-07,
- "deprecation_date": "2025-10-01",
- "input_cost_per_token": 1e-06,
- "litellm_provider": "anthropic",
- "max_input_tokens": 200000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_token": 5e-06,
- "search_context_cost_per_query": {
- "search_context_size_high": 0.01,
- "search_context_size_low": 0.01,
- "search_context_size_medium": 0.01
- },
- "supports_assistant_prefill": true,
- "supports_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true,
- "tool_use_system_prompt_tokens": 264
- },
- "claude-3-5-sonnet-20240620": {
- "cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_1hr": 6e-06,
- "cache_read_input_token_cost": 3e-07,
- "deprecation_date": "2025-06-01",
- "input_cost_per_token": 3e-06,
- "litellm_provider": "anthropic",
- "max_input_tokens": 200000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_token": 1.5e-05,
- "supports_assistant_prefill": true,
- "supports_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "tool_use_system_prompt_tokens": 159
- },
- "claude-3-5-sonnet-20241022": {
- "cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_1hr": 6e-06,
- "cache_read_input_token_cost": 3e-07,
- "deprecation_date": "2025-10-01",
- "input_cost_per_token": 3e-06,
- "litellm_provider": "anthropic",
- "max_input_tokens": 200000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_token": 1.5e-05,
- "search_context_cost_per_query": {
- "search_context_size_high": 0.01,
- "search_context_size_low": 0.01,
- "search_context_size_medium": 0.01
- },
- "supports_assistant_prefill": true,
- "supports_computer_use": true,
- "supports_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true,
- "tool_use_system_prompt_tokens": 159
- },
- "claude-3-5-sonnet-latest": {
- "cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_1hr": 6e-06,
- "cache_read_input_token_cost": 3e-07,
- "deprecation_date": "2025-06-01",
- "input_cost_per_token": 3e-06,
- "litellm_provider": "anthropic",
- "max_input_tokens": 200000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_token": 1.5e-05,
- "search_context_cost_per_query": {
- "search_context_size_high": 0.01,
- "search_context_size_low": 0.01,
- "search_context_size_medium": 0.01
- },
- "supports_assistant_prefill": true,
- "supports_computer_use": true,
- "supports_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true,
- "tool_use_system_prompt_tokens": 159
- },
"claude-3-7-sonnet-20250219": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@@ -159,34 +28,6 @@
"supports_web_search": true,
"tool_use_system_prompt_tokens": 159
},
- "claude-3-7-sonnet-latest": {
- "cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_1hr": 6e-06,
- "cache_read_input_token_cost": 3e-07,
- "deprecation_date": "2025-06-01",
- "input_cost_per_token": 3e-06,
- "litellm_provider": "anthropic",
- "max_input_tokens": 200000,
- "max_output_tokens": 64000,
- "max_tokens": 64000,
- "mode": "chat",
- "output_cost_per_token": 1.5e-05,
- "search_context_cost_per_query": {
- "search_context_size_high": 0.01,
- "search_context_size_low": 0.01,
- "search_context_size_medium": 0.01
- },
- "supports_assistant_prefill": true,
- "supports_computer_use": true,
- "supports_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "tool_use_system_prompt_tokens": 159
- },
"claude-3-haiku-20240307": {
"cache_creation_input_token_cost": 3e-07,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@@ -226,28 +67,9 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 395
},
- "claude-3-opus-latest": {
- "cache_creation_input_token_cost": 1.875e-05,
- "cache_creation_input_token_cost_above_1hr": 6e-06,
- "cache_read_input_token_cost": 1.5e-06,
- "deprecation_date": "2025-03-01",
- "input_cost_per_token": 1.5e-05,
- "litellm_provider": "anthropic",
- "max_input_tokens": 200000,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_token": 7.5e-05,
- "supports_assistant_prefill": true,
- "supports_function_calling": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "tool_use_system_prompt_tokens": 395
- },
"claude-4-opus-20250514": {
"cache_creation_input_token_cost": 1.875e-05,
+ "cache_creation_input_token_cost_above_1hr": 3e-05,
"cache_read_input_token_cost": 1.5e-06,
"input_cost_per_token": 1.5e-05,
"litellm_provider": "anthropic",
@@ -274,6 +96,7 @@
},
"claude-4-sonnet-20250514": {
"cache_creation_input_token_cost": 3.75e-06,
+ "cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
@@ -447,6 +270,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@@ -474,6 +298,7 @@
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@@ -485,18 +310,14 @@
"claude-opus-4-6": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "anthropic",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"provider_specific_entry": {
"fast": 6.0,
"us": 1.1
@@ -506,9 +327,13 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
+ "supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
+ "supports_max_reasoning_effort": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_output_config": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@@ -520,18 +345,14 @@
"claude-opus-4-6-20260205": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "anthropic",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"provider_specific_entry": {
"fast": 6.0,
"us": 1.1
@@ -541,9 +362,13 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
+ "supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
+ "supports_max_reasoning_effort": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_output_config": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@@ -555,18 +380,14 @@
"claude-opus-4-6-thinking": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
- "cache_creation_input_token_cost_above_200k_tokens": 1.25e-05,
"cache_read_input_token_cost": 5e-07,
- "cache_read_input_token_cost_above_200k_tokens": 1e-06,
"input_cost_per_token": 5e-06,
- "input_cost_per_token_above_200k_tokens": 1e-05,
"litellm_provider": "anthropic",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
- "output_cost_per_token_above_200k_tokens": 3.75e-05,
"provider_specific_entry": {
"fast": 6.0,
"us": 1.1
@@ -576,9 +397,13 @@
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
+ "supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
+ "supports_max_reasoning_effort": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_output_config": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@@ -587,6 +412,114 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
+ "claude-opus-4-7": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_1hr": 1e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "input_cost_per_token": 5e-06,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "provider_specific_entry": {
+ "fast": 6.0,
+ "us": 1.1
+ },
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_adaptive_thinking": true,
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_max_reasoning_effort": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_output_config": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_xhigh_reasoning_effort": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "claude-opus-4-7-20260416": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_1hr": 1e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "input_cost_per_token": 5e-06,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "provider_specific_entry": {
+ "fast": 6.0,
+ "us": 1.1
+ },
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_adaptive_thinking": true,
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_max_reasoning_effort": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_output_config": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_xhigh_reasoning_effort": true,
+ "tool_use_system_prompt_tokens": 346
+ },
+ "claude-opus-4-8": {
+ "cache_creation_input_token_cost": 6.25e-06,
+ "cache_creation_input_token_cost_above_1hr": 1e-05,
+ "cache_read_input_token_cost": 5e-07,
+ "input_cost_per_token": 5e-06,
+ "litellm_provider": "anthropic",
+ "max_input_tokens": 1000000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 2.5e-05,
+ "provider_specific_entry": {
+ "fast": 6.0,
+ "us": 1.1
+ },
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.01,
+ "search_context_size_low": 0.01,
+ "search_context_size_medium": 0.01
+ },
+ "supports_adaptive_thinking": true,
+ "supports_assistant_prefill": false,
+ "supports_computer_use": true,
+ "supports_function_calling": true,
+ "supports_max_reasoning_effort": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_output_config": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_xhigh_reasoning_effort": true,
+ "tool_use_system_prompt_tokens": 346
+ },
"claude-sonnet-4-20250514": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@@ -621,6 +554,7 @@
},
"claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
+ "cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
@@ -651,6 +585,7 @@
},
"claude-sonnet-4-5-20250929": {
"cache_creation_input_token_cost": 3.75e-06,
+ "cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
@@ -682,6 +617,7 @@
},
"claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 3.75e-06,
+ "cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost": 3e-07,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
@@ -707,26 +643,27 @@
},
"claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
- "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
+ "cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
- "cache_read_input_token_cost_above_200k_tokens": 6e-07,
"input_cost_per_token": 3e-06,
- "input_cost_per_token_above_200k_tokens": 6e-06,
"litellm_provider": "anthropic",
- "max_input_tokens": 200000,
+ "max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
- "output_cost_per_token_above_200k_tokens": 2.25e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
+ "supports_adaptive_thinking": true,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
+ "supports_max_reasoning_effort": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_output_config": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
@@ -735,6 +672,54 @@
"supports_vision": true,
"tool_use_system_prompt_tokens": 346
},
+ "codex-auto-review": {
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_272k_tokens": 1e-06,
+ "cache_read_input_token_cost_flex": 2.5e-07,
+ "cache_read_input_token_cost_priority": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_272k_tokens": 1e-05,
+ "input_cost_per_token_batches": 2.5e-06,
+ "input_cost_per_token_flex": 2.5e-06,
+ "input_cost_per_token_priority": 1e-05,
+ "litellm_provider": "openai",
+ "max_input_tokens": 1050000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 3e-05,
+ "output_cost_per_token_above_272k_tokens": 4.5e-05,
+ "output_cost_per_token_batches": 1.5e-05,
+ "output_cost_per_token_flex": 1.5e-05,
+ "output_cost_per_token_priority": 6e-05,
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/batch",
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": false,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
"deepseek-chat": {
"cache_read_input_token_cost": 2.8e-08,
"input_cost_per_token": 2.8e-07,
@@ -792,478 +777,9 @@
"supports_reasoning": true,
"supports_tool_choice": true
},
- "gemini-1.0-pro": {
- "input_cost_per_character": 1.25e-07,
- "input_cost_per_image": 0.0025,
- "input_cost_per_token": 5e-07,
- "input_cost_per_video_per_second": 0.002,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 32760,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 3.75e-07,
- "output_cost_per_token": 1.5e-06,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#google_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true
- },
- "gemini-1.0-pro-001": {
- "deprecation_date": "2025-04-09",
- "input_cost_per_character": 1.25e-07,
- "input_cost_per_image": 0.0025,
- "input_cost_per_token": 5e-07,
- "input_cost_per_video_per_second": 0.002,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 32760,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 3.75e-07,
- "output_cost_per_token": 1.5e-06,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true
- },
- "gemini-1.0-pro-002": {
- "deprecation_date": "2025-04-09",
- "input_cost_per_character": 1.25e-07,
- "input_cost_per_image": 0.0025,
- "input_cost_per_token": 5e-07,
- "input_cost_per_video_per_second": 0.002,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 32760,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 3.75e-07,
- "output_cost_per_token": 1.5e-06,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true
- },
- "gemini-1.0-pro-vision": {
- "input_cost_per_image": 0.0025,
- "input_cost_per_token": 5e-07,
- "litellm_provider": "vertex_ai-vision-models",
- "max_images_per_prompt": 16,
- "max_input_tokens": 16384,
- "max_output_tokens": 2048,
- "max_tokens": 2048,
- "max_video_length": 2,
- "max_videos_per_prompt": 1,
- "mode": "chat",
- "output_cost_per_token": 1.5e-06,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.0-pro-vision-001": {
- "deprecation_date": "2025-04-09",
- "input_cost_per_image": 0.0025,
- "input_cost_per_token": 5e-07,
- "litellm_provider": "vertex_ai-vision-models",
- "max_images_per_prompt": 16,
- "max_input_tokens": 16384,
- "max_output_tokens": 2048,
- "max_tokens": 2048,
- "max_video_length": 2,
- "max_videos_per_prompt": 1,
- "mode": "chat",
- "output_cost_per_token": 1.5e-06,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.0-ultra": {
- "input_cost_per_character": 1.25e-07,
- "input_cost_per_image": 0.0025,
- "input_cost_per_token": 5e-07,
- "input_cost_per_video_per_second": 0.002,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 8192,
- "max_output_tokens": 2048,
- "max_tokens": 2048,
- "mode": "chat",
- "output_cost_per_character": 3.75e-07,
- "output_cost_per_token": 1.5e-06,
- "source": "As of Jun, 2024. There is no available doc on vertex ai pricing gemini-1.0-ultra-001. Using gemini-1.0-pro pricing. Got max_tokens info here: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true
- },
- "gemini-1.0-ultra-001": {
- "input_cost_per_character": 1.25e-07,
- "input_cost_per_image": 0.0025,
- "input_cost_per_token": 5e-07,
- "input_cost_per_video_per_second": 0.002,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 8192,
- "max_output_tokens": 2048,
- "max_tokens": 2048,
- "mode": "chat",
- "output_cost_per_character": 3.75e-07,
- "output_cost_per_token": 1.5e-06,
- "source": "As of Jun, 2024. There is no available doc on vertex ai pricing gemini-1.0-ultra-001. Using gemini-1.0-pro pricing. Got max_tokens info here: https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true
- },
- "gemini-1.5-flash": {
- "deprecation_date": "2025-09-29",
- "input_cost_per_audio_per_second": 2e-06,
- "input_cost_per_audio_per_second_above_128k_tokens": 4e-06,
- "input_cost_per_character": 1.875e-08,
- "input_cost_per_character_above_128k_tokens": 2.5e-07,
- "input_cost_per_image": 2e-05,
- "input_cost_per_image_above_128k_tokens": 4e-05,
- "input_cost_per_token": 7.5e-08,
- "input_cost_per_token_above_128k_tokens": 1e-06,
- "input_cost_per_video_per_second": 2e-05,
- "input_cost_per_video_per_second_above_128k_tokens": 4e-05,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1000000,
- "max_output_tokens": 8192,
- "max_pdf_size_mb": 30,
- "max_tokens": 8192,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_character": 7.5e-08,
- "output_cost_per_character_above_128k_tokens": 1.5e-07,
- "output_cost_per_token": 3e-07,
- "output_cost_per_token_above_128k_tokens": 6e-07,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.5-flash-001": {
- "deprecation_date": "2025-05-24",
- "input_cost_per_audio_per_second": 2e-06,
- "input_cost_per_audio_per_second_above_128k_tokens": 4e-06,
- "input_cost_per_character": 1.875e-08,
- "input_cost_per_character_above_128k_tokens": 2.5e-07,
- "input_cost_per_image": 2e-05,
- "input_cost_per_image_above_128k_tokens": 4e-05,
- "input_cost_per_token": 7.5e-08,
- "input_cost_per_token_above_128k_tokens": 1e-06,
- "input_cost_per_video_per_second": 2e-05,
- "input_cost_per_video_per_second_above_128k_tokens": 4e-05,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1000000,
- "max_output_tokens": 8192,
- "max_pdf_size_mb": 30,
- "max_tokens": 8192,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_character": 7.5e-08,
- "output_cost_per_character_above_128k_tokens": 1.5e-07,
- "output_cost_per_token": 3e-07,
- "output_cost_per_token_above_128k_tokens": 6e-07,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.5-flash-002": {
- "deprecation_date": "2025-09-24",
- "input_cost_per_audio_per_second": 2e-06,
- "input_cost_per_audio_per_second_above_128k_tokens": 4e-06,
- "input_cost_per_character": 1.875e-08,
- "input_cost_per_character_above_128k_tokens": 2.5e-07,
- "input_cost_per_image": 2e-05,
- "input_cost_per_image_above_128k_tokens": 4e-05,
- "input_cost_per_token": 7.5e-08,
- "input_cost_per_token_above_128k_tokens": 1e-06,
- "input_cost_per_video_per_second": 2e-05,
- "input_cost_per_video_per_second_above_128k_tokens": 4e-05,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 8192,
- "max_pdf_size_mb": 30,
- "max_tokens": 8192,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_character": 7.5e-08,
- "output_cost_per_character_above_128k_tokens": 1.5e-07,
- "output_cost_per_token": 3e-07,
- "output_cost_per_token_above_128k_tokens": 6e-07,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-1.5-flash",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.5-flash-exp-0827": {
- "deprecation_date": "2025-09-29",
- "input_cost_per_audio_per_second": 2e-06,
- "input_cost_per_audio_per_second_above_128k_tokens": 4e-06,
- "input_cost_per_character": 1.875e-08,
- "input_cost_per_character_above_128k_tokens": 2.5e-07,
- "input_cost_per_image": 2e-05,
- "input_cost_per_image_above_128k_tokens": 4e-05,
- "input_cost_per_token": 4.688e-09,
- "input_cost_per_token_above_128k_tokens": 1e-06,
- "input_cost_per_video_per_second": 2e-05,
- "input_cost_per_video_per_second_above_128k_tokens": 4e-05,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1000000,
- "max_output_tokens": 8192,
- "max_pdf_size_mb": 30,
- "max_tokens": 8192,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_character": 1.875e-08,
- "output_cost_per_character_above_128k_tokens": 3.75e-08,
- "output_cost_per_token": 4.6875e-09,
- "output_cost_per_token_above_128k_tokens": 9.375e-09,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.5-flash-preview-0514": {
- "deprecation_date": "2025-09-29",
- "input_cost_per_audio_per_second": 2e-06,
- "input_cost_per_audio_per_second_above_128k_tokens": 4e-06,
- "input_cost_per_character": 1.875e-08,
- "input_cost_per_character_above_128k_tokens": 2.5e-07,
- "input_cost_per_image": 2e-05,
- "input_cost_per_image_above_128k_tokens": 4e-05,
- "input_cost_per_token": 7.5e-08,
- "input_cost_per_token_above_128k_tokens": 1e-06,
- "input_cost_per_video_per_second": 2e-05,
- "input_cost_per_video_per_second_above_128k_tokens": 4e-05,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1000000,
- "max_output_tokens": 8192,
- "max_pdf_size_mb": 30,
- "max_tokens": 8192,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_character": 1.875e-08,
- "output_cost_per_character_above_128k_tokens": 3.75e-08,
- "output_cost_per_token": 4.6875e-09,
- "output_cost_per_token_above_128k_tokens": 9.375e-09,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.5-pro": {
- "deprecation_date": "2025-09-29",
- "input_cost_per_audio_per_second": 3.125e-05,
- "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05,
- "input_cost_per_character": 3.125e-07,
- "input_cost_per_character_above_128k_tokens": 6.25e-07,
- "input_cost_per_image": 0.00032875,
- "input_cost_per_image_above_128k_tokens": 0.0006575,
- "input_cost_per_token": 1.25e-06,
- "input_cost_per_token_above_128k_tokens": 2.5e-06,
- "input_cost_per_video_per_second": 0.00032875,
- "input_cost_per_video_per_second_above_128k_tokens": 0.0006575,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 2097152,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 1.25e-06,
- "output_cost_per_character_above_128k_tokens": 2.5e-06,
- "output_cost_per_token": 5e-06,
- "output_cost_per_token_above_128k_tokens": 1e-05,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.5-pro-001": {
- "deprecation_date": "2025-05-24",
- "input_cost_per_audio_per_second": 3.125e-05,
- "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05,
- "input_cost_per_character": 3.125e-07,
- "input_cost_per_character_above_128k_tokens": 6.25e-07,
- "input_cost_per_image": 0.00032875,
- "input_cost_per_image_above_128k_tokens": 0.0006575,
- "input_cost_per_token": 1.25e-06,
- "input_cost_per_token_above_128k_tokens": 2.5e-06,
- "input_cost_per_video_per_second": 0.00032875,
- "input_cost_per_video_per_second_above_128k_tokens": 0.0006575,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 1000000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 1.25e-06,
- "output_cost_per_character_above_128k_tokens": 2.5e-06,
- "output_cost_per_token": 5e-06,
- "output_cost_per_token_above_128k_tokens": 1e-05,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.5-pro-002": {
- "deprecation_date": "2025-09-24",
- "input_cost_per_audio_per_second": 3.125e-05,
- "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05,
- "input_cost_per_character": 3.125e-07,
- "input_cost_per_character_above_128k_tokens": 6.25e-07,
- "input_cost_per_image": 0.00032875,
- "input_cost_per_image_above_128k_tokens": 0.0006575,
- "input_cost_per_token": 1.25e-06,
- "input_cost_per_token_above_128k_tokens": 2.5e-06,
- "input_cost_per_video_per_second": 0.00032875,
- "input_cost_per_video_per_second_above_128k_tokens": 0.0006575,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 2097152,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 1.25e-06,
- "output_cost_per_character_above_128k_tokens": 2.5e-06,
- "output_cost_per_token": 5e-06,
- "output_cost_per_token_above_128k_tokens": 1e-05,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-1.5-pro",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gemini-1.5-pro-preview-0215": {
- "deprecation_date": "2025-09-29",
- "input_cost_per_audio_per_second": 3.125e-05,
- "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05,
- "input_cost_per_character": 3.125e-07,
- "input_cost_per_character_above_128k_tokens": 6.25e-07,
- "input_cost_per_image": 0.00032875,
- "input_cost_per_image_above_128k_tokens": 0.0006575,
- "input_cost_per_token": 7.8125e-08,
- "input_cost_per_token_above_128k_tokens": 1.5625e-07,
- "input_cost_per_video_per_second": 0.00032875,
- "input_cost_per_video_per_second_above_128k_tokens": 0.0006575,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 1000000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 1.25e-06,
- "output_cost_per_character_above_128k_tokens": 2.5e-06,
- "output_cost_per_token": 3.125e-07,
- "output_cost_per_token_above_128k_tokens": 6.25e-07,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
- "gemini-1.5-pro-preview-0409": {
- "deprecation_date": "2025-09-29",
- "input_cost_per_audio_per_second": 3.125e-05,
- "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05,
- "input_cost_per_character": 3.125e-07,
- "input_cost_per_character_above_128k_tokens": 6.25e-07,
- "input_cost_per_image": 0.00032875,
- "input_cost_per_image_above_128k_tokens": 0.0006575,
- "input_cost_per_token": 7.8125e-08,
- "input_cost_per_token_above_128k_tokens": 1.5625e-07,
- "input_cost_per_video_per_second": 0.00032875,
- "input_cost_per_video_per_second_above_128k_tokens": 0.0006575,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 1000000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 1.25e-06,
- "output_cost_per_character_above_128k_tokens": 2.5e-06,
- "output_cost_per_token": 3.125e-07,
- "output_cost_per_token_above_128k_tokens": 6.25e-07,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_tool_choice": true
- },
- "gemini-1.5-pro-preview-0514": {
- "deprecation_date": "2025-09-29",
- "input_cost_per_audio_per_second": 3.125e-05,
- "input_cost_per_audio_per_second_above_128k_tokens": 6.25e-05,
- "input_cost_per_character": 3.125e-07,
- "input_cost_per_character_above_128k_tokens": 6.25e-07,
- "input_cost_per_image": 0.00032875,
- "input_cost_per_image_above_128k_tokens": 0.0006575,
- "input_cost_per_token": 7.8125e-08,
- "input_cost_per_token_above_128k_tokens": 1.5625e-07,
- "input_cost_per_video_per_second": 0.00032875,
- "input_cost_per_video_per_second_above_128k_tokens": 0.0006575,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 1000000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 1.25e-06,
- "output_cost_per_character_above_128k_tokens": 2.5e-06,
- "output_cost_per_token": 3.125e-07,
- "output_cost_per_token_above_128k_tokens": 6.25e-07,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
"gemini-2.0-flash": {
"cache_read_input_token_cost": 2.5e-08,
- "deprecation_date": "2026-03-31",
+ "deprecation_date": "2026-06-01",
"input_cost_per_audio_token": 7e-07,
"input_cost_per_token": 1e-07,
"litellm_provider": "vertex_ai-language-models",
@@ -1278,6 +794,11 @@
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_token": 4e-07,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://ai.google.dev/pricing#2_0flash",
"supported_modalities": [
"text",
@@ -1303,7 +824,7 @@
},
"gemini-2.0-flash-001": {
"cache_read_input_token_cost": 3.75e-08,
- "deprecation_date": "2026-03-31",
+ "deprecation_date": "2026-06-01",
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 1.5e-07,
"litellm_provider": "vertex_ai-language-models",
@@ -1318,54 +839,11 @@
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_token": 6e-07,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text",
- "image"
- ],
- "supports_audio_output": true,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true
- },
- "gemini-2.0-flash-exp": {
- "cache_read_input_token_cost": 3.75e-08,
- "input_cost_per_audio_per_second": 0,
- "input_cost_per_audio_per_second_above_128k_tokens": 0,
- "input_cost_per_character": 0,
- "input_cost_per_character_above_128k_tokens": 0,
- "input_cost_per_image": 0,
- "input_cost_per_image_above_128k_tokens": 0,
- "input_cost_per_token": 1.5e-07,
- "input_cost_per_token_above_128k_tokens": 0,
- "input_cost_per_video_per_second": 0,
- "input_cost_per_video_per_second_above_128k_tokens": 0,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 8192,
- "max_pdf_size_mb": 30,
- "max_tokens": 8192,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_character": 0,
- "output_cost_per_character_above_128k_tokens": 0,
- "output_cost_per_token": 6e-07,
- "output_cost_per_token_above_128k_tokens": 0,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_modalities": [
"text",
@@ -1410,7 +888,7 @@
},
"gemini-2.0-flash-lite": {
"cache_read_input_token_cost": 1.875e-08,
- "deprecation_date": "2026-03-31",
+ "deprecation_date": "2026-06-01",
"input_cost_per_audio_token": 7.5e-08,
"input_cost_per_token": 7.5e-08,
"litellm_provider": "vertex_ai-language-models",
@@ -1424,6 +902,11 @@
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_token": 3e-07,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash",
"supported_modalities": [
"text",
@@ -1446,7 +929,7 @@
},
"gemini-2.0-flash-lite-001": {
"cache_read_input_token_cost": 1.875e-08,
- "deprecation_date": "2026-03-31",
+ "deprecation_date": "2026-06-01",
"input_cost_per_audio_token": 7.5e-08,
"input_cost_per_token": 7.5e-08,
"litellm_provider": "vertex_ai-language-models",
@@ -1460,6 +943,11 @@
"max_videos_per_prompt": 10,
"mode": "chat",
"output_cost_per_token": 3e-07,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash",
"supported_modalities": [
"text",
@@ -1480,235 +968,6 @@
"supports_vision": true,
"supports_web_search": true
},
- "gemini-2.0-flash-live-preview-04-09": {
- "cache_read_input_token_cost": 7.5e-08,
- "input_cost_per_audio_token": 3e-06,
- "input_cost_per_image": 3e-06,
- "input_cost_per_token": 5e-07,
- "input_cost_per_video_per_second": 3e-06,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 65535,
- "max_pdf_size_mb": 30,
- "max_tokens": 65535,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_audio_token": 1.2e-05,
- "output_cost_per_token": 2e-06,
- "rpm": 10,
- "source": "https://cloud.google.com/vertex-ai/docs/generative-ai/model-reference/gemini#gemini-2-0-flash-live-preview-04-09",
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions"
- ],
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text",
- "audio"
- ],
- "supports_audio_output": true,
- "supports_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_url_context": true,
- "supports_vision": true,
- "supports_web_search": true,
- "tpm": 250000
- },
- "gemini-2.0-flash-preview-image-generation": {
- "cache_read_input_token_cost": 2.5e-08,
- "deprecation_date": "2025-11-14",
- "input_cost_per_audio_token": 7e-07,
- "input_cost_per_token": 1e-07,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 8192,
- "max_pdf_size_mb": 30,
- "max_tokens": 8192,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_token": 4e-07,
- "source": "https://ai.google.dev/pricing#2_0flash",
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text",
- "image"
- ],
- "supports_audio_input": true,
- "supports_audio_output": true,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true
- },
- "gemini-2.0-flash-thinking-exp": {
- "cache_read_input_token_cost": 0.0,
- "deprecation_date": "2025-12-02",
- "input_cost_per_audio_per_second": 0,
- "input_cost_per_audio_per_second_above_128k_tokens": 0,
- "input_cost_per_character": 0,
- "input_cost_per_character_above_128k_tokens": 0,
- "input_cost_per_image": 0,
- "input_cost_per_image_above_128k_tokens": 0,
- "input_cost_per_token": 0,
- "input_cost_per_token_above_128k_tokens": 0,
- "input_cost_per_video_per_second": 0,
- "input_cost_per_video_per_second_above_128k_tokens": 0,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 8192,
- "max_pdf_size_mb": 30,
- "max_tokens": 8192,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_character": 0,
- "output_cost_per_character_above_128k_tokens": 0,
- "output_cost_per_token": 0,
- "output_cost_per_token_above_128k_tokens": 0,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash",
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text",
- "image"
- ],
- "supports_audio_output": true,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true
- },
- "gemini-2.0-flash-thinking-exp-01-21": {
- "cache_read_input_token_cost": 0.0,
- "deprecation_date": "2025-12-02",
- "input_cost_per_audio_per_second": 0,
- "input_cost_per_audio_per_second_above_128k_tokens": 0,
- "input_cost_per_character": 0,
- "input_cost_per_character_above_128k_tokens": 0,
- "input_cost_per_image": 0,
- "input_cost_per_image_above_128k_tokens": 0,
- "input_cost_per_token": 0,
- "input_cost_per_token_above_128k_tokens": 0,
- "input_cost_per_video_per_second": 0,
- "input_cost_per_video_per_second_above_128k_tokens": 0,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 65536,
- "max_pdf_size_mb": 30,
- "max_tokens": 65536,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_character": 0,
- "output_cost_per_character_above_128k_tokens": 0,
- "output_cost_per_token": 0,
- "output_cost_per_token_above_128k_tokens": 0,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#gemini-2.0-flash",
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text",
- "image"
- ],
- "supports_audio_output": false,
- "supports_function_calling": false,
- "supports_parallel_function_calling": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": false,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true
- },
- "gemini-2.0-pro-exp-02-05": {
- "cache_read_input_token_cost": 3.125e-07,
- "input_cost_per_token": 1.25e-06,
- "input_cost_per_token_above_200k_tokens": 2.5e-06,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 2097152,
- "max_output_tokens": 8192,
- "max_pdf_size_mb": 30,
- "max_tokens": 8192,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_token": 1e-05,
- "output_cost_per_token_above_200k_tokens": 1.5e-05,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions"
- ],
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text"
- ],
- "supports_audio_input": true,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_video_input": true,
- "supports_vision": true,
- "supports_web_search": true
- },
"gemini-2.5-computer-use-preview-10-2025": {
"input_cost_per_token": 1.25e-06,
"input_cost_per_token_above_200k_tokens": 2.5e-06,
@@ -1751,6 +1010,11 @@
"mode": "chat",
"output_cost_per_reasoning_token": 2.5e-06,
"output_cost_per_token": 2.5e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
"supported_endpoints": [
"/v1/chat/completions",
@@ -1773,6 +1037,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
+ "supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_url_context": true,
@@ -1821,6 +1086,7 @@
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
+ "supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_url_context": true,
@@ -1828,57 +1094,6 @@
"supports_web_search": false,
"tpm": 8000000
},
- "gemini-2.5-flash-image-preview": {
- "cache_read_input_token_cost": 7.5e-08,
- "deprecation_date": "2026-01-15",
- "input_cost_per_audio_token": 1e-06,
- "input_cost_per_image_token": 3e-07,
- "input_cost_per_token": 3e-07,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 65535,
- "max_pdf_size_mb": 30,
- "max_tokens": 65535,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "image_generation",
- "output_cost_per_image": 0.039,
- "output_cost_per_image_token": 3e-05,
- "output_cost_per_reasoning_token": 3e-05,
- "output_cost_per_token": 3e-05,
- "rpm": 100000,
- "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions",
- "/v1/batch"
- ],
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text",
- "image"
- ],
- "supports_audio_output": false,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_url_context": true,
- "supports_vision": true,
- "supports_web_search": true,
- "tpm": 8000000
- },
"gemini-2.5-flash-lite": {
"cache_read_input_token_cost": 1e-08,
"input_cost_per_audio_token": 3e-07,
@@ -1896,6 +1111,11 @@
"mode": "chat",
"output_cost_per_reasoning_token": 4e-07,
"output_cost_per_token": 4e-07,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
"supported_endpoints": [
"/v1/chat/completions",
@@ -1918,6 +1138,7 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
+ "supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_url_context": true,
@@ -1942,6 +1163,11 @@
"mode": "chat",
"output_cost_per_reasoning_token": 4e-07,
"output_cost_per_token": 4e-07,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
"supported_endpoints": [
"/v1/chat/completions",
@@ -1987,6 +1213,11 @@
"mode": "chat",
"output_cost_per_reasoning_token": 4e-07,
"output_cost_per_token": 4e-07,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2087,96 +1318,6 @@
"supports_audio_input": true,
"supports_audio_output": true
},
- "gemini-2.5-flash-preview-04-17": {
- "cache_read_input_token_cost": 3.75e-08,
- "input_cost_per_audio_token": 1e-06,
- "input_cost_per_token": 1.5e-07,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 65535,
- "max_pdf_size_mb": 30,
- "max_tokens": 65535,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_reasoning_token": 3.5e-06,
- "output_cost_per_token": 6e-07,
- "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions",
- "/v1/batch"
- ],
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text"
- ],
- "supports_audio_output": false,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true
- },
- "gemini-2.5-flash-preview-05-20": {
- "cache_read_input_token_cost": 7.5e-08,
- "deprecation_date": "2025-11-18",
- "input_cost_per_audio_token": 1e-06,
- "input_cost_per_token": 3e-07,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 65535,
- "max_pdf_size_mb": 30,
- "max_tokens": 65535,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_reasoning_token": 2.5e-06,
- "output_cost_per_token": 2.5e-06,
- "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions",
- "/v1/batch"
- ],
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text"
- ],
- "supports_audio_output": false,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_url_context": true,
- "supports_vision": true,
- "supports_web_search": true
- },
"gemini-2.5-flash-preview-09-2025": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_audio_token": 1e-06,
@@ -2194,6 +1335,11 @@
"mode": "chat",
"output_cost_per_reasoning_token": 2.5e-06,
"output_cost_per_token": 2.5e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2251,6 +1397,11 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_200k_tokens": 1.5e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2271,199 +1422,13 @@
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
+ "supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
"supports_web_search": true
},
- "gemini-2.5-pro-exp-03-25": {
- "cache_read_input_token_cost": 1.25e-07,
- "cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
- "input_cost_per_token": 1.25e-06,
- "input_cost_per_token_above_200k_tokens": 2.5e-06,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 65535,
- "max_pdf_size_mb": 30,
- "max_tokens": 65535,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_token": 1e-05,
- "output_cost_per_token_above_200k_tokens": 1.5e-05,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions"
- ],
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text"
- ],
- "supports_audio_input": true,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_video_input": true,
- "supports_vision": true,
- "supports_web_search": true
- },
- "gemini-2.5-pro-preview-03-25": {
- "cache_read_input_token_cost": 1.25e-07,
- "cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
- "deprecation_date": "2025-12-02",
- "input_cost_per_audio_token": 1.25e-06,
- "input_cost_per_token": 1.25e-06,
- "input_cost_per_token_above_200k_tokens": 2.5e-06,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 65535,
- "max_pdf_size_mb": 30,
- "max_tokens": 65535,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_token": 1e-05,
- "output_cost_per_token_above_200k_tokens": 1.5e-05,
- "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions",
- "/v1/batch"
- ],
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text"
- ],
- "supports_audio_output": false,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true
- },
- "gemini-2.5-pro-preview-05-06": {
- "cache_read_input_token_cost": 1.25e-07,
- "cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
- "deprecation_date": "2025-12-02",
- "input_cost_per_audio_token": 1.25e-06,
- "input_cost_per_token": 1.25e-06,
- "input_cost_per_token_above_200k_tokens": 2.5e-06,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 65535,
- "max_pdf_size_mb": 30,
- "max_tokens": 65535,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_token": 1e-05,
- "output_cost_per_token_above_200k_tokens": 1.5e-05,
- "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions",
- "/v1/batch"
- ],
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text"
- ],
- "supported_regions": [
- "global"
- ],
- "supports_audio_output": false,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true
- },
- "gemini-2.5-pro-preview-06-05": {
- "cache_read_input_token_cost": 1.25e-07,
- "cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
- "input_cost_per_audio_token": 1.25e-06,
- "input_cost_per_token": 1.25e-06,
- "input_cost_per_token_above_200k_tokens": 2.5e-06,
- "litellm_provider": "vertex_ai-language-models",
- "max_audio_length_hours": 8.4,
- "max_audio_per_prompt": 1,
- "max_images_per_prompt": 3000,
- "max_input_tokens": 1048576,
- "max_output_tokens": 65535,
- "max_pdf_size_mb": 30,
- "max_tokens": 65535,
- "max_video_length": 1,
- "max_videos_per_prompt": 10,
- "mode": "chat",
- "output_cost_per_token": 1e-05,
- "output_cost_per_token_above_200k_tokens": 1.5e-05,
- "source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions",
- "/v1/batch"
- ],
- "supported_modalities": [
- "text",
- "image",
- "audio",
- "video"
- ],
- "supported_output_modalities": [
- "text"
- ],
- "supports_audio_output": false,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true,
- "supports_web_search": true
- },
"gemini-2.5-pro-preview-tts": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
@@ -2483,6 +1448,11 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_200k_tokens": 1.5e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-2.5-pro-preview",
"supported_modalities": [
"text"
@@ -2500,7 +1470,7 @@
"supports_vision": true,
"supports_web_search": true
},
- "gemini-3-flash-preview": {
+ "gemini-3-flash": {
"cache_read_input_token_cost": 5e-08,
"cache_read_input_token_cost_priority": 9e-08,
"input_cost_per_audio_token": 1e-06,
@@ -2521,6 +1491,11 @@
"output_cost_per_reasoning_token": 3e-06,
"output_cost_per_token": 3e-06,
"output_cost_per_token_priority": 5.4e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
"source": "https://ai.google.dev/pricing/gemini-3",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2549,7 +1524,65 @@
"supports_tool_choice": true,
"supports_url_context": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
+ },
+ "gemini-3-flash-preview": {
+ "cache_read_input_token_cost": 5e-08,
+ "cache_read_input_token_cost_priority": 9e-08,
+ "input_cost_per_audio_token": 1e-06,
+ "input_cost_per_audio_token_priority": 1.8e-06,
+ "input_cost_per_token": 5e-07,
+ "input_cost_per_token_priority": 9e-07,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_audio_length_hours": 8.4,
+ "max_audio_per_prompt": 1,
+ "max_images_per_prompt": 3000,
+ "max_input_tokens": 1048576,
+ "max_output_tokens": 65535,
+ "max_pdf_size_mb": 30,
+ "max_tokens": 65535,
+ "max_video_length": 1,
+ "max_videos_per_prompt": 10,
+ "mode": "chat",
+ "output_cost_per_reasoning_token": 3e-06,
+ "output_cost_per_token": 3e-06,
+ "output_cost_per_token_priority": 5.4e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
+ "source": "https://ai.google.dev/pricing/gemini-3",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio",
+ "video"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_audio_output": false,
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_url_context": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
},
"gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
@@ -2564,6 +1597,11 @@
"output_cost_per_image_token": 0.00012,
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_batches": 6e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2581,9 +1619,11 @@
"supports_function_calling": false,
"supports_prompt_caching": true,
"supports_response_schema": true,
+ "supports_service_tier": true,
"supports_system_messages": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
},
"gemini-3-pro-preview": {
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@@ -2591,6 +1631,7 @@
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
"cache_read_input_token_cost_priority": 3.6e-07,
+ "deprecation_date": "2026-03-26",
"input_cost_per_token": 2e-06,
"input_cost_per_token_above_200k_tokens": 4e-06,
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
@@ -2612,6 +1653,11 @@
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"output_cost_per_token_priority": 2.16e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2639,6 +1685,240 @@
"supports_tool_choice": true,
"supports_video_input": true,
"supports_vision": true,
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
+ },
+ "gemini-3.1-flash-image": {
+ "input_cost_per_image": 0.00056,
+ "input_cost_per_token": 5e-07,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_input_tokens": 65536,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "image_generation",
+ "output_cost_per_image": 0.0672,
+ "output_cost_per_image_token": 6e-05,
+ "output_cost_per_token": 3e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
+ "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "image"
+ ],
+ "supports_function_calling": false,
+ "supports_prompt_caching": true,
+ "supports_response_schema": true,
+ "supports_system_messages": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
+ },
+ "gemini-3.1-flash-image-preview": {
+ "input_cost_per_image": 0.00056,
+ "input_cost_per_token": 5e-07,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_input_tokens": 65536,
+ "max_output_tokens": 32768,
+ "max_tokens": 32768,
+ "mode": "image_generation",
+ "output_cost_per_image": 0.0672,
+ "output_cost_per_image_token": 6e-05,
+ "output_cost_per_token": 3e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
+ "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "image"
+ ],
+ "supports_function_calling": false,
+ "supports_prompt_caching": true,
+ "supports_response_schema": true,
+ "supports_system_messages": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
+ },
+ "gemini-3.1-flash-lite": {
+ "cache_read_input_token_cost": 2.5e-08,
+ "cache_read_input_token_cost_batches": 1.25e-08,
+ "cache_read_input_token_cost_flex": 1.25e-08,
+ "cache_read_input_token_cost_per_audio_token": 5e-08,
+ "cache_read_input_token_cost_priority": 4.5e-08,
+ "input_cost_per_audio_token": 5e-07,
+ "input_cost_per_token": 2.5e-07,
+ "input_cost_per_token_batches": 1.25e-07,
+ "input_cost_per_token_flex": 1.25e-07,
+ "input_cost_per_token_priority": 4.5e-07,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_audio_length_hours": 8.4,
+ "max_audio_per_prompt": 1,
+ "max_images_per_prompt": 3000,
+ "max_input_tokens": 1048576,
+ "max_output_tokens": 65536,
+ "max_pdf_size_mb": 30,
+ "max_tokens": 65536,
+ "max_video_length": 1,
+ "max_videos_per_prompt": 10,
+ "mode": "chat",
+ "output_cost_per_reasoning_token": 1.5e-06,
+ "output_cost_per_token": 1.5e-06,
+ "output_cost_per_token_batches": 7.5e-07,
+ "output_cost_per_token_flex": 7.5e-07,
+ "output_cost_per_token_priority": 2.7e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
+ "source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-lite",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio",
+ "video"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": false,
+ "supports_code_execution": true,
+ "supports_file_search": true,
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_url_context": true,
+ "supports_video_input": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
+ },
+ "gemini-3.1-flash-lite-preview": {
+ "cache_read_input_token_cost": 2.5e-08,
+ "cache_read_input_token_cost_per_audio_token": 5e-08,
+ "input_cost_per_audio_token": 5e-07,
+ "input_cost_per_token": 2.5e-07,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_audio_length_hours": 8.4,
+ "max_audio_per_prompt": 1,
+ "max_images_per_prompt": 3000,
+ "max_input_tokens": 1048576,
+ "max_output_tokens": 65536,
+ "max_pdf_size_mb": 30,
+ "max_tokens": 65536,
+ "max_video_length": 1,
+ "max_videos_per_prompt": 10,
+ "mode": "chat",
+ "output_cost_per_reasoning_token": 1.5e-06,
+ "output_cost_per_token": 1.5e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
+ "source": "https://ai.google.dev/gemini-api/docs/models",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio",
+ "video"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": false,
+ "supports_code_execution": true,
+ "supports_file_search": true,
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_url_context": true,
+ "supports_video_input": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
+ },
+ "gemini-3.1-flash-live-preview": {
+ "input_cost_per_audio_token": 3e-06,
+ "input_cost_per_image_token": 1e-06,
+ "input_cost_per_token": 7.5e-07,
+ "input_cost_per_video_per_second": 3.3333333333333335e-05,
+ "litellm_provider": "gemini",
+ "max_input_tokens": 131072,
+ "max_output_tokens": 65536,
+ "max_tokens": 65536,
+ "mode": "chat",
+ "output_cost_per_audio_token": 1.2e-05,
+ "output_cost_per_token": 4.5e-06,
+ "source": "https://ai.google.dev/gemini-api/docs/pricing",
+ "supported_endpoints": [
+ "/v1/realtime"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio",
+ "video"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "audio"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": true,
+ "supports_function_calling": true,
+ "supports_vision": true,
"supports_web_search": true
},
"gemini-3.1-pro-high": {
@@ -2669,6 +1949,11 @@
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"output_cost_per_token_priority": 2.16e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2697,7 +1982,8 @@
"supports_url_context": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
},
"gemini-3.1-pro-low": {
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@@ -2727,6 +2013,11 @@
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"output_cost_per_token_priority": 2.16e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2755,7 +2046,8 @@
"supports_url_context": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
},
"gemini-3.1-pro-preview": {
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@@ -2785,6 +2077,11 @@
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
"output_cost_per_token_batches": 6e-06,
"output_cost_per_token_priority": 2.16e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2813,7 +2110,8 @@
"supports_url_context": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
},
"gemini-3.1-pro-preview-customtools": {
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@@ -2837,6 +2135,11 @@
"output_cost_per_token": 1.2e-05,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"output_cost_per_token_batches": 6e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2864,7 +2167,67 @@
"supports_url_context": true,
"supports_video_input": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
+ },
+ "gemini-3.5-flash": {
+ "cache_read_input_token_cost": 1.5e-07,
+ "cache_read_input_token_cost_priority": 2.7e-07,
+ "input_cost_per_audio_token": 1e-06,
+ "input_cost_per_audio_token_priority": 1.8e-06,
+ "input_cost_per_token": 1.5e-06,
+ "input_cost_per_token_priority": 2.7e-06,
+ "litellm_provider": "vertex_ai-language-models",
+ "max_audio_length_hours": 8.4,
+ "max_audio_per_prompt": 1,
+ "max_images_per_prompt": 3000,
+ "max_input_tokens": 1048576,
+ "max_output_tokens": 65535,
+ "max_pdf_size_mb": 30,
+ "max_tokens": 65535,
+ "max_video_length": 1,
+ "max_videos_per_prompt": 10,
+ "mode": "chat",
+ "output_cost_per_reasoning_token": 9e-06,
+ "output_cost_per_token": 9e-06,
+ "output_cost_per_token_priority": 1.62e-05,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.014,
+ "search_context_size_low": 0.014,
+ "search_context_size_medium": 0.014
+ },
+ "source": "https://ai.google.dev/pricing/gemini-3",
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/completions",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio",
+ "video"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": false,
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_url_context": true,
+ "supports_video_input": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "web_search_billing_unit": "per_query"
},
"gemini-embedding-001": {
"input_cost_per_token": 1.5e-07,
@@ -2876,6 +2239,35 @@
"output_vector_size": 3072,
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models"
},
+ "gemini-embedding-2": {
+ "input_cost_per_audio_per_second": 0.00016,
+ "input_cost_per_image": 0.00012,
+ "input_cost_per_token": 2e-07,
+ "input_cost_per_video_per_second": 0.00079,
+ "litellm_provider": "vertex_ai-embedding-models",
+ "max_input_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "embedding",
+ "output_cost_per_token": 0,
+ "output_vector_size": 3072,
+ "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "supports_multimodal": true,
+ "uses_embed_content": true
+ },
+ "gemini-embedding-2-preview": {
+ "input_cost_per_audio_per_second": 0.00016,
+ "input_cost_per_image": 0.00012,
+ "input_cost_per_token": 2e-07,
+ "input_cost_per_video_per_second": 0.00079,
+ "litellm_provider": "vertex_ai-embedding-models",
+ "max_input_tokens": 8192,
+ "max_tokens": 8192,
+ "mode": "embedding",
+ "output_cost_per_token": 0,
+ "output_vector_size": 3072,
+ "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "uses_embed_content": true
+ },
"gemini-exp-1206": {
"cache_read_input_token_cost": 3e-08,
"input_cost_per_audio_token": 1e-06,
@@ -2894,6 +2286,11 @@
"output_cost_per_reasoning_token": 2.5e-06,
"output_cost_per_token": 2.5e-06,
"rpm": 100000,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
"supported_endpoints": [
"/v1/chat/completions",
@@ -2930,13 +2327,11 @@
"max_input_tokens": 1000000,
"max_output_tokens": 8192,
"max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 0,
+ "mode": "embedding",
"output_cost_per_token": 0,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/gemini-experimental",
- "supports_function_calling": false,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true
+ "output_vector_size": 3072,
+ "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
+ "uses_embed_content": true
},
"gemini-flash-latest": {
"cache_read_input_token_cost": 3e-08,
@@ -2956,6 +2351,11 @@
"output_cost_per_reasoning_token": 2.5e-06,
"output_cost_per_token": 2.5e-06,
"rpm": 100000,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-preview",
"supported_endpoints": [
"/v1/chat/completions",
@@ -3003,6 +2403,11 @@
"output_cost_per_reasoning_token": 4e-07,
"output_cost_per_token": 4e-07,
"rpm": 15,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://ai.google.dev/gemini-api/docs/models#gemini-2.5-flash-lite",
"supported_endpoints": [
"/v1/chat/completions",
@@ -3046,13 +2451,17 @@
"max_tokens": 65535,
"max_video_length": 1,
"max_videos_per_prompt": 10,
- "mode": "chat",
+ "mode": "realtime",
"output_cost_per_audio_token": 1.2e-05,
"output_cost_per_token": 2e-06,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_endpoints": [
- "/v1/chat/completions",
- "/v1/completions"
+ "/vertex_ai/live"
],
"supported_modalities": [
"text",
@@ -3077,38 +2486,6 @@
"supports_vision": true,
"supports_web_search": true
},
- "gemini-pro": {
- "input_cost_per_character": 1.25e-07,
- "input_cost_per_image": 0.0025,
- "input_cost_per_token": 5e-07,
- "input_cost_per_video_per_second": 0.002,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 32760,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 3.75e-07,
- "output_cost_per_token": 1.5e-06,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true
- },
- "gemini-pro-experimental": {
- "input_cost_per_character": 0,
- "input_cost_per_token": 0,
- "litellm_provider": "vertex_ai-language-models",
- "max_input_tokens": 1000000,
- "max_output_tokens": 8192,
- "max_tokens": 8192,
- "mode": "chat",
- "output_cost_per_character": 0,
- "output_cost_per_token": 0,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/multimodal/gemini-experimental",
- "supports_function_calling": false,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true
- },
"gemini-pro-latest": {
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
@@ -3128,6 +2505,11 @@
"output_cost_per_token": 1e-05,
"output_cost_per_token_above_200k_tokens": 1.5e-05,
"rpm": 2000,
+ "search_context_cost_per_query": {
+ "search_context_size_high": 0.035,
+ "search_context_size_low": 0.035,
+ "search_context_size_medium": 0.035
+ },
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@@ -3155,24 +2537,6 @@
"supports_web_search": true,
"tpm": 800000
},
- "gemini-pro-vision": {
- "input_cost_per_image": 0.0025,
- "input_cost_per_token": 5e-07,
- "litellm_provider": "vertex_ai-vision-models",
- "max_images_per_prompt": 16,
- "max_input_tokens": 16384,
- "max_output_tokens": 2048,
- "max_tokens": 2048,
- "max_video_length": 2,
- "max_videos_per_prompt": 1,
- "mode": "chat",
- "output_cost_per_token": 1.5e-06,
- "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
"gemini-robotics-er-1.5-preview": {
"cache_read_input_token_cost": 0,
"input_cost_per_audio_token": 1e-06,
@@ -3236,31 +2600,6 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
- "gpt-3.5-turbo-0301": {
- "input_cost_per_token": 1.5e-06,
- "litellm_provider": "openai",
- "max_input_tokens": 4097,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_token": 2e-06,
- "supports_prompt_caching": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
- "gpt-3.5-turbo-0613": {
- "input_cost_per_token": 1.5e-06,
- "litellm_provider": "openai",
- "max_input_tokens": 4097,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_token": 2e-06,
- "supports_function_calling": true,
- "supports_prompt_caching": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
"gpt-3.5-turbo-1106": {
"deprecation_date": "2026-09-28",
"input_cost_per_token": 1e-06,
@@ -3288,18 +2627,6 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
- "gpt-3.5-turbo-16k-0613": {
- "input_cost_per_token": 3e-06,
- "litellm_provider": "openai",
- "max_input_tokens": 16385,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_token": 4e-06,
- "supports_prompt_caching": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
"gpt-3.5-turbo-instruct": {
"input_cost_per_token": 1.5e-06,
"litellm_provider": "text-completion-openai",
@@ -3347,6 +2674,7 @@
"supports_tool_choice": true
},
"gpt-4-0314": {
+ "deprecation_date": "2026-03-26",
"input_cost_per_token": 3e-05,
"litellm_provider": "openai",
"max_input_tokens": 8192,
@@ -3354,7 +2682,6 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 6e-05,
- "supports_prompt_caching": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
@@ -3387,57 +2714,6 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
- "gpt-4-1106-vision-preview": {
- "deprecation_date": "2024-12-06",
- "input_cost_per_token": 1e-05,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_token": 3e-05,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gpt-4-32k": {
- "input_cost_per_token": 6e-05,
- "litellm_provider": "openai",
- "max_input_tokens": 32768,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_token": 0.00012,
- "supports_prompt_caching": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
- "gpt-4-32k-0314": {
- "input_cost_per_token": 6e-05,
- "litellm_provider": "openai",
- "max_input_tokens": 32768,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_token": 0.00012,
- "supports_prompt_caching": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
- "gpt-4-32k-0613": {
- "input_cost_per_token": 6e-05,
- "litellm_provider": "openai",
- "max_input_tokens": 32768,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_token": 0.00012,
- "supports_prompt_caching": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
"gpt-4-turbo": {
"input_cost_per_token": 1e-05,
"litellm_provider": "openai",
@@ -3485,21 +2761,6 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
- "gpt-4-vision-preview": {
- "deprecation_date": "2024-12-06",
- "input_cost_per_token": 1e-05,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_token": 3e-05,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
"gpt-4.1": {
"cache_read_input_token_cost": 5e-07,
"cache_read_input_token_cost_priority": 8.75e-07,
@@ -3535,7 +2796,8 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"gpt-4.1-2025-04-14": {
"cache_read_input_token_cost": 5e-07,
@@ -3569,7 +2831,8 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"gpt-4.1-mini": {
"cache_read_input_token_cost": 1e-07,
@@ -3606,7 +2869,8 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"gpt-4.1-mini-2025-04-14": {
"cache_read_input_token_cost": 1e-07,
@@ -3640,7 +2904,8 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"gpt-4.1-nano": {
"cache_read_input_token_cost": 2.5e-08,
@@ -3713,47 +2978,6 @@
"supports_tool_choice": true,
"supports_vision": true
},
- "gpt-4.5-preview": {
- "cache_read_input_token_cost": 3.75e-05,
- "input_cost_per_token": 7.5e-05,
- "input_cost_per_token_batches": 3.75e-05,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 16384,
- "max_tokens": 16384,
- "mode": "chat",
- "output_cost_per_token": 0.00015,
- "output_cost_per_token_batches": 7.5e-05,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gpt-4.5-preview-2025-02-27": {
- "cache_read_input_token_cost": 3.75e-05,
- "deprecation_date": "2025-07-14",
- "input_cost_per_token": 7.5e-05,
- "input_cost_per_token_batches": 3.75e-05,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 16384,
- "max_tokens": 16384,
- "mode": "chat",
- "output_cost_per_token": 0.00015,
- "output_cost_per_token_batches": 7.5e-05,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_response_schema": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
"gpt-4o": {
"cache_read_input_token_cost": 1.25e-06,
"cache_read_input_token_cost_priority": 2.125e-06,
@@ -3857,23 +3081,6 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
- "gpt-4o-audio-preview-2024-10-01": {
- "input_cost_per_audio_token": 4e-05,
- "input_cost_per_token": 2.5e-06,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 16384,
- "max_tokens": 16384,
- "mode": "chat",
- "output_cost_per_audio_token": 8e-05,
- "output_cost_per_token": 1e-05,
- "supports_audio_input": true,
- "supports_audio_output": true,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
"gpt-4o-audio-preview-2024-12-17": {
"input_cost_per_audio_token": 4e-05,
"input_cost_per_token": 2.5e-06,
@@ -4077,7 +3284,7 @@
"supports_vision": true
},
"gpt-4o-mini-transcribe": {
- "input_cost_per_audio_token": 3e-06,
+ "input_cost_per_audio_token": 1.25e-06,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 16000,
@@ -4089,7 +3296,7 @@
]
},
"gpt-4o-mini-transcribe-2025-03-20": {
- "input_cost_per_audio_token": 3e-06,
+ "input_cost_per_audio_token": 1.25e-06,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 16000,
@@ -4101,7 +3308,7 @@
]
},
"gpt-4o-mini-transcribe-2025-12-15": {
- "input_cost_per_audio_token": 3e-06,
+ "input_cost_per_audio_token": 1.25e-06,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "openai",
"max_input_tokens": 16000,
@@ -4184,25 +3391,6 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
- "gpt-4o-realtime-preview-2024-10-01": {
- "cache_creation_input_audio_token_cost": 2e-05,
- "cache_read_input_token_cost": 2.5e-06,
- "input_cost_per_audio_token": 0.0001,
- "input_cost_per_token": 5e-06,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 4096,
- "max_tokens": 4096,
- "mode": "chat",
- "output_cost_per_audio_token": 0.0002,
- "output_cost_per_token": 2e-05,
- "supports_audio_input": true,
- "supports_audio_output": true,
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_system_messages": true,
- "supports_tool_choice": true
- },
"gpt-4o-realtime-preview-2024-12-17": {
"cache_read_input_token_cost": 2.5e-06,
"input_cost_per_audio_token": 4e-05,
@@ -4286,7 +3474,7 @@
"supports_vision": true
},
"gpt-4o-transcribe": {
- "input_cost_per_audio_token": 6e-06,
+ "input_cost_per_audio_token": 2.5e-06,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 16000,
@@ -4298,7 +3486,7 @@
]
},
"gpt-4o-transcribe-diarize": {
- "input_cost_per_audio_token": 6e-06,
+ "input_cost_per_audio_token": 2.5e-06,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 16000,
@@ -4337,7 +3525,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4346,7 +3536,9 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-2025-08-07": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4376,7 +3568,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4385,7 +3579,9 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-chat": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4409,7 +3605,9 @@
"text"
],
"supports_function_calling": false,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": false,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4417,7 +3615,8 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": false,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-chat-latest": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4441,7 +3640,9 @@
"text"
],
"supports_function_calling": false,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": false,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4449,7 +3650,8 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": false,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-codex": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4471,7 +3673,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4479,7 +3683,9 @@
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-mini": {
"cache_read_input_token_cost": 2.5e-08,
@@ -4509,7 +3715,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4518,7 +3726,9 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-mini-2025-08-07": {
"cache_read_input_token_cost": 2.5e-08,
@@ -4548,7 +3758,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4557,7 +3769,9 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-nano": {
"cache_read_input_token_cost": 5e-09,
@@ -4585,7 +3799,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4593,7 +3809,9 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-nano-2025-08-07": {
"cache_read_input_token_cost": 5e-09,
@@ -4620,7 +3838,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4628,7 +3848,9 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-pro": {
"input_cost_per_token": 1.5e-05,
@@ -4652,7 +3874,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": false,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4661,7 +3885,8 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-pro-2025-10-06": {
"input_cost_per_token": 1.5e-05,
@@ -4685,7 +3910,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": false,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4694,7 +3921,8 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-search-api": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4706,6 +3934,8 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4713,7 +3943,8 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5-search-api-2025-10-14": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4725,6 +3956,7 @@
"mode": "chat",
"output_cost_per_token": 1e-05,
"supports_function_calling": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4732,7 +3964,8 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5.1": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4759,7 +3992,9 @@
"image"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4768,7 +4003,9 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5.1-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4795,7 +4032,9 @@
"image"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4804,7 +4043,9 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5.1-chat-latest": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4831,7 +4072,9 @@
"image"
],
"supports_function_calling": false,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
"supports_parallel_function_calling": false,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4839,7 +4082,9 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": false,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5.1-codex": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4864,7 +4109,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4872,7 +4119,9 @@
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5.1-codex-max": {
"cache_read_input_token_cost": 1.25e-07,
@@ -4894,7 +4143,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4902,7 +4153,9 @@
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
},
"gpt-5.1-codex-mini": {
"cache_read_input_token_cost": 2.5e-08,
@@ -4927,7 +4180,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4935,7 +4190,9 @@
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5.2": {
"cache_read_input_token_cost": 1.75e-07,
@@ -4963,7 +4220,9 @@
"image"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -4972,7 +4231,9 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
},
"gpt-5.2-2025-12-11": {
"cache_read_input_token_cost": 1.75e-07,
@@ -5000,7 +4261,9 @@
"image"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -5009,7 +4272,9 @@
"supports_service_tier": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
},
"gpt-5.2-chat-latest": {
"cache_read_input_token_cost": 1.75e-07,
@@ -5035,7 +4300,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -5043,7 +4310,9 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5.2-codex": {
"cache_read_input_token_cost": 1.75e-07,
@@ -5068,7 +4337,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -5076,7 +4347,9 @@
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
},
"gpt-5.2-pro": {
"input_cost_per_token": 2.1e-05,
@@ -5098,7 +4371,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -5107,7 +4382,8 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
},
"gpt-5.2-pro-2025-12-11": {
"input_cost_per_token": 2.1e-05,
@@ -5129,7 +4405,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -5138,17 +4416,21 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
- "supports_web_search": true
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
},
- "gpt-5.4": {
- "cache_read_input_token_cost": 2.5e-07,
- "input_cost_per_token": 2.5e-06,
+ "gpt-5.3-chat-latest": {
+ "cache_read_input_token_cost": 1.75e-07,
+ "cache_read_input_token_cost_priority": 3.5e-07,
+ "input_cost_per_token": 1.75e-06,
+ "input_cost_per_token_priority": 3.5e-06,
"litellm_provider": "openai",
- "max_input_tokens": 1050000,
- "max_output_tokens": 128000,
- "max_tokens": 128000,
+ "max_input_tokens": 128000,
+ "max_output_tokens": 16384,
+ "max_tokens": 16384,
"mode": "chat",
- "output_cost_per_token": 1.5e-05,
+ "output_cost_per_token": 1.4e-05,
+ "output_cost_per_token_priority": 2.8e-05,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
@@ -5157,111 +4439,13 @@
"text",
"image"
],
- "supported_output_modalities": [
- "text",
- "image"
- ],
- "supports_function_calling": true,
- "supports_native_streaming": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_service_tier": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "codex-auto-review": {
- "cache_read_input_token_cost": 2.5e-07,
- "input_cost_per_token": 2.5e-06,
- "litellm_provider": "openai",
- "max_input_tokens": 1050000,
- "max_output_tokens": 128000,
- "max_tokens": 128000,
- "mode": "chat",
- "output_cost_per_token": 1.5e-05,
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/responses"
- ],
- "supported_modalities": [
- "text",
- "image"
- ],
- "supported_output_modalities": [
- "text",
- "image"
- ],
- "supports_function_calling": true,
- "supports_native_streaming": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_service_tier": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gpt-5.4-mini": {
- "cache_read_input_token_cost": 7.5e-08,
- "input_cost_per_token": 7.5e-07,
- "litellm_provider": "openai",
- "max_input_tokens": 400000,
- "max_output_tokens": 128000,
- "max_tokens": 128000,
- "mode": "chat",
- "output_cost_per_token": 4.5e-06,
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/batch",
- "/v1/responses"
- ],
- "supported_modalities": [
- "text",
- "image"
- ],
- "supported_output_modalities": [
- "text"
- ],
- "supports_function_calling": true,
- "supports_native_streaming": true,
- "supports_parallel_function_calling": true,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_response_schema": true,
- "supports_service_tier": true,
- "supports_system_messages": true,
- "supports_tool_choice": true,
- "supports_vision": true
- },
- "gpt-5.4-nano": {
- "cache_read_input_token_cost": 2e-08,
- "input_cost_per_token": 2e-07,
- "litellm_provider": "openai",
- "max_input_tokens": 400000,
- "max_output_tokens": 128000,
- "max_tokens": 128000,
- "mode": "chat",
- "output_cost_per_token": 1.25e-06,
- "supported_endpoints": [
- "/v1/chat/completions",
- "/v1/batch",
- "/v1/responses"
- ],
- "supported_modalities": [
- "text",
- "image"
- ],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -5269,7 +4453,9 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
},
"gpt-5.3-codex": {
"cache_read_input_token_cost": 1.75e-07,
@@ -5294,7 +4480,9 @@
"text"
],
"supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
"supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
@@ -5302,8 +4490,586 @@
"supports_response_schema": true,
"supports_system_messages": false,
"supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
+ },
+ "gpt-5.3-codex-spark": {
+ "cache_read_input_token_cost": 1.75e-07,
+ "cache_read_input_token_cost_priority": 3.5e-07,
+ "input_cost_per_token": 1.75e-06,
+ "input_cost_per_token_priority": 3.5e-06,
+ "litellm_provider": "openai",
+ "max_input_tokens": 272000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "responses",
+ "output_cost_per_token": 1.4e-05,
+ "output_cost_per_token_priority": 2.8e-05,
+ "supported_endpoints": [
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_system_messages": false,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": false
+ },
+ "gpt-5.4": {
+ "cache_read_input_token_cost": 2.5e-07,
+ "cache_read_input_token_cost_above_272k_tokens": 5e-07,
+ "cache_read_input_token_cost_flex": 1.3e-07,
+ "cache_read_input_token_cost_priority": 5e-07,
+ "input_cost_per_token": 2.5e-06,
+ "input_cost_per_token_above_272k_tokens": 5e-06,
+ "input_cost_per_token_batches": 1.25e-06,
+ "input_cost_per_token_flex": 1.25e-06,
+ "input_cost_per_token_priority": 5e-06,
+ "litellm_provider": "openai",
+ "max_input_tokens": 1050000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 1.5e-05,
+ "output_cost_per_token_above_272k_tokens": 2.25e-05,
+ "output_cost_per_token_batches": 7.5e-06,
+ "output_cost_per_token_flex": 7.5e-06,
+ "output_cost_per_token_priority": 3e-05,
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/batch",
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.4-2026-03-05": {
+ "cache_read_input_token_cost": 2.5e-07,
+ "cache_read_input_token_cost_above_272k_tokens": 5e-07,
+ "cache_read_input_token_cost_flex": 1.3e-07,
+ "cache_read_input_token_cost_priority": 5e-07,
+ "input_cost_per_token": 2.5e-06,
+ "input_cost_per_token_above_272k_tokens": 5e-06,
+ "input_cost_per_token_batches": 1.25e-06,
+ "input_cost_per_token_flex": 1.25e-06,
+ "input_cost_per_token_priority": 5e-06,
+ "litellm_provider": "openai",
+ "max_input_tokens": 1050000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 1.5e-05,
+ "output_cost_per_token_above_272k_tokens": 2.25e-05,
+ "output_cost_per_token_batches": 7.5e-06,
+ "output_cost_per_token_flex": 7.5e-06,
+ "output_cost_per_token_priority": 3e-05,
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/batch",
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
"supports_vision": true
},
+ "gpt-5.4-mini": {
+ "cache_read_input_token_cost": 7.5e-08,
+ "cache_read_input_token_cost_batches": 3.75e-08,
+ "cache_read_input_token_cost_flex": 3.75e-08,
+ "cache_read_input_token_cost_priority": 1.5e-07,
+ "input_cost_per_token": 7.5e-07,
+ "input_cost_per_token_batches": 3.75e-07,
+ "input_cost_per_token_flex": 3.75e-07,
+ "input_cost_per_token_priority": 1.5e-06,
+ "litellm_provider": "openai",
+ "max_input_tokens": 400000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 4.5e-06,
+ "output_cost_per_token_batches": 2.25e-06,
+ "output_cost_per_token_flex": 2.25e-06,
+ "output_cost_per_token_priority": 9e-06,
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/batch",
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": false,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.4-mini-2026-03-17": {
+ "cache_read_input_token_cost": 7.5e-08,
+ "cache_read_input_token_cost_batches": 3.75e-08,
+ "cache_read_input_token_cost_flex": 3.75e-08,
+ "cache_read_input_token_cost_priority": 1.5e-07,
+ "input_cost_per_token": 7.5e-07,
+ "input_cost_per_token_batches": 3.75e-07,
+ "input_cost_per_token_flex": 3.75e-07,
+ "input_cost_per_token_priority": 1.5e-06,
+ "litellm_provider": "openai",
+ "max_input_tokens": 272000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 4.5e-06,
+ "output_cost_per_token_batches": 2.25e-06,
+ "output_cost_per_token_flex": 2.25e-06,
+ "output_cost_per_token_priority": 9e-06,
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/batch",
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": false,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.4-nano": {
+ "cache_read_input_token_cost": 2e-08,
+ "cache_read_input_token_cost_batches": 1e-08,
+ "cache_read_input_token_cost_flex": 1e-08,
+ "input_cost_per_token": 2e-07,
+ "input_cost_per_token_batches": 1e-07,
+ "input_cost_per_token_flex": 1e-07,
+ "litellm_provider": "openai",
+ "max_input_tokens": 400000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 1.25e-06,
+ "output_cost_per_token_batches": 6.25e-07,
+ "output_cost_per_token_flex": 6.25e-07,
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/batch",
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": false,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.4-nano-2026-03-17": {
+ "cache_read_input_token_cost": 2e-08,
+ "cache_read_input_token_cost_batches": 1e-08,
+ "cache_read_input_token_cost_flex": 1e-08,
+ "input_cost_per_token": 2e-07,
+ "input_cost_per_token_batches": 1e-07,
+ "input_cost_per_token_flex": 1e-07,
+ "litellm_provider": "openai",
+ "max_input_tokens": 272000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 1.25e-06,
+ "output_cost_per_token_batches": 6.25e-07,
+ "output_cost_per_token_flex": 6.25e-07,
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/batch",
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": false,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.4-pro": {
+ "cache_read_input_token_cost": 3e-06,
+ "cache_read_input_token_cost_above_272k_tokens": 6e-06,
+ "input_cost_per_token": 3e-05,
+ "input_cost_per_token_above_272k_tokens": 6e-05,
+ "input_cost_per_token_batches": 1.5e-05,
+ "input_cost_per_token_flex": 1.5e-05,
+ "litellm_provider": "openai",
+ "max_input_tokens": 1050000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "responses",
+ "output_cost_per_token": 0.00018,
+ "output_cost_per_token_above_272k_tokens": 0.00027,
+ "output_cost_per_token_batches": 9e-05,
+ "output_cost_per_token_flex": 9e-05,
+ "supported_endpoints": [
+ "/v1/responses",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": false,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.4-pro-2026-03-05": {
+ "cache_read_input_token_cost": 3e-06,
+ "cache_read_input_token_cost_above_272k_tokens": 6e-06,
+ "input_cost_per_token": 3e-05,
+ "input_cost_per_token_above_272k_tokens": 6e-05,
+ "input_cost_per_token_batches": 1.5e-05,
+ "input_cost_per_token_flex": 1.5e-05,
+ "litellm_provider": "openai",
+ "max_input_tokens": 1050000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "responses",
+ "output_cost_per_token": 0.00018,
+ "output_cost_per_token_above_272k_tokens": 0.00027,
+ "output_cost_per_token_batches": 9e-05,
+ "output_cost_per_token_flex": 9e-05,
+ "supported_endpoints": [
+ "/v1/responses",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": true,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": false,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.5": {
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_272k_tokens": 1e-06,
+ "cache_read_input_token_cost_flex": 2.5e-07,
+ "cache_read_input_token_cost_priority": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_272k_tokens": 1e-05,
+ "input_cost_per_token_batches": 2.5e-06,
+ "input_cost_per_token_flex": 2.5e-06,
+ "input_cost_per_token_priority": 1e-05,
+ "litellm_provider": "openai",
+ "max_input_tokens": 1050000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 3e-05,
+ "output_cost_per_token_above_272k_tokens": 4.5e-05,
+ "output_cost_per_token_batches": 1.5e-05,
+ "output_cost_per_token_flex": 1.5e-05,
+ "output_cost_per_token_priority": 6e-05,
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/batch",
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": false,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.5-2026-04-23": {
+ "cache_read_input_token_cost": 5e-07,
+ "cache_read_input_token_cost_above_272k_tokens": 1e-06,
+ "cache_read_input_token_cost_flex": 2.5e-07,
+ "cache_read_input_token_cost_priority": 1e-06,
+ "input_cost_per_token": 5e-06,
+ "input_cost_per_token_above_272k_tokens": 1e-05,
+ "input_cost_per_token_batches": 2.5e-06,
+ "input_cost_per_token_flex": 2.5e-06,
+ "input_cost_per_token_priority": 1e-05,
+ "litellm_provider": "openai",
+ "max_input_tokens": 1050000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "chat",
+ "output_cost_per_token": 3e-05,
+ "output_cost_per_token_above_272k_tokens": 4.5e-05,
+ "output_cost_per_token_batches": 1.5e-05,
+ "output_cost_per_token_flex": 1.5e-05,
+ "output_cost_per_token_priority": 6e-05,
+ "supported_endpoints": [
+ "/v1/chat/completions",
+ "/v1/batch",
+ "/v1/responses"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_minimal_reasoning_effort": false,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": true,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": true,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.5-pro": {
+ "cache_read_input_token_cost": 3e-06,
+ "cache_read_input_token_cost_above_272k_tokens": 6e-06,
+ "input_cost_per_token": 3e-05,
+ "input_cost_per_token_above_272k_tokens": 6e-05,
+ "input_cost_per_token_batches": 1.5e-05,
+ "input_cost_per_token_flex": 1.5e-05,
+ "litellm_provider": "openai",
+ "max_input_tokens": 1050000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "responses",
+ "output_cost_per_token": 0.00018,
+ "output_cost_per_token_above_272k_tokens": 0.00027,
+ "output_cost_per_token_batches": 9e-05,
+ "output_cost_per_token_flex": 9e-05,
+ "supported_endpoints": [
+ "/v1/responses",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_low_reasoning_effort": false,
+ "supports_minimal_reasoning_effort": false,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": false,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
+ "gpt-5.5-pro-2026-04-23": {
+ "cache_read_input_token_cost": 3e-06,
+ "cache_read_input_token_cost_above_272k_tokens": 6e-06,
+ "input_cost_per_token": 3e-05,
+ "input_cost_per_token_above_272k_tokens": 6e-05,
+ "input_cost_per_token_batches": 1.5e-05,
+ "input_cost_per_token_flex": 1.5e-05,
+ "litellm_provider": "openai",
+ "max_input_tokens": 1050000,
+ "max_output_tokens": 128000,
+ "max_tokens": 128000,
+ "mode": "responses",
+ "output_cost_per_token": 0.00018,
+ "output_cost_per_token_above_272k_tokens": 0.00027,
+ "output_cost_per_token_batches": 9e-05,
+ "output_cost_per_token_flex": 9e-05,
+ "supported_endpoints": [
+ "/v1/responses",
+ "/v1/batch"
+ ],
+ "supported_modalities": [
+ "text",
+ "image"
+ ],
+ "supported_output_modalities": [
+ "text"
+ ],
+ "supports_function_calling": true,
+ "supports_low_reasoning_effort": false,
+ "supports_minimal_reasoning_effort": false,
+ "supports_native_streaming": true,
+ "supports_none_reasoning_effort": false,
+ "supports_parallel_function_calling": true,
+ "supports_pdf_input": true,
+ "supports_prompt_caching": true,
+ "supports_reasoning": true,
+ "supports_response_schema": false,
+ "supports_service_tier": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": true,
+ "supports_web_search": true,
+ "supports_xhigh_reasoning_effort": true
+ },
"gpt-audio": {
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_token": 2.5e-06,
@@ -5340,6 +5106,39 @@
"supports_tool_choice": true,
"supports_vision": false
},
+ "gpt-audio-1.5": {
+ "input_cost_per_audio_token": 3.2e-05,
+ "input_cost_per_token": 2.5e-06,
+ "litellm_provider": "openai",
+ "max_input_tokens": 128000,
+ "max_output_tokens": 16384,
+ "max_tokens": 16384,
+ "mode": "chat",
+ "output_cost_per_audio_token": 6.4e-05,
+ "output_cost_per_token": 1e-05,
+ "supported_endpoints": [
+ "/v1/chat/completions"
+ ],
+ "supported_modalities": [
+ "text",
+ "audio"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "audio"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": true,
+ "supports_function_calling": true,
+ "supports_native_streaming": true,
+ "supports_parallel_function_calling": true,
+ "supports_prompt_caching": false,
+ "supports_reasoning": false,
+ "supports_response_schema": false,
+ "supports_system_messages": true,
+ "supports_tool_choice": true,
+ "supports_vision": false
+ },
"gpt-audio-2025-08-28": {
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_token": 2.5e-06,
@@ -5540,6 +5339,38 @@
"supports_pdf_input": true,
"supports_vision": true
},
+ "gpt-image-2": {
+ "cache_read_input_image_token_cost": 2e-06,
+ "cache_read_input_token_cost": 1.25e-06,
+ "input_cost_per_image_token": 8e-06,
+ "input_cost_per_token": 5e-06,
+ "litellm_provider": "openai",
+ "mode": "image_generation",
+ "output_cost_per_image_token": 3e-05,
+ "output_cost_per_token": 1e-05,
+ "supported_endpoints": [
+ "/v1/images/generations",
+ "/v1/images/edits"
+ ],
+ "supports_pdf_input": true,
+ "supports_vision": true
+ },
+ "gpt-image-2-2026-04-21": {
+ "cache_read_input_image_token_cost": 2e-06,
+ "cache_read_input_token_cost": 1.25e-06,
+ "input_cost_per_image_token": 8e-06,
+ "input_cost_per_token": 5e-06,
+ "litellm_provider": "openai",
+ "mode": "image_generation",
+ "output_cost_per_image_token": 3e-05,
+ "output_cost_per_token": 1e-05,
+ "supported_endpoints": [
+ "/v1/images/generations",
+ "/v1/images/edits"
+ ],
+ "supports_pdf_input": true,
+ "supports_vision": true
+ },
"gpt-realtime": {
"cache_creation_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
@@ -5572,6 +5403,70 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
+ "gpt-realtime-1.5": {
+ "cache_creation_input_audio_token_cost": 4e-07,
+ "cache_read_input_token_cost": 4e-07,
+ "input_cost_per_audio_token": 3.2e-05,
+ "input_cost_per_image": 5e-06,
+ "input_cost_per_token": 4e-06,
+ "litellm_provider": "openai",
+ "max_input_tokens": 32000,
+ "max_output_tokens": 4096,
+ "max_tokens": 4096,
+ "mode": "chat",
+ "output_cost_per_audio_token": 6.4e-05,
+ "output_cost_per_token": 1.6e-05,
+ "supported_endpoints": [
+ "/v1/realtime"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "audio"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": true,
+ "supports_function_calling": true,
+ "supports_parallel_function_calling": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true
+ },
+ "gpt-realtime-2": {
+ "cache_creation_input_audio_token_cost": 4e-07,
+ "cache_read_input_token_cost": 4e-07,
+ "input_cost_per_audio_token": 3.2e-05,
+ "input_cost_per_image": 5e-06,
+ "input_cost_per_token": 4e-06,
+ "litellm_provider": "openai",
+ "max_input_tokens": 32000,
+ "max_output_tokens": 4096,
+ "max_tokens": 4096,
+ "mode": "chat",
+ "output_cost_per_audio_token": 6.4e-05,
+ "output_cost_per_token": 1.6e-05,
+ "supported_endpoints": [
+ "/v1/realtime"
+ ],
+ "supported_modalities": [
+ "text",
+ "image",
+ "audio"
+ ],
+ "supported_output_modalities": [
+ "text",
+ "audio"
+ ],
+ "supports_audio_input": true,
+ "supports_audio_output": true,
+ "supports_function_calling": true,
+ "supports_parallel_function_calling": true,
+ "supports_system_messages": true,
+ "supports_tool_choice": true
+ },
"gpt-realtime-2025-08-28": {
"cache_creation_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
@@ -5720,62 +5615,6 @@
"supports_tool_choice": true,
"supports_vision": true
},
- "o1-mini": {
- "cache_read_input_token_cost": 5.5e-07,
- "input_cost_per_token": 1.1e-06,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 65536,
- "max_tokens": 65536,
- "mode": "chat",
- "output_cost_per_token": 4.4e-06,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_vision": true
- },
- "o1-mini-2024-09-12": {
- "cache_read_input_token_cost": 1.5e-06,
- "deprecation_date": "2025-10-27",
- "input_cost_per_token": 3e-06,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 65536,
- "max_tokens": 65536,
- "mode": "chat",
- "output_cost_per_token": 1.2e-05,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_vision": true
- },
- "o1-preview": {
- "cache_read_input_token_cost": 7.5e-06,
- "input_cost_per_token": 1.5e-05,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 32768,
- "max_tokens": 32768,
- "mode": "chat",
- "output_cost_per_token": 6e-05,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_vision": true
- },
- "o1-preview-2024-09-12": {
- "cache_read_input_token_cost": 7.5e-06,
- "input_cost_per_token": 1.5e-05,
- "litellm_provider": "openai",
- "max_input_tokens": 128000,
- "max_output_tokens": 32768,
- "max_tokens": 32768,
- "mode": "chat",
- "output_cost_per_token": 6e-05,
- "supports_pdf_input": true,
- "supports_prompt_caching": true,
- "supports_reasoning": true,
- "supports_vision": true
- },
"o1-pro": {
"input_cost_per_token": 0.00015,
"input_cost_per_token_batches": 7.5e-05,
@@ -5876,7 +5715,8 @@
"supports_response_schema": true,
"supports_service_tier": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"o3-2025-04-16": {
"cache_read_input_token_cost": 5e-07,
@@ -5908,7 +5748,8 @@
"supports_response_schema": true,
"supports_service_tier": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"o3-deep-research": {
"cache_read_input_token_cost": 2.5e-06,
@@ -5941,7 +5782,8 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"o3-deep-research-2025-06-26": {
"cache_read_input_token_cost": 2.5e-06,
@@ -5974,7 +5816,8 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"o3-mini": {
"cache_read_input_token_cost": 5.5e-07,
@@ -6038,7 +5881,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"o3-pro-2025-06-10": {
"input_cost_per_token": 2e-05,
@@ -6068,7 +5912,8 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"o4-mini": {
"cache_read_input_token_cost": 2.75e-07,
@@ -6093,7 +5938,8 @@
"supports_response_schema": true,
"supports_service_tier": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"o4-mini-2025-04-16": {
"cache_read_input_token_cost": 2.75e-07,
@@ -6112,7 +5958,8 @@
"supports_response_schema": true,
"supports_service_tier": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"o4-mini-deep-research": {
"cache_read_input_token_cost": 5e-07,
@@ -6145,7 +5992,8 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
},
"o4-mini-deep-research-2025-06-26": {
"cache_read_input_token_cost": 5e-07,
@@ -6178,6 +6026,7 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
- "supports_vision": true
+ "supports_vision": true,
+ "supports_web_search": true
}
}
diff --git a/deploy/Dockerfile b/deploy/Dockerfile
index a947158f..d39dd17d 100644
--- a/deploy/Dockerfile
+++ b/deploy/Dockerfile
@@ -7,7 +7,7 @@
# =============================================================================
ARG NODE_IMAGE=node:24-alpine
-ARG GOLANG_IMAGE=golang:1.26.3-alpine
+ARG GOLANG_IMAGE=golang:1.26.4-alpine
ARG ALPINE_IMAGE=alpine:3.20
ARG GOPROXY=https://goproxy.cn,direct
ARG GOSUMDB=sum.golang.google.cn
diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml
index 31b38a19..35c76964 100644
--- a/deploy/config.example.yaml
+++ b/deploy/config.example.yaml
@@ -320,6 +320,14 @@ gateway:
queue: 0.7
error_rate: 0.8
ttft: 0.5
+ # OpenAI 高级调度器补充配置
+ openai_scheduler:
+ # 是否允许 session_hash sticky 在账号健康度恶化时临时逃逸;false 可一键回退旧行为
+ sticky_escape_enabled: true
+ # TTFT EWMA 超过该阈值(毫秒)时跳过 sticky,默认 15s,避免轻微抖动就逃逸
+ sticky_escape_ttft_ms: 15000
+ # 错误率 EWMA 超过该阈值时跳过 sticky,默认 0.5,仅在明显降级时触发
+ sticky_escape_error_rate: 0.5
# OpenAI HTTP upstream protocol strategy.
# OpenAI HTTP 上游协议策略(默认 HTTP/2;代理明确不兼容时可临时回退 HTTP/1.1)。
openai_http2:
diff --git a/frontend/src/api/admin/accounts.ts b/frontend/src/api/admin/accounts.ts
index 6fd23c47..bb75e302 100644
--- a/frontend/src/api/admin/accounts.ts
+++ b/frontend/src/api/admin/accounts.ts
@@ -487,6 +487,23 @@ export async function syncUpstreamModels(id: number): Promise {
+ const { data } = await apiClient.post('/admin/accounts/models/sync-upstream-preview', params)
+ return data
+}
+
export interface CRSPreviewAccount {
crs_account_id: string
kind: string
@@ -703,6 +720,7 @@ export const accountsAPI = {
setSchedulable,
getAvailableModels,
syncUpstreamModels,
+ syncUpstreamModelsPreview,
generateAuthUrl,
exchangeCode,
refreshOpenAIToken,
diff --git a/frontend/src/api/admin/groups.ts b/frontend/src/api/admin/groups.ts
index 6b94b799..b7846efd 100644
--- a/frontend/src/api/admin/groups.ts
+++ b/frontend/src/api/admin/groups.ts
@@ -76,6 +76,23 @@ export async function getById(id: number): Promise {
return data
}
+/**
+ * Get candidate models for custom /v1/models list.
+ * id=0 returns platform default models for create flow.
+ */
+export async function getModelsListCandidates(
+ id: number,
+ platform?: GroupPlatform
+): Promise {
+ const { data } = await apiClient.get<{ models: string[] }>(
+ `/admin/groups/${id}/models-list-candidates`,
+ {
+ params: platform ? { platform } : undefined
+ }
+ )
+ return data.models || []
+}
+
/**
* Create new group
* @param groupData - Group data
@@ -306,6 +323,7 @@ export const groupsAPI = {
getAll,
getByPlatform,
getById,
+ getModelsListCandidates,
create,
update,
delete: deleteGroup,
diff --git a/frontend/src/api/admin/ops.ts b/frontend/src/api/admin/ops.ts
index 69235668..a3a47c2c 100644
--- a/frontend/src/api/admin/ops.ts
+++ b/frontend/src/api/admin/ops.ts
@@ -778,9 +778,15 @@ export interface OpsAlertRuntimeSettings {
thresholds: OpsMetricThresholds // 指标阈值配置
}
+export interface OpsOpenAIAccountQuotaAutoPauseSettings {
+ default_threshold_5h: number // 0~1,0 表示不启用全局默认 5h 阈值
+ default_threshold_7d: number // 0~1,0 表示不启用全局默认 7d 阈值
+}
+
export interface OpsAdvancedSettings {
data_retention: OpsDataRetentionSettings
aggregation: OpsAggregationSettings
+ openai_account_quota_auto_pause: OpsOpenAIAccountQuotaAutoPauseSettings
ignore_count_tokens_errors: boolean
ignore_context_canceled: boolean
ignore_no_available_accounts: boolean
@@ -901,6 +907,9 @@ export interface OpsErrorLog {
user_id?: number | null
user_email: string
api_key_id?: number | null
+ // 关联 api_key 名称(后端 LEFT JOIN api_keys;软删保留 name,故已删 key 仍有原名)。
+ api_key_name?: string
+ api_key_deleted?: boolean
account_id?: number | null
account_name: string
group_id?: number | null
@@ -935,6 +944,15 @@ export interface OpsErrorDetail extends OpsErrorLog {
time_to_first_token_ms?: number | null
is_business_limited: boolean
+
+ // Deleted key owner info (INVALID_API_KEY attribution)
+ attempted_key_prefix?: string | null
+ deleted_key_owner_user_id?: number | null
+ deleted_key_owner_email?: string | null
+ deleted_key_name?: string | null
+
+ // Bound (non-deleted) key prefix, snapshotted at error time
+ api_key_prefix?: string | null
}
export type OpsErrorLogsResponse = PaginatedResponse
@@ -1069,6 +1087,10 @@ export type OpsErrorListQueryParams = {
platform?: string
group_id?: number | null
account_id?: number | null
+ user_id?: number
+ api_key_id?: number
+ // 模型过滤:后端以 COALESCE(requested_model, model) 精确匹配(admin 路径)。
+ model?: string
phase?: string
error_owner?: string
diff --git a/frontend/src/api/admin/riskControl.ts b/frontend/src/api/admin/riskControl.ts
index 521114c2..aefd1618 100644
--- a/frontend/src/api/admin/riskControl.ts
+++ b/frontend/src/api/admin/riskControl.ts
@@ -132,6 +132,16 @@ export interface ContentModerationRuntimeStatus {
dropped: number
processed: number
errors: number
+ pre_block_active: number
+ pre_block_checked: number
+ pre_block_allowed: number
+ pre_block_blocked: number
+ pre_block_errors: number
+ pre_block_avg_latency_ms: number
+ pre_block_api_key_active: number
+ pre_block_api_key_available_count: number
+ pre_block_api_key_total_calls: number
+ pre_block_api_key_loads: ContentModerationAPIKeyLoad[]
api_key_statuses: ContentModerationAPIKeyStatus[]
flagged_hash_count: number
last_cleanup_at?: string
@@ -139,6 +149,20 @@ export interface ContentModerationRuntimeStatus {
last_cleanup_deleted_non_hit: number
}
+export interface ContentModerationAPIKeyLoad {
+ index: number
+ key_hash: string
+ masked: string
+ status: ContentModerationAPIKeyStatusValue
+ active: number
+ total: number
+ success: number
+ errors: number
+ avg_latency_ms: number
+ last_latency_ms: number
+ last_http_status: number
+}
+
export interface ContentModerationLog {
id: number
request_id: string
diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts
index d2b878cc..5be63076 100644
--- a/frontend/src/api/admin/settings.ts
+++ b/frontend/src/api/admin/settings.ts
@@ -560,6 +560,7 @@ export interface SystemSettings {
rewrite_message_cache_control: boolean;
antigravity_user_agent_version: string;
openai_codex_user_agent: string;
+ openai_allow_claude_code_codex_plugin: boolean;
web_search_emulation_enabled?: boolean;
// Payment configuration
@@ -611,6 +612,9 @@ export interface SystemSettings {
// OpenAI fast/flex policy
openai_fast_policy_settings?: OpenAIFastPolicySettings;
+
+ // Allow user view error requests
+ allow_user_view_error_requests: boolean;
}
export interface UpdateSettingsRequest {
@@ -792,6 +796,7 @@ export interface UpdateSettingsRequest {
rewrite_message_cache_control?: boolean;
antigravity_user_agent_version?: string;
openai_codex_user_agent?: string;
+ openai_allow_claude_code_codex_plugin?: boolean;
// Payment configuration
payment_enabled?: boolean;
risk_control_enabled?: boolean;
@@ -840,6 +845,8 @@ export interface UpdateSettingsRequest {
// OpenAI fast/flex policy
openai_fast_policy_settings?: OpenAIFastPolicySettings;
+
+ allow_user_view_error_requests?: boolean;
}
/**
diff --git a/frontend/src/api/admin/usage.ts b/frontend/src/api/admin/usage.ts
index 7ad00742..d933ac63 100644
--- a/frontend/src/api/admin/usage.ts
+++ b/frontend/src/api/admin/usage.ts
@@ -27,6 +27,7 @@ export interface AdminUsageStatsResponse {
export interface SimpleUser {
id: number
email: string
+ deleted: boolean
}
export interface SimpleApiKey {
@@ -120,6 +121,7 @@ export async function getStats(params: {
start_date?: string
end_date?: string
timezone?: string
+ nocache?: number
}): Promise {
const { data } = await apiClient.get('/admin/usage/stats', {
params
diff --git a/frontend/src/api/admin/users.ts b/frontend/src/api/admin/users.ts
index bfe5e3ba..61879d2f 100644
--- a/frontend/src/api/admin/users.ts
+++ b/frontend/src/api/admin/users.ts
@@ -100,10 +100,12 @@ export async function list(
/**
* Get user by ID
* @param id - User ID
+ * @param includeDeleted - Whether to include soft-deleted users
* @returns User details
*/
-export async function getById(id: number): Promise {
- const { data } = await apiClient.get(`/admin/users/${id}`)
+export async function getById(id: number, includeDeleted = false): Promise {
+ const url = includeDeleted ? `/admin/users/${id}?include_deleted=true` : `/admin/users/${id}`
+ const { data } = await apiClient.get(url)
return data
}
@@ -115,8 +117,11 @@ export async function getById(id: number): Promise {
export async function create(userData: {
email: string
password: string
+ username?: string
+ notes?: string
balance?: number
concurrency?: number
+ rpm_limit?: number
allowed_groups?: number[] | null
}): Promise {
const { data } = await apiClient.post('/admin/users', userData)
diff --git a/frontend/src/api/usage.ts b/frontend/src/api/usage.ts
index ee08ee9d..f0aec3d6 100644
--- a/frontend/src/api/usage.ts
+++ b/frontend/src/api/usage.ts
@@ -10,7 +10,10 @@ import type {
UsageStatsResponse,
PaginatedResponse,
TrendDataPoint,
- ModelStat
+ ModelStat,
+ UserErrorRequest,
+ UserErrorRequestDetail,
+ UserErrorListParams
} from '@/types'
// ==================== Dashboard Types ====================
@@ -304,6 +307,22 @@ export async function getDashboardApiKeysUsage(
return data
}
+export async function listMyErrorRequests(
+ params: UserErrorListParams,
+ config: { signal?: AbortSignal } = {}
+): Promise> {
+ const { data } = await apiClient.get>('/usage/errors', {
+ ...config,
+ params
+ })
+ return data
+}
+
+export async function getMyErrorDetail(id: number): Promise {
+ const { data } = await apiClient.get(`/usage/errors/${id}`)
+ return data
+}
+
export const usageAPI = {
list,
query,
@@ -316,7 +335,10 @@ export const usageAPI = {
getDashboardTrend,
getDashboardModels,
getMyApiKeyDailyUsage,
- getDashboardApiKeysUsage
+ getDashboardApiKeysUsage,
+ // Error requests
+ listMyErrorRequests,
+ getMyErrorDetail,
}
export default usageAPI
diff --git a/frontend/src/components/account/AccountStatusIndicator.vue b/frontend/src/components/account/AccountStatusIndicator.vue
index 8438c584..a37b5c8e 100644
--- a/frontend/src/components/account/AccountStatusIndicator.vue
+++ b/frontend/src/components/account/AccountStatusIndicator.vue
@@ -222,6 +222,8 @@ const formatScopeName = (scope: string): string => {
// Claude 系列
'claude-opus-4-6': 'COpus46',
'claude-opus-4-6-thinking': 'COpus46T',
+ 'claude-opus-4-7': 'COpus47',
+ 'claude-opus-4-8': 'COpus48',
'claude-sonnet-4-6': 'CSon46',
'claude-sonnet-4-5': 'CSon45',
'claude-sonnet-4-5-thinking': 'CSon45T',
diff --git a/frontend/src/components/account/AccountUsageCell.vue b/frontend/src/components/account/AccountUsageCell.vue
index 64f1366b..887fbf79 100644
--- a/frontend/src/components/account/AccountUsageCell.vue
+++ b/frontend/src/components/account/AccountUsageCell.vue
@@ -664,6 +664,7 @@ const antigravityClaudeUsageFromAPI = computed(() =>
getAntigravityUsageFromAPI([
'claude-sonnet-4-5', 'claude-opus-4-5-thinking',
'claude-sonnet-4-6', 'claude-opus-4-6', 'claude-opus-4-6-thinking',
+ 'claude-opus-4-7', 'claude-opus-4-8',
])
)
diff --git a/frontend/src/components/account/BulkEditAccountModal.vue b/frontend/src/components/account/BulkEditAccountModal.vue
index c8d53220..6e71fe4b 100644
--- a/frontend/src/components/account/BulkEditAccountModal.vue
+++ b/frontend/src/components/account/BulkEditAccountModal.vue
@@ -742,6 +742,50 @@
+
+
+
+
+
+
+
+
+ {{ t('admin.accounts.openai.codexCLIOnlyAllowClaudeCodeDesc') }}
+
+
+
+
+
@@ -1219,6 +1263,7 @@ const enableOpenAIPassthrough = ref(false)
const enableOpenAIWSMode = ref(false)
const enableOpenAIAPIKeyWSMode = ref(false)
const enableCodexCLIOnly = ref(false)
+const enableCodexCLIOnlyAllowClaudeCode = ref(false)
const enableOpenAICompactMode = ref(false)
const enableOpenAICompactModelMapping = ref(false)
const enableRpmLimit = ref(false)
@@ -1246,6 +1291,7 @@ const openaiPassthroughEnabled = ref(false)
const openaiOAuthResponsesWebSocketV2Mode = ref
(OPENAI_WS_MODE_OFF)
const openaiAPIKeyResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF)
const codexCLIOnlyEnabled = ref(false)
+const codexCLIOnlyAllowClaudeCodeEnabled = ref(false)
const openAICompactMode = ref('auto')
const openAICompactModelMappings = ref([])
const rpmLimitEnabled = ref(false)
@@ -1496,6 +1542,11 @@ const buildUpdatePayload = (): Record | null => {
extra.codex_cli_only = codexCLIOnlyEnabled.value
}
+ if (enableCodexCLIOnlyAllowClaudeCode.value) {
+ const extra = ensureExtra()
+ extra.codex_cli_only_allowed_clients = codexCLIOnlyAllowClaudeCodeEnabled.value ? ['claude_code'] : []
+ }
+
if (enableOpenAICompactMode.value) {
const extra = ensureExtra()
extra.openai_compact_mode = openAICompactMode.value
@@ -1602,6 +1653,7 @@ const handleSubmit = async () => {
enableOpenAIWSMode.value ||
enableOpenAIAPIKeyWSMode.value ||
enableCodexCLIOnly.value ||
+ enableCodexCLIOnlyAllowClaudeCode.value ||
enableOpenAICompactMode.value ||
enableOpenAICompactModelMapping.value ||
enableRpmLimit.value ||
@@ -1704,6 +1756,7 @@ watch(
enableOpenAIWSMode.value = false
enableOpenAIAPIKeyWSMode.value = false
enableCodexCLIOnly.value = false
+ enableCodexCLIOnlyAllowClaudeCode.value = false
enableOpenAICompactMode.value = false
enableOpenAICompactModelMapping.value = false
enableRpmLimit.value = false
@@ -1727,6 +1780,7 @@ watch(
openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF
openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF
codexCLIOnlyEnabled.value = false
+ codexCLIOnlyAllowClaudeCodeEnabled.value = false
openAICompactMode.value = 'auto'
openAICompactModelMappings.value = []
rpmLimitEnabled.value = false
diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue
index 90d5e15a..d1a7729a 100644
--- a/frontend/src/components/account/CreateAccountModal.vue
+++ b/frontend/src/components/account/CreateAccountModal.vue
@@ -1124,7 +1124,7 @@
-
+
{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }}
{{
@@ -1290,6 +1290,18 @@
}}
+
+
+
+
+ {{ t('admin.accounts.poolModeRetryStatusCodesHint', { default: DEFAULT_POOL_MODE_RETRY_STATUS_CODES.join(', ') }) }}
+
+
@@ -1550,7 +1562,7 @@
-
+
{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }}
{{ t('admin.accounts.supportsAllModels') }}
@@ -1635,6 +1647,18 @@
}}
+
+
+
+
+ {{ t('admin.accounts.poolModeRetryStatusCodesHint', { default: DEFAULT_POOL_MODE_RETRY_STATUS_CODES.join(', ') }) }}
+
+
@@ -1789,7 +1813,7 @@
-
+
{{ t('admin.accounts.selectedModels', { count: allowedModels.length }) }}
{{
@@ -2611,6 +2635,32 @@
/>
+
+
+
+
+ {{ t('admin.accounts.openai.codexCLIOnlyAllowClaudeCodeDesc') }}
+
+
+
+
@@ -2655,7 +2705,7 @@
+
+ {{ t('admin.accounts.openai.responsesModeTextDisabledHint') }}
+
+
+
+
+
+
+
{{ t('admin.accounts.openai.endpointCapabilitiesDesc') }}
+
@@ -3148,7 +3226,8 @@ import type {
CreateAccountRequest,
CodexSessionImportMessage,
OpenAICompactMode,
- OpenAIResponsesMode
+ OpenAIResponsesMode,
+ OpenAIEndpointCapability
} from '@/types'
import BaseDialog from '@/components/common/BaseDialog.vue'
import ConfirmDialog from '@/components/common/ConfirmDialog.vue'
@@ -3282,6 +3361,17 @@ const accountCategory = ref<'oauth-based' | 'apikey' | 'bedrock' | 'service_acco
const addMethod = ref
('oauth') // For oauth-based: 'oauth' or 'setup-token'
const apiKeyBaseUrl = ref('https://api.anthropic.com')
const apiKeyValue = ref('')
+
+const syncPreviewCredentials = computed(() => {
+ if (!apiKeyValue.value) return undefined
+ return {
+ platform: form.platform,
+ type: form.type,
+ base_url: apiKeyBaseUrl.value || undefined,
+ api_key: apiKeyValue.value
+ }
+})
+
const editQuotaLimit = ref(null)
const editQuotaDailyLimit = ref(null)
const editQuotaWeeklyLimit = ref(null)
@@ -3297,8 +3387,27 @@ const modelRestrictionMode = ref<'whitelist' | 'mapping'>('whitelist')
const allowedModels = ref([])
const DEFAULT_POOL_MODE_RETRY_COUNT = 3
const MAX_POOL_MODE_RETRY_COUNT = 10
+const DEFAULT_POOL_MODE_RETRY_STATUS_CODES = [401, 403, 429]
const poolModeEnabled = ref(false)
const poolModeRetryCount = ref(DEFAULT_POOL_MODE_RETRY_COUNT)
+const poolModeRetryStatusCodesInput = ref('')
+
+function parsePoolModeRetryStatusCodes(input: string): number[] {
+ if (!input || !input.trim()) return []
+ const seen = new Set()
+ const out: number[] = []
+ for (const token of input.split(/[,\s]+/)) {
+ const trimmed = token.trim()
+ if (!trimmed) continue
+ const n = Number(trimmed)
+ if (!Number.isFinite(n) || !Number.isInteger(n)) continue
+ if (n < 100 || n > 599) continue
+ if (seen.has(n)) continue
+ seen.add(n)
+ out.push(n)
+ }
+ return out.sort((a, b) => a - b)
+}
const customErrorCodesEnabled = ref(false)
const selectedErrorCodes = ref([])
const customErrorCodeInput = ref(null)
@@ -3307,9 +3416,11 @@ const autoPauseOnExpired = ref(true)
const openaiPassthroughEnabled = ref(false)
const openAICompactMode = ref('auto')
const openAIResponsesMode = ref('auto')
+const openAIEndpointCapabilities = ref(['chat_completions', 'embeddings'])
const openaiOAuthResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF)
const openaiAPIKeyResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF)
const codexCLIOnlyEnabled = ref(false)
+const codexCLIOnlyAllowClaudeCodeEnabled = ref(false)
const anthropicPassthroughEnabled = ref(false)
const webSearchEmulationMode = ref('default')
const webSearchGlobalEnabled = ref(false)
@@ -3369,6 +3480,58 @@ const openAIResponsesModeOptions = computed(() => [
{ value: 'force_responses', label: t('admin.accounts.openai.responsesModeForceResponses') },
{ value: 'force_chat_completions', label: t('admin.accounts.openai.responsesModeForceChatCompletions') }
])
+const openAITextEndpointCapabilityLabel = computed(() => {
+ if (openAIResponsesMode.value === 'force_responses') {
+ return t('admin.accounts.openai.capabilityResponses')
+ }
+ if (openAIResponsesMode.value === 'force_chat_completions') {
+ return t('admin.accounts.openai.capabilityChatCompletions')
+ }
+ return t('admin.accounts.openai.capabilityTextAuto')
+})
+const openAIEndpointCapabilityOptions = computed<{ value: OpenAIEndpointCapability; label: string }[]>(() => [
+ { value: 'chat_completions', label: openAITextEndpointCapabilityLabel.value },
+ { value: 'embeddings', label: t('admin.accounts.openai.capabilityEmbeddings') }
+])
+const openAITextGenerationCapabilityEnabled = computed(() =>
+ openAIEndpointCapabilities.value.includes('chat_completions')
+)
+
+const normalizeOpenAIEndpointCapabilities = (values: OpenAIEndpointCapability[]) => {
+ const allowed: OpenAIEndpointCapability[] = ['chat_completions', 'embeddings']
+ const selected = allowed.filter((value) => values.includes(value))
+ return selected.length > 0 ? selected : allowed
+}
+
+const toggleOpenAIEndpointCapability = (capability: OpenAIEndpointCapability, event?: Event) => {
+ if (openAIEndpointCapabilities.value.includes(capability)) {
+ if (openAIEndpointCapabilities.value.length <= 1) {
+ const input = event?.target as HTMLInputElement | null
+ if (input) input.checked = true
+ return
+ }
+ openAIEndpointCapabilities.value = openAIEndpointCapabilities.value.filter(
+ (value) => value !== capability
+ )
+ if (!openAITextGenerationCapabilityEnabled.value) {
+ openAIResponsesMode.value = 'auto'
+ }
+ return
+ }
+ openAIEndpointCapabilities.value = normalizeOpenAIEndpointCapabilities([
+ ...openAIEndpointCapabilities.value,
+ capability
+ ])
+}
+
+const applyOpenAIEndpointCapabilities = (credentials: Record) => {
+ const capabilities = normalizeOpenAIEndpointCapabilities(openAIEndpointCapabilities.value)
+ if (capabilities.length === 2) {
+ delete credentials.openai_capabilities
+ return
+ }
+ credentials.openai_capabilities = capabilities
+}
function buildAntigravityExtra(): Record | undefined {
const extra: Record = {}
@@ -3678,9 +3841,11 @@ watch(
}
if (newPlatform !== 'openai') {
openaiPassthroughEnabled.value = false
+ openAIEndpointCapabilities.value = ['chat_completions', 'embeddings']
openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF
openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF
codexCLIOnlyEnabled.value = false
+ codexCLIOnlyAllowClaudeCodeEnabled.value = false
}
if (newPlatform !== 'anthropic') {
anthropicPassthroughEnabled.value = false
@@ -3701,6 +3866,7 @@ watch(
([category, platform]) => {
if (platform === 'openai' && category !== 'oauth-based') {
codexCLIOnlyEnabled.value = false
+ codexCLIOnlyAllowClaudeCodeEnabled.value = false
}
if (platform !== 'anthropic' || category !== 'apikey') {
anthropicPassthroughEnabled.value = false
@@ -4068,6 +4234,7 @@ const resetForm = () => {
})
poolModeEnabled.value = false
poolModeRetryCount.value = DEFAULT_POOL_MODE_RETRY_COUNT
+ poolModeRetryStatusCodesInput.value = ''
customErrorCodesEnabled.value = false
selectedErrorCodes.value = []
customErrorCodeInput.value = null
@@ -4076,9 +4243,11 @@ const resetForm = () => {
openaiPassthroughEnabled.value = false
openAICompactMode.value = 'auto'
openAIResponsesMode.value = 'auto'
+ openAIEndpointCapabilities.value = ['chat_completions', 'embeddings']
openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF
openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF
codexCLIOnlyEnabled.value = false
+ codexCLIOnlyAllowClaudeCodeEnabled.value = false
anthropicPassthroughEnabled.value = false
webSearchEmulationMode.value = 'default'
// Reset quota control state
@@ -4157,13 +4326,26 @@ const buildOpenAIExtra = (base?: Record): Record {
if (poolModeEnabled.value) {
credentials.pool_mode = true
credentials.pool_mode_retry_count = normalizePoolModeRetryCount(poolModeRetryCount.value)
+ const parsedRetryStatusCodes = parsePoolModeRetryStatusCodes(poolModeRetryStatusCodesInput.value)
+ if (parsedRetryStatusCodes.length > 0) {
+ credentials.pool_mode_retry_status_codes = parsedRetryStatusCodes
+ }
}
applyInterceptWarmup(credentials, interceptWarmupRequests.value, 'create')
@@ -4450,6 +4636,7 @@ const handleSubmit = async () => {
}
}
if (form.platform === 'openai') {
+ applyOpenAIEndpointCapabilities(credentials)
const compactModelMapping = buildOpenAICompactModelMapping()
if (compactModelMapping) {
credentials.compact_model_mapping = compactModelMapping
@@ -4460,6 +4647,10 @@ const handleSubmit = async () => {
if (poolModeEnabled.value) {
credentials.pool_mode = true
credentials.pool_mode_retry_count = normalizePoolModeRetryCount(poolModeRetryCount.value)
+ const parsedRetryStatusCodes = parsePoolModeRetryStatusCodes(poolModeRetryStatusCodesInput.value)
+ if (parsedRetryStatusCodes.length > 0) {
+ credentials.pool_mode_retry_status_codes = parsedRetryStatusCodes
+ }
}
// Add custom error codes if enabled
@@ -4568,6 +4759,9 @@ const createAccountAndFinish = async (
}
}
if (platform === 'openai') {
+ if (type === 'apikey') {
+ applyOpenAIEndpointCapabilities(credentials)
+ }
const compactModelMapping = buildOpenAICompactModelMapping()
if (compactModelMapping) {
credentials.compact_model_mapping = compactModelMapping
diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue
index 070887fe..8dc85d0e 100644
--- a/frontend/src/components/account/EditAccountModal.vue
+++ b/frontend/src/components/account/EditAccountModal.vue
@@ -305,6 +305,18 @@
}}
+
+
+
+
+ {{ t('admin.accounts.poolModeRetryStatusCodesHint', { default: DEFAULT_POOL_MODE_RETRY_STATUS_CODES.join(', ') }) }}
+
+
@@ -973,6 +985,18 @@
}}
+
+
+
+
+ {{ t('admin.accounts.poolModeRetryStatusCodesHint', { default: DEFAULT_POOL_MODE_RETRY_STATUS_CODES.join(', ') }) }}
+
+
@@ -1415,7 +1439,7 @@
-
+
{{ t(openAIResponsesStatusKey) }}
+
+ {{ t('admin.accounts.openai.responsesModeTextDisabledHint') }}
+
+
+
+
+
+
+
{{ t('admin.accounts.openai.endpointCapabilitiesDesc') }}
+
@@ -1618,6 +1673,32 @@
/>
+
+
+
+
+ {{ t('admin.accounts.openai.codexCLIOnlyAllowClaudeCodeDesc') }}
+
+
+
+
+
+
+
+
+
+
+
{{ t('admin.accounts.autoPauseDisabledHint') }}
+
+
+
+
+
{{ t('admin.accounts.autoPauseThresholdHint') }}
+
+
+
+
+
+
+
{{ t('admin.accounts.autoPauseDisabledHint') }}
+
+
+
+
+
{{ t('admin.accounts.autoPauseThresholdHint') }}
+
+
+
('whitelist')
const allowedModels = ref([])
const DEFAULT_POOL_MODE_RETRY_COUNT = 3
const MAX_POOL_MODE_RETRY_COUNT = 10
+const DEFAULT_POOL_MODE_RETRY_STATUS_CODES = [401, 403, 429]
const poolModeEnabled = ref(false)
const poolModeRetryCount = ref(DEFAULT_POOL_MODE_RETRY_COUNT)
+const poolModeRetryStatusCodesInput = ref('')
+
+function parsePoolModeRetryStatusCodes(input: string): number[] {
+ if (!input || !input.trim()) return []
+ const seen = new Set()
+ const out: number[] = []
+ for (const token of input.split(/[,\s]+/)) {
+ const trimmed = token.trim()
+ if (!trimmed) continue
+ const n = Number(trimmed)
+ if (!Number.isFinite(n) || !Number.isInteger(n)) continue
+ if (n < 100 || n > 599) continue
+ if (seen.has(n)) continue
+ seen.add(n)
+ out.push(n)
+ }
+ return out.sort((a, b) => a - b)
+}
+
+function formatPoolModeRetryStatusCodes(value: unknown): string {
+ if (!Array.isArray(value)) return ''
+ const out: number[] = []
+ const seen = new Set()
+ for (const v of value) {
+ const n = typeof v === 'string' ? Number(v.trim()) : Number(v)
+ if (!Number.isFinite(n) || !Number.isInteger(n)) continue
+ if (n < 100 || n > 599) continue
+ if (seen.has(n)) continue
+ seen.add(n)
+ out.push(n)
+ }
+ return out.sort((a, b) => a - b).join(', ')
+}
const customErrorCodesEnabled = ref(false)
const selectedErrorCodes = ref([])
const customErrorCodeInput = ref(null)
const interceptWarmupRequests = ref(false)
const autoPauseOnExpired = ref(false)
+const autoPause5hThreshold = ref(null)
+const autoPause7dThreshold = ref(null)
+const autoPause5hDisabled = ref(false)
+const autoPause7dDisabled = ref(false)
const mixedScheduling = ref(false) // For antigravity accounts: enable mixed scheduling
const allowOverages = ref(false) // For antigravity accounts: enable AI Credits overages
const antigravityModelRestrictionMode = ref<'whitelist' | 'mapping'>('whitelist')
@@ -2375,9 +2580,11 @@ const customBaseUrl = ref('')
const openaiPassthroughEnabled = ref(false)
const openAICompactMode = ref('auto')
const openAIResponsesMode = ref('auto')
+const openAIEndpointCapabilities = ref(['chat_completions', 'embeddings'])
const openaiOAuthResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF)
const openaiAPIKeyResponsesWebSocketV2Mode = ref(OPENAI_WS_MODE_OFF)
const codexCLIOnlyEnabled = ref(false)
+const codexCLIOnlyAllowClaudeCodeEnabled = ref(false)
type CodexImageGenerationBridgeMode = 'inherit' | 'enabled' | 'disabled'
const codexImageGenerationBridgeMode = ref('inherit')
const anthropicPassthroughEnabled = ref(false)
@@ -2481,6 +2688,85 @@ const openAIResponsesModeOptions = computed(() => [
{ value: 'force_responses', label: t('admin.accounts.openai.responsesModeForceResponses') },
{ value: 'force_chat_completions', label: t('admin.accounts.openai.responsesModeForceChatCompletions') }
])
+const openAITextEndpointCapabilityLabel = computed(() => {
+ if (openAIResponsesMode.value === 'force_responses') {
+ return t('admin.accounts.openai.capabilityResponses')
+ }
+ if (openAIResponsesMode.value === 'force_chat_completions') {
+ return t('admin.accounts.openai.capabilityChatCompletions')
+ }
+ const extra = props.account?.extra as Record | undefined
+ if (extra?.openai_responses_supported === true) {
+ return t('admin.accounts.openai.capabilityResponsesAuto')
+ }
+ if (extra?.openai_responses_supported === false) {
+ return t('admin.accounts.openai.capabilityChatCompletionsAuto')
+ }
+ return t('admin.accounts.openai.capabilityTextAuto')
+})
+const openAIEndpointCapabilityOptions = computed<{ value: OpenAIEndpointCapability; label: string }[]>(() => [
+ { value: 'chat_completions', label: openAITextEndpointCapabilityLabel.value },
+ { value: 'embeddings', label: t('admin.accounts.openai.capabilityEmbeddings') }
+])
+const openAITextGenerationCapabilityEnabled = computed(() =>
+ openAIEndpointCapabilities.value.includes('chat_completions')
+)
+
+const normalizeOpenAIEndpointCapabilities = (values: OpenAIEndpointCapability[]) => {
+ const allowed: OpenAIEndpointCapability[] = ['chat_completions', 'embeddings']
+ const selected = allowed.filter((value) => values.includes(value))
+ return selected.length > 0 ? selected : allowed
+}
+
+const readOpenAIEndpointCapabilities = (credentials?: Record): OpenAIEndpointCapability[] => {
+ const raw = credentials?.openai_capabilities
+ if (Array.isArray(raw)) {
+ return normalizeOpenAIEndpointCapabilities(
+ raw.filter((value): value is OpenAIEndpointCapability =>
+ value === 'chat_completions' || value === 'embeddings'
+ )
+ )
+ }
+ if (raw !== null && typeof raw === 'object') {
+ const capabilityMap = raw as Record
+ return normalizeOpenAIEndpointCapabilities(
+ openAIEndpointCapabilityOptions.value
+ .map((option) => option.value)
+ .filter((value) => capabilityMap[value] === true)
+ )
+ }
+ return ['chat_completions', 'embeddings']
+}
+
+const toggleOpenAIEndpointCapability = (capability: OpenAIEndpointCapability, event?: Event) => {
+ if (openAIEndpointCapabilities.value.includes(capability)) {
+ if (openAIEndpointCapabilities.value.length <= 1) {
+ const input = event?.target as HTMLInputElement | null
+ if (input) input.checked = true
+ return
+ }
+ openAIEndpointCapabilities.value = openAIEndpointCapabilities.value.filter(
+ (value) => value !== capability
+ )
+ if (!openAITextGenerationCapabilityEnabled.value) {
+ openAIResponsesMode.value = 'auto'
+ }
+ return
+ }
+ openAIEndpointCapabilities.value = normalizeOpenAIEndpointCapabilities([
+ ...openAIEndpointCapabilities.value,
+ capability
+ ])
+}
+
+const applyOpenAIEndpointCapabilities = (credentials: Record) => {
+ const capabilities = normalizeOpenAIEndpointCapabilities(openAIEndpointCapabilities.value)
+ if (capabilities.length === 2) {
+ delete credentials.openai_capabilities
+ return
+ }
+ credentials.openai_capabilities = capabilities
+}
const normalizeOpenAIResponsesMode = (mode: unknown): OpenAIResponsesMode => {
if (mode === 'force_responses' || mode === 'force_chat_completions') {
return mode
@@ -2658,18 +2944,24 @@ const syncFormFromAccount = (newAccount: Account | null) => {
// Load mixed scheduling setting (only for antigravity accounts)
mixedScheduling.value = false
allowOverages.value = false
- const extra = newAccount.extra as Record | undefined
- mixedScheduling.value = extra?.mixed_scheduling === true
- allowOverages.value = extra?.allow_overages === true
+ const extra = newAccount.extra as Record | undefined
+ mixedScheduling.value = extra?.mixed_scheduling === true
+ allowOverages.value = extra?.allow_overages === true
+ autoPause5hThreshold.value = typeof extra?.auto_pause_5h_threshold === 'number' ? extra.auto_pause_5h_threshold * 100 : null
+ autoPause7dThreshold.value = typeof extra?.auto_pause_7d_threshold === 'number' ? extra.auto_pause_7d_threshold * 100 : null
+ autoPause5hDisabled.value = extra?.auto_pause_5h_disabled === true
+ autoPause7dDisabled.value = extra?.auto_pause_7d_disabled === true
// Load OpenAI passthrough toggle (OpenAI OAuth/API Key)
openaiPassthroughEnabled.value = false
openAICompactMode.value = 'auto'
openAIResponsesMode.value = 'auto'
+ openAIEndpointCapabilities.value = ['chat_completions', 'embeddings']
openAICompactModelMappings.value = []
openaiOAuthResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF
openaiAPIKeyResponsesWebSocketV2Mode.value = OPENAI_WS_MODE_OFF
codexCLIOnlyEnabled.value = false
+ codexCLIOnlyAllowClaudeCodeEnabled.value = false
codexImageGenerationBridgeMode.value = 'inherit'
anthropicPassthroughEnabled.value = false
webSearchEmulationMode.value = 'default'
@@ -2678,6 +2970,12 @@ const syncFormFromAccount = (newAccount: Account | null) => {
openAICompactMode.value = (extra?.openai_compact_mode as OpenAICompactMode) || 'auto'
if (newAccount.type === 'apikey') {
openAIResponsesMode.value = normalizeOpenAIResponsesMode(extra?.openai_responses_mode)
+ openAIEndpointCapabilities.value = readOpenAIEndpointCapabilities(
+ newAccount.credentials as Record | undefined
+ )
+ if (!openAITextGenerationCapabilityEnabled.value) {
+ openAIResponsesMode.value = 'auto'
+ }
}
const codexImageGenerationBridgeValue = typeof extra?.codex_image_generation_bridge === 'boolean'
? extra.codex_image_generation_bridge
@@ -2701,6 +2999,9 @@ const syncFormFromAccount = (newAccount: Account | null) => {
})
if (newAccount.type === 'oauth') {
codexCLIOnlyEnabled.value = extra?.codex_cli_only === true
+ codexCLIOnlyAllowClaudeCodeEnabled.value =
+ Array.isArray(extra?.codex_cli_only_allowed_clients) &&
+ (extra.codex_cli_only_allowed_clients as unknown[]).includes('claude_code')
}
const credentials = newAccount.credentials as Record | undefined
const compactMappings = credentials?.compact_model_mapping as Record | undefined
@@ -2807,6 +3108,7 @@ const syncFormFromAccount = (newAccount: Account | null) => {
poolModeRetryCount.value = normalizePoolModeRetryCount(
Number(credentials.pool_mode_retry_count ?? DEFAULT_POOL_MODE_RETRY_COUNT)
)
+ poolModeRetryStatusCodesInput.value = formatPoolModeRetryStatusCodes(credentials.pool_mode_retry_status_codes)
// Load custom error codes
customErrorCodesEnabled.value = credentials.custom_error_codes_enabled === true
@@ -2834,6 +3136,7 @@ const syncFormFromAccount = (newAccount: Account | null) => {
poolModeEnabled.value = bedrockCreds.pool_mode === true
const retryCount = bedrockCreds.pool_mode_retry_count
poolModeRetryCount.value = (typeof retryCount === 'number' && retryCount >= 0) ? retryCount : DEFAULT_POOL_MODE_RETRY_COUNT
+ poolModeRetryStatusCodesInput.value = formatPoolModeRetryStatusCodes(bedrockCreds.pool_mode_retry_status_codes)
// Load quota limits for bedrock
const bedrockExtra = (newAccount.extra as Record) || {}
@@ -2876,6 +3179,7 @@ const syncFormFromAccount = (newAccount: Account | null) => {
}
poolModeEnabled.value = false
poolModeRetryCount.value = DEFAULT_POOL_MODE_RETRY_COUNT
+ poolModeRetryStatusCodesInput.value = ''
customErrorCodesEnabled.value = false
selectedErrorCodes.value = []
}
@@ -3415,6 +3719,7 @@ const handleSubmit = async () => {
newCredentials.model_mapping = currentCredentials.model_mapping
}
if (props.account.platform === 'openai') {
+ applyOpenAIEndpointCapabilities(newCredentials)
const compactModelMapping = buildModelMappingObject('mapping', [], openAICompactModelMappings.value)
if (compactModelMapping) {
newCredentials.compact_model_mapping = compactModelMapping
@@ -3427,9 +3732,16 @@ const handleSubmit = async () => {
if (poolModeEnabled.value) {
newCredentials.pool_mode = true
newCredentials.pool_mode_retry_count = normalizePoolModeRetryCount(poolModeRetryCount.value)
+ const parsedRetryStatusCodes = parsePoolModeRetryStatusCodes(poolModeRetryStatusCodesInput.value)
+ if (parsedRetryStatusCodes.length > 0) {
+ newCredentials.pool_mode_retry_status_codes = parsedRetryStatusCodes
+ } else {
+ delete newCredentials.pool_mode_retry_status_codes
+ }
} else {
delete newCredentials.pool_mode
delete newCredentials.pool_mode_retry_count
+ delete newCredentials.pool_mode_retry_status_codes
}
// Add custom error codes if enabled
@@ -3545,9 +3857,16 @@ const handleSubmit = async () => {
if (poolModeEnabled.value) {
newCredentials.pool_mode = true
newCredentials.pool_mode_retry_count = normalizePoolModeRetryCount(poolModeRetryCount.value)
+ const parsedRetryStatusCodes = parsePoolModeRetryStatusCodes(poolModeRetryStatusCodesInput.value)
+ if (parsedRetryStatusCodes.length > 0) {
+ newCredentials.pool_mode_retry_status_codes = parsedRetryStatusCodes
+ } else {
+ delete newCredentials.pool_mode_retry_status_codes
+ }
} else {
delete newCredentials.pool_mode
delete newCredentials.pool_mode_retry_count
+ delete newCredentials.pool_mode_retry_status_codes
}
// Model mapping
@@ -3754,9 +4073,9 @@ const handleSubmit = async () => {
}
// For OpenAI OAuth/API Key accounts, handle passthrough mode in extra
- if (props.account.platform === 'openai' && (props.account.type === 'oauth' || props.account.type === 'apikey')) {
- const currentExtra = (props.account.extra as Record) || {}
- const newExtra: Record = { ...currentExtra }
+ if (props.account.platform === 'openai' && (props.account.type === 'oauth' || props.account.type === 'apikey')) {
+ const currentExtra = (props.account.extra as Record) || {}
+ const newExtra: Record = { ...currentExtra }
const hadCodexCLIOnlyEnabled = currentExtra.codex_cli_only === true
if (props.account.type === 'oauth') {
newExtra.openai_oauth_responses_websockets_v2_mode = openaiOAuthResponsesWebSocketV2Mode.value
@@ -3778,15 +4097,35 @@ const handleSubmit = async () => {
} else {
newExtra.openai_compact_mode = openAICompactMode.value
}
- if (props.account.type === 'apikey') {
- if (openAIResponsesMode.value === 'auto') {
+ if (props.account.type === 'apikey') {
+ if (!openAITextGenerationCapabilityEnabled.value || openAIResponsesMode.value === 'auto') {
delete newExtra.openai_responses_mode
} else {
newExtra.openai_responses_mode = openAIResponsesMode.value
}
- }
+ }
+ if (autoPause5hThreshold.value != null && autoPause5hThreshold.value > 0) {
+ newExtra.auto_pause_5h_threshold = autoPause5hThreshold.value / 100
+ } else {
+ delete newExtra.auto_pause_5h_threshold
+ }
+ if (autoPause7dThreshold.value != null && autoPause7dThreshold.value > 0) {
+ newExtra.auto_pause_7d_threshold = autoPause7dThreshold.value / 100
+ } else {
+ delete newExtra.auto_pause_7d_threshold
+ }
+ if (autoPause5hDisabled.value) {
+ newExtra.auto_pause_5h_disabled = true
+ } else {
+ delete newExtra.auto_pause_5h_disabled
+ }
+ if (autoPause7dDisabled.value) {
+ newExtra.auto_pause_7d_disabled = true
+ } else {
+ delete newExtra.auto_pause_7d_disabled
+ }
- delete newExtra.codex_image_generation_bridge_enabled
+ delete newExtra.codex_image_generation_bridge_enabled
if (codexImageGenerationBridgeMode.value === 'inherit') {
delete newExtra.codex_image_generation_bridge
} else {
@@ -3802,6 +4141,12 @@ const handleSubmit = async () => {
} else {
delete newExtra.codex_cli_only
}
+ // 仅当 codex_cli_only 开启且子开关开启时写入 Claude Code 插件白名单,否则清除避免孤立字段
+ if (codexCLIOnlyEnabled.value && codexCLIOnlyAllowClaudeCodeEnabled.value) {
+ newExtra.codex_cli_only_allowed_clients = ['claude_code']
+ } else {
+ delete newExtra.codex_cli_only_allowed_clients
+ }
}
updatePayload.extra = newExtra
diff --git a/frontend/src/components/account/ModelWhitelistSelector.vue b/frontend/src/components/account/ModelWhitelistSelector.vue
index 9a0d6af8..d4d726d0 100644
--- a/frontend/src/components/account/ModelWhitelistSelector.vue
+++ b/frontend/src/components/account/ModelWhitelistSelector.vue
@@ -133,6 +133,7 @@ import { ref, computed } from 'vue'
import { useI18n } from 'vue-i18n'
import { useAppStore } from '@/stores/app'
import { accountsAPI } from '@/api/admin/accounts'
+import type { SyncUpstreamPreviewParams } from '@/api/admin/accounts'
import ModelIcon from '@/components/common/ModelIcon.vue'
import Icon from '@/components/icons/Icon.vue'
import { allModels, getModelsByPlatform } from '@/composables/useModelWhitelist'
@@ -144,6 +145,12 @@ const props = defineProps<{
platform?: string
platforms?: string[]
accountId?: number
+ syncCredentials?: {
+ platform: string
+ type: string
+ base_url?: string
+ api_key: string
+ }
}>()
const emit = defineEmits<{
@@ -176,9 +183,14 @@ const normalizedPlatforms = computed(() => {
const upstreamSyncPlatforms = new Set(['anthropic', 'openai', 'gemini', 'antigravity'])
const canSyncUpstream = computed(() => {
- if (!props.accountId) return false
- if (normalizedPlatforms.value.length === 0) return true
- return normalizedPlatforms.value.some(platform => upstreamSyncPlatforms.has(platform.toLowerCase()))
+ if (props.accountId) {
+ if (normalizedPlatforms.value.length === 0) return true
+ return normalizedPlatforms.value.some(platform => upstreamSyncPlatforms.has(platform.toLowerCase()))
+ }
+ if (props.syncCredentials) {
+ return upstreamSyncPlatforms.has(props.syncCredentials.platform.toLowerCase())
+ }
+ return false
})
const availableOptions = computed(() => {
@@ -249,11 +261,20 @@ const fillRelated = () => {
}
const syncUpstreamModels = async () => {
- if (!props.accountId || isSyncingUpstream.value) return
+ if (isSyncingUpstream.value) return
+ if (!props.accountId && !props.syncCredentials) return
isSyncingUpstream.value = true
try {
- const result = await accountsAPI.syncUpstreamModels(props.accountId)
+ let result
+ if (props.accountId) {
+ result = await accountsAPI.syncUpstreamModels(props.accountId)
+ } else if (props.syncCredentials) {
+ result = await accountsAPI.syncUpstreamModelsPreview(props.syncCredentials as SyncUpstreamPreviewParams)
+ } else {
+ return
+ }
+
const upstreamModels = result.models.map(model => model.trim()).filter(Boolean)
if (upstreamModels.length === 0) {
appStore.showInfo(t('admin.accounts.syncUpstreamModelsEmpty'))
diff --git a/frontend/src/components/account/__tests__/BulkEditAccountModal.spec.ts b/frontend/src/components/account/__tests__/BulkEditAccountModal.spec.ts
index caa307fc..3ae75ee9 100644
--- a/frontend/src/components/account/__tests__/BulkEditAccountModal.spec.ts
+++ b/frontend/src/components/account/__tests__/BulkEditAccountModal.spec.ts
@@ -197,6 +197,25 @@ describe('BulkEditAccountModal', () => {
})
})
+ it('OpenAI OAuth 批量编辑应提交 codex_cli_only_allowed_clients 字段', async () => {
+ const wrapper = mountModal({
+ selectedPlatforms: ['openai'],
+ selectedTypes: ['oauth']
+ })
+
+ await wrapper.get('#bulk-edit-openai-codex-allow-claude-code-enabled').setValue(true)
+ await wrapper.get('#bulk-edit-openai-codex-allow-claude-code-toggle').trigger('click')
+ await wrapper.get('#bulk-edit-account-form').trigger('submit.prevent')
+ await flushPromises()
+
+ expect(adminAPI.accounts.bulkUpdate).toHaveBeenCalledTimes(1)
+ expect(adminAPI.accounts.bulkUpdate).toHaveBeenCalledWith([1, 2], {
+ extra: {
+ codex_cli_only_allowed_clients: ['claude_code']
+ }
+ })
+ })
+
it('OpenAI API Key 批量编辑应提交 API Key 专属 WS mode 字段', async () => {
const wrapper = mountModal({
selectedPlatforms: ['openai'],
diff --git a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts
index 0b8e939c..f4865de9 100644
--- a/frontend/src/components/account/__tests__/EditAccountModal.spec.ts
+++ b/frontend/src/components/account/__tests__/EditAccountModal.spec.ts
@@ -310,6 +310,137 @@ describe('EditAccountModal', () => {
expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_responses_supported).toBe(true)
})
+ it('submits OpenAI APIKey endpoint capabilities from credentials', async () => {
+ const account = buildAccount()
+ account.credentials.openai_capabilities = ['chat_completions']
+ updateAccountMock.mockReset()
+ checkMixedChannelRiskMock.mockReset()
+ checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false })
+ updateAccountMock.mockResolvedValue(account)
+
+ const wrapper = mountModal(account)
+
+ expect(wrapper.findAll('input[type="checkbox"]').some((input) => (input.element as HTMLInputElement).checked)).toBe(true)
+
+ await wrapper.get('form#edit-account-form').trigger('submit.prevent')
+
+ expect(updateAccountMock).toHaveBeenCalledTimes(1)
+ expect(updateAccountMock.mock.calls[0]?.[1]?.credentials?.openai_capabilities).toEqual([
+ 'chat_completions'
+ ])
+ })
+
+ it('submits OpenAI quota auto-pause thresholds in extra', async () => {
+ const account = buildAccount()
+ account.extra = {
+ auto_pause_5h_threshold: 0.9,
+ auto_pause_7d_threshold: 0.8
+ }
+ updateAccountMock.mockReset()
+ checkMixedChannelRiskMock.mockReset()
+ checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false })
+ updateAccountMock.mockResolvedValue(account)
+
+ const wrapper = mountModal(account)
+
+ await wrapper.get('[data-testid="auto-pause-5h-threshold"]').setValue('95')
+ await wrapper.get('[data-testid="auto-pause-7d-threshold"]').setValue('96')
+ await wrapper.get('form#edit-account-form').trigger('submit.prevent')
+
+ expect(updateAccountMock).toHaveBeenCalledTimes(1)
+ expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.auto_pause_5h_threshold).toBe(0.95)
+ expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.auto_pause_7d_threshold).toBe(0.96)
+ })
+
+ it('submits OpenAI quota auto-pause disable flag in extra', async () => {
+ // Toggling the per-account disable flag must persist as auto_pause_5h_disabled
+ // so an admin can exempt one account from auto-pause even when a global default
+ // threshold is configured (otherwise leaving the threshold blank would silently
+ // fall back to the global default).
+ const account = buildAccount()
+ updateAccountMock.mockReset()
+ checkMixedChannelRiskMock.mockReset()
+ checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false })
+ updateAccountMock.mockResolvedValue(account)
+
+ const wrapper = mountModal(account)
+
+ await wrapper.get('[data-testid="auto-pause-5h-disabled"]').trigger('click')
+ await wrapper.get('form#edit-account-form').trigger('submit.prevent')
+
+ expect(updateAccountMock).toHaveBeenCalledTimes(1)
+ expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.auto_pause_5h_disabled).toBe(true)
+ expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.auto_pause_7d_disabled).toBeUndefined()
+ })
+
+ it('keeps at least one OpenAI APIKey endpoint capability selected', async () => {
+ const account = buildAccount()
+ updateAccountMock.mockReset()
+ checkMixedChannelRiskMock.mockReset()
+ checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false })
+ updateAccountMock.mockResolvedValue(account)
+
+ const wrapper = mountModal(account)
+
+ const chatCheckbox = wrapper.get(
+ '[data-testid="openai-endpoint-capability-chat_completions"]'
+ )
+ const embeddingsCheckbox = wrapper.get(
+ '[data-testid="openai-endpoint-capability-embeddings"]'
+ )
+
+ expect(chatCheckbox.element.checked).toBe(true)
+ expect(embeddingsCheckbox.element.checked).toBe(true)
+
+ await embeddingsCheckbox.setValue(false)
+
+ expect(chatCheckbox.element.checked).toBe(true)
+ expect(embeddingsCheckbox.element.checked).toBe(false)
+
+ await chatCheckbox.setValue(false)
+
+ expect(chatCheckbox.element.checked).toBe(true)
+ expect(embeddingsCheckbox.element.checked).toBe(false)
+
+ await wrapper.get('form#edit-account-form').trigger('submit.prevent')
+
+ expect(updateAccountMock).toHaveBeenCalledTimes(1)
+ expect(updateAccountMock.mock.calls[0]?.[1]?.credentials?.openai_capabilities).toEqual([
+ 'chat_completions'
+ ])
+ })
+
+ it('disables text generation protocol when only embeddings requests are accepted', async () => {
+ const account = buildAccount()
+ account.credentials.openai_capabilities = ['embeddings']
+ account.extra = {
+ openai_responses_mode: 'force_responses',
+ openai_responses_supported: true
+ }
+ updateAccountMock.mockReset()
+ checkMixedChannelRiskMock.mockReset()
+ checkMixedChannelRiskMock.mockResolvedValue({ has_risk: false })
+ updateAccountMock.mockResolvedValue(account)
+
+ const wrapper = mountModal(account)
+
+ const responsesModeSelect = wrapper.get(
+ '[data-testid="openai-responses-mode-select"]'
+ )
+
+ expect(responsesModeSelect.element.disabled).toBe(true)
+ expect(wrapper.find('[data-testid="openai-responses-mode-not-applicable"]').exists()).toBe(true)
+
+ await wrapper.get('form#edit-account-form').trigger('submit.prevent')
+
+ expect(updateAccountMock).toHaveBeenCalledTimes(1)
+ expect(updateAccountMock.mock.calls[0]?.[1]?.credentials?.openai_capabilities).toEqual([
+ 'embeddings'
+ ])
+ expect(updateAccountMock.mock.calls[0]?.[1]?.extra).not.toHaveProperty('openai_responses_mode')
+ expect(updateAccountMock.mock.calls[0]?.[1]?.extra?.openai_responses_supported).toBe(true)
+ })
+
it('submits account-level Codex image generation bridge override', async () => {
const account = buildAccount()
account.extra = {
diff --git a/frontend/src/components/admin/usage/UsageFilters.vue b/frontend/src/components/admin/usage/UsageFilters.vue
index 66c2b4fa..a800f190 100644
--- a/frontend/src/components/admin/usage/UsageFilters.vue
+++ b/frontend/src/components/admin/usage/UsageFilters.vue
@@ -35,7 +35,7 @@
@click="selectUser(u)"
class="w-full px-4 py-2 text-left hover:bg-gray-100 dark:hover:bg-gray-700"
>
- {{ u.email }}
+ {{ u.email }}({{ t('admin.usage.userDeletedBadge') }})
#{{ u.id }}
@@ -168,7 +168,7 @@
diff --git a/frontend/src/components/user/UserErrorRequestsTable.vue b/frontend/src/components/user/UserErrorRequestsTable.vue
new file mode 100644
index 00000000..4bbb28cd
--- /dev/null
+++ b/frontend/src/components/user/UserErrorRequestsTable.vue
@@ -0,0 +1,173 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ | {{ t('usage.errors.model') }} |
+ {{ t('usage.errors.keyName') }} |
+ {{ t('usage.errors.endpoint') }} |
+ {{ t('usage.errors.status') }} |
+ {{ t('usage.errors.category') }} |
+ {{ t('usage.errors.message') }} |
+ {{ t('usage.errors.platform') }} |
+ {{ t('usage.errors.time') }} |
+
+
+
+
+ | {{ row.model || '-' }} |
+
+ {{ row.key_name || '-' }}
+ {{ t('usage.errors.keyDeleted') }}
+ |
+ {{ row.inbound_endpoint || '-' }} |
+ {{ row.status_code || '-' }} |
+ {{ t('usage.errors.categories.' + row.category) }} |
+ {{ row.message || '-' }} |
+ {{ row.platform || '-' }} |
+ {{ formatDateTime(row.created_at) }} |
+
+
+ | {{ t('usage.errors.empty') }} |
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/frontend/src/composables/__tests__/useModelWhitelist.spec.ts b/frontend/src/composables/__tests__/useModelWhitelist.spec.ts
index f2e19dc1..e64250a4 100644
--- a/frontend/src/composables/__tests__/useModelWhitelist.spec.ts
+++ b/frontend/src/composables/__tests__/useModelWhitelist.spec.ts
@@ -35,6 +35,11 @@ describe('useModelWhitelist', () => {
expect(models).toContain('gemini-3-pro-image')
})
+ it('Claude 模型列表包含 Opus 4.8', () => {
+ expect(getModelsByPlatform('claude')).toContain('claude-opus-4-8')
+ expect(getModelsByPlatform('antigravity')).toContain('claude-opus-4-8')
+ })
+
it('gemini 模型列表包含原生生图模型', () => {
const models = getModelsByPlatform('gemini')
diff --git a/frontend/src/composables/useModelWhitelist.ts b/frontend/src/composables/useModelWhitelist.ts
index 1f747595..dc192180 100644
--- a/frontend/src/composables/useModelWhitelist.ts
+++ b/frontend/src/composables/useModelWhitelist.ts
@@ -29,6 +29,7 @@ export const claudeModels = [
'claude-opus-4-5-20251101',
'claude-opus-4-6',
'claude-opus-4-7',
+ 'claude-opus-4-8',
'claude-sonnet-4-6'
]
@@ -53,6 +54,7 @@ const antigravityModels = [
'claude-opus-4-6',
'claude-opus-4-6-thinking',
'claude-opus-4-7',
+ 'claude-opus-4-8',
'claude-opus-4-5-thinking',
'claude-sonnet-4-6',
'claude-sonnet-4-5',
@@ -238,6 +240,7 @@ const anthropicPresetMappings = [
{ label: 'Opus 4.5', from: 'claude-opus-4-5-20251101', to: 'claude-opus-4-5-20251101', color: 'bg-purple-100 text-purple-700 hover:bg-purple-200 dark:bg-purple-900/30 dark:text-purple-400' },
{ label: 'Opus 4.6', from: 'claude-opus-4-6', to: 'claude-opus-4-6', color: 'bg-purple-100 text-purple-700 hover:bg-purple-200 dark:bg-purple-900/30 dark:text-purple-400' },
{ label: 'Opus 4.7', from: 'claude-opus-4-7', to: 'claude-opus-4-7', color: 'bg-purple-100 text-purple-700 hover:bg-purple-200 dark:bg-purple-900/30 dark:text-purple-400' },
+ { label: 'Opus 4.8', from: 'claude-opus-4-8', to: 'claude-opus-4-8', color: 'bg-purple-100 text-purple-700 hover:bg-purple-200 dark:bg-purple-900/30 dark:text-purple-400' },
{ label: 'Haiku 3.5', from: 'claude-3-5-haiku-20241022', to: 'claude-3-5-haiku-20241022', color: 'bg-green-100 text-green-700 hover:bg-green-200 dark:bg-green-900/30 dark:text-green-400' },
{ label: 'Haiku 4.5', from: 'claude-haiku-4-5-20251001', to: 'claude-haiku-4-5-20251001', color: 'bg-emerald-100 text-emerald-700 hover:bg-emerald-200 dark:bg-emerald-900/30 dark:text-emerald-400' },
{ label: 'Opus->Sonnet', from: 'claude-opus-4-6', to: 'claude-sonnet-4-5-20250929', color: 'bg-amber-100 text-amber-700 hover:bg-amber-200 dark:bg-amber-900/30 dark:text-amber-400' }
@@ -297,13 +300,15 @@ const antigravityPresetMappings = [
{ label: 'Sonnet 4.5', from: 'claude-sonnet-4-5', to: 'claude-sonnet-4-5', color: 'bg-cyan-100 text-cyan-700 hover:bg-cyan-200 dark:bg-cyan-900/30 dark:text-cyan-400' },
{ label: 'Opus 4.6', from: 'claude-opus-4-6', to: 'claude-opus-4-6-thinking', color: 'bg-pink-100 text-pink-700 hover:bg-pink-200 dark:bg-pink-900/30 dark:text-pink-400' },
{ label: 'Opus 4.6-thinking', from: 'claude-opus-4-6-thinking', to: 'claude-opus-4-6-thinking', color: 'bg-pink-100 text-pink-700 hover:bg-pink-200 dark:bg-pink-900/30 dark:text-pink-400' },
- { label: 'Opus 4.7', from: 'claude-opus-4-7', to: 'claude-opus-4-7', color: 'bg-pink-100 text-pink-700 hover:bg-pink-200 dark:bg-pink-900/30 dark:text-pink-400' }
+ { label: 'Opus 4.7', from: 'claude-opus-4-7', to: 'claude-opus-4-7', color: 'bg-pink-100 text-pink-700 hover:bg-pink-200 dark:bg-pink-900/30 dark:text-pink-400' },
+ { label: 'Opus 4.8', from: 'claude-opus-4-8', to: 'claude-opus-4-8', color: 'bg-pink-100 text-pink-700 hover:bg-pink-200 dark:bg-pink-900/30 dark:text-pink-400' }
]
// Bedrock 预设映射(与后端 DefaultBedrockModelMapping 保持一致)
const bedrockPresetMappings = [
{ label: 'Opus 4.6', from: 'claude-opus-4-6', to: 'us.anthropic.claude-opus-4-6-v1', color: 'bg-pink-100 text-pink-700 hover:bg-pink-200 dark:bg-pink-900/30 dark:text-pink-400' },
{ label: 'Opus 4.7', from: 'claude-opus-4-7', to: 'us.anthropic.claude-opus-4-7-v1', color: 'bg-pink-100 text-pink-700 hover:bg-pink-200 dark:bg-pink-900/30 dark:text-pink-400' },
+ { label: 'Opus 4.8', from: 'claude-opus-4-8', to: 'us.anthropic.claude-opus-4-8-v1', color: 'bg-pink-100 text-pink-700 hover:bg-pink-200 dark:bg-pink-900/30 dark:text-pink-400' },
{ label: 'Sonnet 4.6', from: 'claude-sonnet-4-6', to: 'us.anthropic.claude-sonnet-4-6', color: 'bg-cyan-100 text-cyan-700 hover:bg-cyan-200 dark:bg-cyan-900/30 dark:text-cyan-400' },
{ label: 'Opus 4.5', from: 'claude-opus-4-5-thinking', to: 'us.anthropic.claude-opus-4-5-20251101-v1:0', color: 'bg-pink-100 text-pink-700 hover:bg-pink-200 dark:bg-pink-900/30 dark:text-pink-400' },
{ label: 'Sonnet 4.5', from: 'claude-sonnet-4-5', to: 'us.anthropic.claude-sonnet-4-5-20250929-v1:0', color: 'bg-cyan-100 text-cyan-700 hover:bg-cyan-200 dark:bg-cyan-900/30 dark:text-cyan-400' },
diff --git a/frontend/src/i18n/__tests__/riskControlLocales.spec.ts b/frontend/src/i18n/__tests__/riskControlLocales.spec.ts
new file mode 100644
index 00000000..eab94fe6
--- /dev/null
+++ b/frontend/src/i18n/__tests__/riskControlLocales.spec.ts
@@ -0,0 +1,24 @@
+import { describe, expect, it } from 'vitest'
+
+import en from '../locales/en'
+import zh from '../locales/zh'
+
+describe('risk control locale copy', () => {
+ it('describes worker runtime as audit and pre-block record processing', () => {
+ expect(zh.admin.riskControl.workerStatusHint).toContain('前置拦截记录任务')
+ expect(zh.admin.riskControl.workerStatusHint).not.toContain('异步观察任务')
+ expect(en.admin.riskControl.workerStatusHint).toContain('pre-block record tasks')
+ expect(en.admin.riskControl.workerStatusHint).not.toContain('observation tasks')
+ })
+
+ it('keeps pre-block audit key summary aware of async worker load', () => {
+ expect(zh.admin.riskControl.preBlockAPIKeyLoadSummary).toContain('worker:{workerActive} / {workerTotal}')
+ expect(en.admin.riskControl.preBlockAPIKeyLoadSummary).toContain('worker: {workerActive} / {workerTotal}')
+ })
+
+ it('does not describe pre-block audit key polling as bypassing the worker pool', () => {
+ expect(zh.admin.riskControl.preBlockAPIKeyLoadHint).toBe('同步前置拦截直接轮询可用审核 Key。')
+ expect(zh.admin.riskControl.preBlockAPIKeyLoadHint).not.toContain('Worker 池')
+ expect(en.admin.riskControl.preBlockAPIKeyLoadHint).not.toContain('worker pool')
+ })
+})
diff --git a/frontend/src/i18n/locales/en.ts b/frontend/src/i18n/locales/en.ts
index 386fe73c..3ee4921b 100644
--- a/frontend/src/i18n/locales/en.ts
+++ b/frontend/src/i18n/locales/en.ts
@@ -933,7 +933,26 @@ export default {
exportExcelSuccess: 'Usage data exported successfully (Excel format)',
exportExcelFailed: 'Failed to export usage data',
imageUnit: ' images',
- userAgent: 'User-Agent'
+ userAgent: 'User-Agent',
+ tabs: { usage: 'Usage', errors: 'Error Requests' },
+ errors: {
+ time: 'Time', model: 'Model', endpoint: 'Endpoint', status: 'Status',
+ category: 'Category', platform: 'Platform', message: 'Message',
+ keyName: 'Key Name', keyDeleted: 'Deleted', allKeys: 'All keys',
+ modelPlaceholder: 'Search model', allCategories: 'All categories',
+ empty: 'No error requests', failedToLoad: 'Failed to load error requests',
+ categories: {
+ auth: 'Auth failed', rate_limit: 'Rate limited', quota: 'Balance/Subscription',
+ invalid_request: 'Invalid request', service_unavailable: 'Service unavailable',
+ upstream: 'Upstream error', internal: 'Platform error', other: 'Other',
+ },
+ detail: {
+ title: 'Error Request Detail',
+ responseBody: 'Response Body',
+ upstreamStatus: 'Upstream Status',
+ loadFailed: 'Failed to load detail, please try again',
+ },
+ },
},
// Shared keys for channel monitor (admin + user views)
@@ -2179,6 +2198,12 @@ export default {
finalPricePreview: 'Final per-image price preview',
notConfigured: 'Not configured'
},
+ modelsList: {
+ title: 'Custom /v1/models Model List',
+ hint: 'Only changes the /v1/models response. Whitelist model calls and account routing are unchanged.',
+ loading: 'Loading model list...',
+ empty: 'No displayable models'
+ },
claudeCode: {
title: 'Claude Code Client Restriction',
tooltip: 'When enabled, this group only allows official Claude Code clients. Non-Claude Code requests will be rejected or fallback to the specified group.',
@@ -2593,14 +2618,37 @@ export default {
modelFilterIncludeSummary: 'Applies to {count} models',
modelFilterExcludeSummary: 'Excludes {count} models',
emptyLogs: 'No audit records',
+ preBlockSyncStatus: 'Pre-Block Sync Status',
+ preBlockSyncHint: 'Live counters for the synchronous moderation path, excluding async record tasks.',
+ preBlockActive: 'Sync Processing',
+ preBlockActiveHint: 'Currently checking',
+ preBlockChecked: 'Checked',
+ preBlockCheckedHint: 'Entered pre-block path',
+ preBlockAllowed: 'Allowed',
+ preBlockAllowedHint: 'No block triggered',
+ preBlockBlocked: 'Blocked',
+ preBlockBlockedHint: 'Rejected after hit',
+ preBlockErrors: 'Audit Errors',
+ preBlockErrorsHint: 'Failed or no usable key',
+ preBlockAvgLatency: 'Avg Latency',
+ preBlockAvgLatencyHint: 'Synchronous path average',
+ preBlockAPIKeyLoad: 'Audit Key Load',
+ preBlockAPIKeyLoadHint: 'Synchronous pre-block checks round-robin usable audit keys directly.',
+ preBlockAPIKeyLoadSummary: 'Sync active {active} / usable keys {available}, {total} total, worker: {workerActive} / {workerTotal}',
+ preBlockAPIKeyTotals: 'Total {total}, success {success}, errors {errors}',
+ preBlockAPIKeyLoadEmpty: 'No audit key load data yet',
+ preBlockKeyActiveShort: 'Active',
+ preBlockKeyTotalShort: 'Total',
+ preBlockKeyAvgShort: 'Avg',
+ preBlockKeyLastShort: 'Last',
workerStatus: 'Worker Runtime',
- workerStatusHint: 'Queue and worker pool status for asynchronous observation tasks.',
+ workerStatusHint: 'Queue and worker pool status for async audit tasks and pre-block record tasks, excluding synchronous pre-block checks.',
workerPool: 'Worker Pool',
workerPoolMeta: '{active} processing, {idle} idle and ready, {total} total',
queueUsage: 'Queue Usage',
activeWorkers: 'Processing',
idleWorkers: 'Idle Ready',
- workerActive: 'Processing an asynchronous audit task',
+ workerActive: 'Processing an async audit or record task',
workerIdle: 'Started, idle and ready',
workerDisabled: 'Risk control or content audit is disabled',
processed: 'Processed',
@@ -3072,9 +3120,11 @@ export default {
usageWindows: 'Usage Windows',
proxy: 'Proxy',
lastUsed: 'Last Used',
+ createdAt: 'Created',
expiresAt: 'Expires At',
actions: 'Actions'
},
+ usageWindowsHint: '"5h / 7d" are the upstream account\'s official rolling usage windows (e.g. OpenAI ChatGPT, Claude). They are imposed by the upstream provider on the account itself — not configured by sub2api, and unrelated to the models you map. Usage resets automatically once each window rolls over, and the limit cannot be lifted from within sub2api.',
allPrivacyModes: 'All Privacy States',
privacyUnset: 'Unset',
privacyTrainingOff: 'Training data sharing disabled',
@@ -3319,10 +3369,21 @@ export default {
'Automatic passthrough is currently enabled: it only affects HTTP passthrough and does not disable WS mode.',
responsesMode: 'Responses API support',
responsesModeDesc:
- 'Only applies to OpenAI API Key accounts. Auto follows probe results; force modes override probing.',
+ 'Only applies to the OpenAI API Key text forwarding path. Auto follows probe results; force modes override probing.',
responsesModeAuto: 'Auto',
responsesModeForceResponses: 'Force Responses',
responsesModeForceChatCompletions: 'Force Chat Completions',
+ responsesModeTextDisabledHint:
+ 'Not applicable when the Responses / Chat Completions endpoint is not enabled.',
+ endpointCapabilities: 'Endpoint capabilities',
+ endpointCapabilitiesDesc:
+ 'Used by account routing. The text endpoint follows the Responses API support setting above and is shown as Responses, Chat Completions, or auto mode; Embeddings independently controls /v1/embeddings.',
+ capabilityResponses: 'Responses',
+ capabilityTextAuto: 'Responses / Chat Completions (Auto)',
+ capabilityResponsesAuto: 'Responses (auto probe)',
+ capabilityChatCompletions: 'Chat Completions',
+ capabilityChatCompletionsAuto: 'Chat Completions (auto probe)',
+ capabilityEmbeddings: 'Embeddings',
responsesStatusAutoSupported: 'Auto probe: Responses',
responsesStatusAutoUnsupported: 'Auto probe: Chat Completions',
responsesStatusAutoUnknown: 'Auto probe: unknown',
@@ -3331,6 +3392,9 @@ export default {
codexCLIOnly: 'Codex official clients only',
codexCLIOnlyDesc:
'Only applies to OpenAI OAuth. When enabled, only Codex official client families are allowed; when disabled, the gateway bypasses this restriction and keeps existing behavior.',
+ codexCLIOnlyAllowClaudeCode: "Also allow Claude Code's Codex plugin",
+ codexCLIOnlyAllowClaudeCodeDesc:
+ 'Only takes effect when the switch above is on. Additionally allows requests from the Claude Code Codex plugin (exact match on originator=Claude Code) without weakening blocking of other non-official clients.',
codexImageGenerationBridge: 'Codex image-generation bridge',
codexImageGenerationBridgeDesc:
'Account policy takes precedence over channel and global settings. Only controls whether Codex requests through the /responses text endpoint receive the image_generation tool; standalone image-generation endpoints are unaffected.',
@@ -3410,6 +3474,9 @@ export default {
poolModeRetryCount: 'Same-Account Retries',
poolModeRetryCountHint:
'Only applies in pool mode. Use 0 to disable in-place retry. Default {default}, maximum {max}.',
+ poolModeRetryStatusCodes: 'Retry Status Codes',
+ poolModeRetryStatusCodesHint:
+ 'Comma-separated HTTP status codes (100-599) that trigger same-account retry in pool mode. Leave blank to use defaults ({default}).',
customErrorCodes: 'Custom Error Codes',
customErrorCodesHint: 'Only stop scheduling for selected error codes',
customErrorCodesWarning:
@@ -3428,6 +3495,12 @@ export default {
'When enabled, warmup requests like title generation will return mock responses without consuming upstream tokens',
autoPauseOnExpired: 'Auto Pause On Expired',
autoPauseOnExpiredDesc: 'When enabled, the account will auto pause scheduling after it expires',
+ autoPause5hThreshold: '5h Usage Threshold (%)',
+ autoPause7dThreshold: '7d Usage Threshold (%)',
+ autoPauseThresholdHint: 'Leave empty or set 0 to use the global default threshold (configured in Ops settings); set a value to override the global default. Reaching the threshold only skips the account during scheduling and does not modify schedulable.',
+ autoPause5hDisabled: 'Disable 5h auto-pause',
+ autoPause7dDisabled: 'Disable 7d auto-pause',
+ autoPauseDisabledHint: 'When enabled, this account is never auto-paused (even if a global default threshold is configured).',
// Quota control (Anthropic OAuth/SetupToken only)
quotaControl: {
title: 'Quota Control',
@@ -4471,6 +4544,7 @@ export default {
ipAddress: 'IP',
clickToViewBalance: 'Click to view balance history',
failedToLoadUser: 'Failed to load user info',
+ userDeletedBadge: 'Deleted',
cleanup: {
button: 'Cleanup',
title: 'Cleanup Usage Records',
@@ -4702,6 +4776,8 @@ export default {
group: 'Group',
user: 'User',
userId: 'User ID',
+ apiKey: 'API Key',
+ keyDeletedBadge: 'Key Deleted',
account: 'Account',
accountId: 'Account ID',
status: 'Status',
@@ -4828,7 +4904,11 @@ export default {
suggestRequest: 'Client request error: ask customer to fix request parameters',
suggestAuth: 'Auth failed: verify API key/credentials',
suggestPlatform: 'Platform error: prioritize investigation and fix',
- suggestGeneric: 'See details for more context'
+ suggestGeneric: 'See details for more context',
+ apiKeyPrefix: 'Key Prefix',
+ attemptedKeyPrefix: 'Attempted Key Prefix',
+ deletedKeyOwner: 'Deleted Key Owner',
+ keyDeletedBadge: 'Key Deleted'
},
requestDetails: {
title: 'Request Details',
@@ -5143,6 +5223,11 @@ export default {
aggregation: 'Pre-aggregation Tasks',
enableAggregation: 'Enable Pre-aggregation',
aggregationHint: 'Pre-aggregation improves query performance for long time windows',
+ openaiQuotaAutoPause: 'OpenAI Account Quota Auto-pause',
+ openaiQuotaAutoPauseHint: 'When an OpenAI account reaches its 5h / 7d usage threshold, the scheduler skips it automatically and resumes once the window rolls over. Per-account thresholds take precedence over this global default.',
+ openaiQuotaAutoPauseDefault5h: 'Default 5h usage threshold (%)',
+ openaiQuotaAutoPauseDefault7d: 'Default 7d usage threshold (%)',
+ openaiQuotaAutoPauseThresholdHint: 'Value 0-100; leave blank or 0 to disable the global default threshold.',
errorFiltering: 'Error Filtering',
ignoreCountTokensErrors: 'Ignore count_tokens errors',
ignoreCountTokensErrorsHint: 'When enabled, errors from count_tokens requests will not be written to the error log.',
@@ -5173,7 +5258,8 @@ export default {
slaMinPercentRange: 'SLA minimum percentage must be between 0 and 100',
ttftP99MaxRange: 'TTFT P99 maximum must be a number ≥ 0',
requestErrorRateMaxRange: 'Request error rate maximum must be between 0 and 100',
- upstreamErrorRateMaxRange: 'Upstream error rate maximum must be between 0 and 100'
+ upstreamErrorRateMaxRange: 'Upstream error rate maximum must be between 0 and 100',
+ openaiQuotaAutoPauseRange: 'OpenAI quota auto-pause threshold must be between 0 and 100'
}
},
concurrency: {
@@ -5567,6 +5653,9 @@ export default {
openaiCodexUserAgent: 'OpenAI Codex UA',
openaiCodexUserAgentPlaceholder: 'codex-tui/0.125.0 (Ubuntu 22.4.0; x86_64) xterm-256color (codex-tui; 0.125.0)',
openaiCodexUserAgentHint: 'Used to bypass Cloudflare browser-UA challenges on the OpenAI upstream. Only applies when the client User-Agent is detected as a browser (Mozilla/...). Leave empty to use the built-in default.',
+ openaiAllowClaudeCodeCodexPlugin: "Allow using the Codex plugin in Claude Code",
+ openaiAllowClaudeCodeCodexPluginDesc:
+ "Global switch; only affects OpenAI OAuth accounts that have 'Codex official clients only' enabled. When on, all such accounts additionally allow requests from the Claude Code Codex plugin (exact match on originator=Claude Code) without per-account config; upstream requests remain pass-through.",
},
webSearchEmulation: {
title: 'Web Search Emulation',
@@ -6266,6 +6355,14 @@ export default {
title: 'OpenAI experimental scheduler policy',
description: "Disabled by default. When enabled, this only changes the gateway's experimental account-selection policy for OpenAI traffic; it does not indicate an upstream OpenAI capability."
},
+ usageRecords: {
+ title: 'Usage Records',
+ description: 'Settings for usage and failed-request records visible to end users.',
+ },
+ user_error_view: {
+ label: 'Allow users to view their own error requests',
+ description: 'When enabled, users can see a redacted view of their failed requests on the usage page (no internal/upstream details). Requires ops monitoring enabled to have data.',
+ },
saveSettings: 'Save Settings',
saving: 'Saving...',
settingsSaved: 'Settings saved successfully',
diff --git a/frontend/src/i18n/locales/zh.ts b/frontend/src/i18n/locales/zh.ts
index 7f4ed46e..53a715b4 100644
--- a/frontend/src/i18n/locales/zh.ts
+++ b/frontend/src/i18n/locales/zh.ts
@@ -937,7 +937,26 @@ export default {
exportExcelSuccess: '使用数据导出成功(Excel格式)',
exportExcelFailed: '使用数据导出失败',
imageUnit: '张',
- userAgent: 'User-Agent'
+ userAgent: 'User-Agent',
+ tabs: { usage: '用量明细', errors: '错误请求' },
+ errors: {
+ time: '时间', model: '模型', endpoint: '端点', status: '状态码',
+ category: '分类', platform: '平台', message: '错误信息',
+ keyName: 'Key 名称', keyDeleted: '已删除', allKeys: '全部 Key',
+ modelPlaceholder: '搜索模型', allCategories: '全部分类',
+ empty: '暂无错误请求', failedToLoad: '加载错误请求失败',
+ categories: {
+ auth: '认证失败', rate_limit: '限流', quota: '余额/订阅',
+ invalid_request: '参数错误', service_unavailable: '服务暂时不可用',
+ upstream: '上游错误', internal: '平台错误', other: '其他',
+ },
+ detail: {
+ title: '错误请求详情',
+ responseBody: '上游响应内容',
+ upstreamStatus: '上游状态码',
+ loadFailed: '加载详情失败,请稍后重试',
+ },
+ },
},
// Shared keys for channel monitor (admin + user views)
@@ -2262,6 +2281,12 @@ export default {
finalPricePreview: '最终单张价格预览',
notConfigured: '未配置'
},
+ modelsList: {
+ title: '自定义 /v1/models 模型列表',
+ hint: '仅影响 /v1/models 展示结果,不影响白名单模型调用和账号调度。',
+ loading: '正在加载模型列表...',
+ empty: '暂无可展示模型'
+ },
claudeCode: {
title: 'Claude Code 客户端限制',
tooltip:
@@ -2670,14 +2695,37 @@ export default {
modelFilterIncludeSummary: '仅 {count} 个模型生效',
modelFilterExcludeSummary: '排除 {count} 个模型',
emptyLogs: '暂无审核记录',
+ preBlockSyncStatus: '前置拦截同步状态',
+ preBlockSyncHint: '同步审核链路的实时计数,不包含异步写记录任务。',
+ preBlockActive: '同步处理中',
+ preBlockActiveHint: '当前正在审核',
+ preBlockChecked: '已检查',
+ preBlockCheckedHint: '进入前置拦截链路',
+ preBlockAllowed: '已放行',
+ preBlockAllowedHint: '未触发拦截',
+ preBlockBlocked: '已拦截',
+ preBlockBlockedHint: '命中后拒绝请求',
+ preBlockErrors: '审核异常',
+ preBlockErrorsHint: '失败或无可用 Key',
+ preBlockAvgLatency: '平均耗时',
+ preBlockAvgLatencyHint: '同步链路平均值',
+ preBlockAPIKeyLoad: '审核 Key 负载',
+ preBlockAPIKeyLoadHint: '同步前置拦截直接轮询可用审核 Key。',
+ preBlockAPIKeyLoadSummary: '同步并发 {active} / 可用 Key {available},累计 {total} 次,worker:{workerActive} / {workerTotal}',
+ preBlockAPIKeyTotals: '累计 {total},成功 {success},异常 {errors}',
+ preBlockAPIKeyLoadEmpty: '暂无审核 Key 负载数据',
+ preBlockKeyActiveShort: '并发',
+ preBlockKeyTotalShort: '累计',
+ preBlockKeyAvgShort: '平均',
+ preBlockKeyLastShort: '最近',
workerStatus: 'Worker 运行状态',
- workerStatusHint: '异步观察任务的队列和 worker 池状态。',
+ workerStatusHint: '异步审计任务和前置拦截记录任务的队列与 Worker 池状态,不包含同步前置拦截审核请求。',
workerPool: 'Worker 池',
workerPoolMeta: '{active} 个处理中,{idle} 个空闲可用,共 {total} 个',
queueUsage: '队列占用',
activeWorkers: '处理中',
idleWorkers: '空闲可用',
- workerActive: '正在处理异步审计任务',
+ workerActive: '正在处理异步审计或记录任务',
workerIdle: '已启动,当前空闲可用',
workerDisabled: '风控或内容审计未启用',
processed: '已处理',
@@ -3110,9 +3158,11 @@ export default {
usageWindows: '用量窗口',
proxy: '代理',
lastUsed: '最近使用',
+ createdAt: '创建时间',
expiresAt: '过期时间',
actions: '操作'
},
+ usageWindowsHint: '“5h / 7d”是上游账号(如 OpenAI ChatGPT、Claude)官方的滚动用量窗口限制,由上游对账号设定,并非 sub2api 配置,也与你映射的模型无关。窗口滚动到期后用量会自动重置,无法在 sub2api 端解除该限制。',
allPrivacyModes: '全部Privacy状态',
privacyUnset: '未设置',
privacyTrainingOff: '已关闭训练数据共享',
@@ -3465,10 +3515,20 @@ export default {
responsesWebsocketsV2PassthroughHint: '当前已开启自动透传:仅影响 HTTP 透传链路,不影响 WS mode。',
responsesMode: 'Responses API 支持',
responsesModeDesc:
- '仅对 OpenAI API Key 生效。自动跟随探测结果,强制模式会覆盖自动探测。',
+ '仅对 OpenAI API Key 的文本转发链路生效。自动跟随探测结果,强制模式会覆盖自动探测。',
responsesModeAuto: '自动',
responsesModeForceResponses: '强制 Responses',
responsesModeForceChatCompletions: '强制 Chat Completions',
+ responsesModeTextDisabledHint: '未启用 Responses / Chat Completions 端点时,此设置不适用。',
+ endpointCapabilities: '端点能力',
+ endpointCapabilitiesDesc:
+ '用于调度筛选。文本端点会跟随上方 Responses API 支持显示为 Responses、Chat Completions 或自动模式;Embeddings 独立控制 /v1/embeddings。',
+ capabilityResponses: 'Responses',
+ capabilityTextAuto: 'Responses / Chat Completions(自动)',
+ capabilityResponsesAuto: 'Responses(自动探测)',
+ capabilityChatCompletions: 'Chat Completions',
+ capabilityChatCompletionsAuto: 'Chat Completions(自动探测)',
+ capabilityEmbeddings: 'Embeddings',
responsesStatusAutoSupported: '自动探测:Responses',
responsesStatusAutoUnsupported: '自动探测:Chat Completions',
responsesStatusAutoUnknown: '自动探测:未探测',
@@ -3476,6 +3536,8 @@ export default {
responsesStatusForcedChatCompletions: '已强制 Chat Completions',
codexCLIOnly: '仅允许 Codex 官方客户端',
codexCLIOnlyDesc: '仅对 OpenAI OAuth 生效。开启后仅允许 Codex 官方客户端家族访问;关闭后完全绕过并保持原逻辑。',
+ codexCLIOnlyAllowClaudeCode: '额外放行 Claude Code 的 Codex 插件',
+ codexCLIOnlyAllowClaudeCodeDesc: '仅在上方开关开启时生效。额外放行通过 Claude Code 的 Codex 插件发起的请求(精确匹配 originator=Claude Code),不影响对其他非官方客户端的拦截。',
codexImageGenerationBridge: 'Codex 图片生成桥接',
codexImageGenerationBridgeDesc:
'账号级策略优先于渠道和全局配置。仅控制 Codex 走 /responses 文本端点时是否注入 image_generation 工具;不影响独立图片生成接口。',
@@ -3553,6 +3615,8 @@ export default {
'启用后,上游 429/403/401 错误将自动重试而不标记账号限流或错误,适用于上游指向另一个 sub2api 实例的场景。',
poolModeRetryCount: '同账号重试次数',
poolModeRetryCountHint: '仅在池模式下生效。0 表示不原地重试;默认 {default},最大 {max}。',
+ poolModeRetryStatusCodes: '同账号重试状态码',
+ poolModeRetryStatusCodesHint: '仅在池模式下生效。以英文逗号分隔的 HTTP 状态码(100-599),命中时触发同账号重试。留空使用默认值({default})。',
customErrorCodes: '自定义错误码',
customErrorCodesHint: '仅对选中的错误码停止调度',
customErrorCodesWarning: '仅选中的错误码会停止调度,其他错误将返回 500。',
@@ -3569,6 +3633,12 @@ export default {
interceptWarmupRequestsDesc: '启用后,标题生成等预热请求将返回 mock 响应,不消耗上游 token',
autoPauseOnExpired: '过期自动暂停调度',
autoPauseOnExpiredDesc: '启用后,账号过期将自动暂停调度',
+ autoPause5hThreshold: '5h 用量阈值(%)',
+ autoPause7dThreshold: '7d 用量阈值(%)',
+ autoPauseThresholdHint: '留空或填 0 表示使用全局默认阈值(在运维设置中配置);填具体值则覆盖全局默认。达到阈值后仅在调度时跳过账号,不修改 schedulable。',
+ autoPause5hDisabled: '禁用 5h 自动暂停',
+ autoPause7dDisabled: '禁用 7d 自动暂停',
+ autoPauseDisabledHint: '开启后该账号永不进入自动暂停(即使全局默认阈值已配置)。',
// Quota control (Anthropic OAuth/SetupToken only)
quotaControl: {
title: '配额控制',
@@ -4627,6 +4697,7 @@ export default {
ipAddress: 'IP',
clickToViewBalance: '点击查看充值记录',
failedToLoadUser: '加载用户信息失败',
+ userDeletedBadge: '已删除',
cleanup: {
button: '清理',
title: '清理使用记录',
@@ -4864,6 +4935,8 @@ export default {
group: '分组',
user: '用户',
userId: '用户 ID',
+ apiKey: 'API Key',
+ keyDeletedBadge: 'Key 已删除',
account: '账号',
accountId: '账号 ID',
status: '状态码',
@@ -4990,7 +5063,11 @@ export default {
suggestRequest: '⚠️ 客户端请求错误,建议:联系客户修正请求参数 / 手动标记已解决',
suggestAuth: '⚠️ 认证失败,建议:检查 API Key 是否有效 / 联系客户更新凭证',
suggestPlatform: '🚨 平台错误,建议立即排查修复',
- suggestGeneric: '查看详情了解更多信息'
+ suggestGeneric: '查看详情了解更多信息',
+ apiKeyPrefix: 'Key 前缀',
+ attemptedKeyPrefix: '尝试的 Key 前缀',
+ deletedKeyOwner: '已删除 Key 所有者',
+ keyDeletedBadge: 'Key 已删除'
},
requestDetails: {
title: '请求明细',
@@ -5305,6 +5382,11 @@ export default {
aggregation: '预聚合任务',
enableAggregation: '启用预聚合任务',
aggregationHint: '预聚合可提升长时间窗口查询性能',
+ openaiQuotaAutoPause: 'OpenAI 账号配额自动暂停',
+ openaiQuotaAutoPauseHint: '当 OpenAI 账号 5h / 7d 用量达到阈值时,调度会自动跳过该账号;窗口滚动后自动恢复。账号级阈值优先于此全局默认值。',
+ openaiQuotaAutoPauseDefault5h: '默认 5h 用量阈值 (%)',
+ openaiQuotaAutoPauseDefault7d: '默认 7d 用量阈值 (%)',
+ openaiQuotaAutoPauseThresholdHint: '取值 0-100,留空或 0 表示不启用全局默认阈值。',
errorFiltering: '错误过滤',
ignoreCountTokensErrors: '忽略 count_tokens 错误',
ignoreCountTokensErrorsHint: '启用后,count_tokens 请求的错误将不会写入错误日志。',
@@ -5336,7 +5418,8 @@ export default {
slaMinPercentRange: 'SLA最低百分比必须在0-100之间',
ttftP99MaxRange: 'TTFT P99最大值必须大于等于0',
requestErrorRateMaxRange: '请求错误率最大值必须在0-100之间',
- upstreamErrorRateMaxRange: '上游错误率最大值必须在0-100之间'
+ upstreamErrorRateMaxRange: '上游错误率最大值必须在0-100之间',
+ openaiQuotaAutoPauseRange: 'OpenAI 配额自动暂停阈值必须在 0-100 之间'
}
},
concurrency: {
@@ -5724,6 +5807,9 @@ export default {
openaiCodexUserAgent: 'OpenAI Codex UA',
openaiCodexUserAgentPlaceholder: 'codex-tui/0.125.0 (Ubuntu 22.4.0; x86_64) xterm-256color (codex-tui; 0.125.0)',
openaiCodexUserAgentHint: '用于规避 OpenAI 上游 Cloudflare 对浏览器 UA 的访问质询。仅在检测到客户端 User-Agent 为浏览器(Mozilla/...)时生效,其他客户端原样透传。留空使用内置默认值。',
+ openaiAllowClaudeCodeCodexPlugin: '允许在 Claude Code 中使用 Codex 插件',
+ openaiAllowClaudeCodeCodexPluginDesc:
+ '全局开关,仅对已开启「仅允许 Codex 官方客户端」的 OpenAI OAuth 账号生效。开启后,所有此类账号都额外放行通过 Claude Code 的 Codex 插件发起的请求(精确匹配 originator=Claude Code),无需逐账号配置;上游请求仍保持透传。',
},
webSearchEmulation: {
title: 'Web Search 模拟',
@@ -6424,6 +6510,14 @@ export default {
title: 'OpenAI 实验调度策略',
description: '默认关闭。开启后仅影响本网关在 OpenAI 账号间的实验性调度选择逻辑,不代表上游 OpenAI 官方能力。'
},
+ usageRecords: {
+ title: '使用记录',
+ description: '与终端用户可见的用量及失败请求记录相关的设置。',
+ },
+ user_error_view: {
+ label: '允许用户查看自己的错误请求',
+ description: '开启后,用户可在用量页查看自己失败请求的精简信息(不含内部/上游错误细节)。需运维监控开启才有数据。',
+ },
saveSettings: '保存设置',
saving: '保存中...',
settingsSaved: '设置保存成功',
diff --git a/frontend/src/stores/app.ts b/frontend/src/stores/app.ts
index 2d2237fe..a01779c0 100644
--- a/frontend/src/stores/app.ts
+++ b/frontend/src/stores/app.ts
@@ -360,6 +360,7 @@ export const useAppStore = defineStore('app', () => {
available_channels_enabled: false,
risk_control_enabled: false,
affiliate_enabled: false,
+ allow_user_view_error_requests: false,
}
}
diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts
index 6bd016ac..221f6619 100644
--- a/frontend/src/types/index.ts
+++ b/frontend/src/types/index.ts
@@ -97,6 +97,7 @@ export interface User {
last_active_at?: string | null
created_at: string
updated_at: string
+ deleted_at?: string | null
}
export interface AdminUser extends User {
@@ -234,6 +235,7 @@ export interface PublicSettings {
channel_monitor_default_interval_seconds: number
available_channels_enabled: boolean
affiliate_enabled: boolean
+ allow_user_view_error_requests?: boolean
}
export interface AuthResponse {
@@ -549,11 +551,17 @@ export interface AdminGroup extends Group {
// OpenAI Messages 调度配置(仅 openai 平台使用)
default_mapped_model?: string
messages_dispatch_model_config?: OpenAIMessagesDispatchModelConfig
+ models_list_config?: ModelsListConfig
// 分组排序
sort_order: number
}
+export interface ModelsListConfig {
+ enabled: boolean
+ models: string[]
+}
+
export interface ApiKey {
id: number
user_id: number
@@ -633,6 +641,13 @@ export interface CreateGroupRequest {
fallback_group_id_on_invalid_request?: number | null
mcp_xml_inject?: boolean
supported_model_scopes?: string[]
+ models_list_config?: ModelsListConfig
+ allow_messages_dispatch?: boolean
+ default_mapped_model?: string
+ messages_dispatch_model_config?: OpenAIMessagesDispatchModelConfig
+ model_routing?: Record | null
+ model_routing_enabled?: boolean
+ rpm_limit?: number
require_oauth_only?: boolean
require_privacy_set?: boolean
// 从指定分组复制账号
@@ -661,6 +676,13 @@ export interface UpdateGroupRequest {
fallback_group_id_on_invalid_request?: number | null
mcp_xml_inject?: boolean
supported_model_scopes?: string[]
+ models_list_config?: ModelsListConfig
+ allow_messages_dispatch?: boolean
+ default_mapped_model?: string
+ messages_dispatch_model_config?: OpenAIMessagesDispatchModelConfig
+ model_routing?: Record | null
+ model_routing_enabled?: boolean
+ rpm_limit?: number
require_oauth_only?: boolean
require_privacy_set?: boolean
copy_accounts_from_group_ids?: number[]
@@ -978,6 +1000,7 @@ export interface CodexUsageSnapshot {
export type OpenAICompactMode = 'auto' | 'force_on' | 'force_off'
export type OpenAIResponsesMode = 'auto' | 'force_responses' | 'force_chat_completions'
+export type OpenAIEndpointCapability = 'chat_completions' | 'embeddings'
export interface OpenAICompactState {
openai_compact_mode?: OpenAICompactMode
@@ -1204,7 +1227,7 @@ export interface UsageLog {
request_type?: UsageRequestType
stream: boolean
openai_ws_mode?: boolean
- duration_ms: number
+ duration_ms: number | null
first_token_ms: number | null
// 图片生成字段
@@ -1571,6 +1594,36 @@ export interface ExtendSubscriptionRequest {
// ==================== Query Parameters ====================
+export interface UserErrorRequest {
+ id: number
+ created_at: string
+ model: string
+ inbound_endpoint: string
+ status_code: number
+ category: string
+ platform: string
+ message: string
+ key_name: string
+ key_deleted: boolean
+}
+
+export interface UserErrorRequestDetail extends UserErrorRequest {
+ error_body: string
+ upstream_status_code?: number
+}
+
+export interface UserErrorListParams {
+ page?: number
+ page_size?: number
+ start_date?: string
+ end_date?: string
+ timezone?: string
+ model?: string
+ status_code?: number
+ category?: string
+ api_key_id?: number
+}
+
export interface UsageQueryParams {
page?: number
page_size?: number
diff --git a/frontend/src/views/admin/AccountsView.vue b/frontend/src/views/admin/AccountsView.vue
index 51137c8b..04b46a8d 100644
--- a/frontend/src/views/admin/AccountsView.vue
+++ b/frontend/src/views/admin/AccountsView.vue
@@ -273,6 +273,12 @@
+
+
+ {{ column.label }}
+
+
+
{{ formatRelativeTime(value) }}
+
+ {{ formatDateTime(value) }}
+
{{ formatExpiresAt(value) }}
@@ -387,6 +396,7 @@ import { useTableSelection } from '@/composables/useTableSelection'
import AppLayout from '@/components/layout/AppLayout.vue'
import TablePageLayout from '@/components/layout/TablePageLayout.vue'
import DataTable from '@/components/common/DataTable.vue'
+import HelpTooltip from '@/components/common/HelpTooltip.vue'
import Pagination from '@/components/common/Pagination.vue'
import ConfirmDialog from '@/components/common/ConfirmDialog.vue'
import { CreateAccountModal, EditAccountModal, BulkEditAccountModal, SyncFromCrsModal, TempUnschedStatusModal } from '@/components/account'
@@ -509,6 +519,7 @@ const ACCOUNT_SORTABLE_KEYS = new Set([
'priority',
'rate_multiplier',
'last_used_at',
+ 'created_at',
'expires_at'
])
const loadInitialAccountSortState = (): AccountSortState => {
@@ -1127,6 +1138,7 @@ const allColumns = computed(() => {
{ key: 'priority', label: t('admin.accounts.columns.priority'), sortable: true },
{ key: 'rate_multiplier', label: t('admin.accounts.columns.billingRateMultiplier'), sortable: true },
{ key: 'last_used_at', label: t('admin.accounts.columns.lastUsed'), sortable: true },
+ { key: 'created_at', label: t('admin.accounts.columns.createdAt'), sortable: true },
{ key: 'expires_at', label: t('admin.accounts.columns.expiresAt'), sortable: true },
{ key: 'notes', label: t('admin.accounts.columns.notes'), sortable: false },
{ key: 'actions', label: t('admin.accounts.columns.actions'), sortable: false }
diff --git a/frontend/src/views/admin/GroupsView.vue b/frontend/src/views/admin/GroupsView.vue
index ebb57bd9..0b583a09 100644
--- a/frontend/src/views/admin/GroupsView.vue
+++ b/frontend/src/views/admin/GroupsView.vue
@@ -69,7 +69,7 @@
{{ t("admin.groups.sortOrder") }}
+
+
+
+
+
+ {{ t("admin.groups.modelsList.hint") }}
+
+
+
+
+
+
+
+ 已选 {{ createModelsListSelectedCount }} /
+ {{ createModelsListState.items.length }}
+
+
+
+
+
+
+
+
+ {{ t("admin.groups.modelsList.loading") }}
+
+
+ {{ t("admin.groups.modelsList.empty") }}
+
+
+
+
+ {{ item.id }}
+
+
+
+
+
+
+
+
+
+
+
+
+ {{ t("admin.groups.modelsList.hint") }}
+
+
+
+
+
+
+
+ 已选 {{ editModelsListSelectedCount }} /
+ {{ editModelsListState.items.length }}
+
+
+
+
+
+
+
+
+ {{ t("admin.groups.modelsList.loading") }}
+
+
+ {{ t("admin.groups.modelsList.empty") }}
+
+
+
+
+ {{ item.id }}
+
+
+
+
+
+
+
+
(null);
const sortableGroups = ref
([]);
const createMessagesDispatchDefaults = createDefaultMessagesDispatchFormState();
const editMessagesDispatchDefaults = createDefaultMessagesDispatchFormState();
+const createModelsListState = reactive(createInitialModelsListState());
+const editModelsListState = reactive(createInitialModelsListState());
+const createModelsListLoading = ref(false);
+const editModelsListLoading = ref(false);
+const modelsListCandidatesTracker = createModelsListCandidatesTracker();
+const createModelsListSelectedCount = computed(
+ () => createModelsListState.items.filter((item) => item.selected).length,
+);
+const editModelsListSelectedCount = computed(
+ () => editModelsListState.items.filter((item) => item.selected).length,
+);
const createForm = reactive({
name: "",
@@ -3335,6 +3561,52 @@ const removeEditRoutingRule = (rule: ModelRoutingRule) => {
editModelRoutingRules.value.splice(index, 1);
};
+const resetModelsListState = (
+ state: typeof createModelsListState,
+ config?: Parameters[0],
+) => {
+ const fresh = createInitialModelsListState(config);
+ state.enabled = fresh.enabled;
+ state.savedModels = fresh.savedModels;
+ state.items = fresh.items;
+};
+
+const loadModelsListCandidates = async (
+ mode: "create" | "edit",
+ groupID: number,
+ platform: GroupPlatform,
+) => {
+ const request = { mode, groupID, platform };
+ const requestID = modelsListCandidatesTracker.next(request);
+ const state = mode === "create" ? createModelsListState : editModelsListState;
+ const loadingRef = mode === "create" ? createModelsListLoading : editModelsListLoading;
+ loadingRef.value = true;
+ try {
+ const models = await adminAPI.groups.getModelsListCandidates(groupID, platform);
+ if (!modelsListCandidatesTracker.isCurrent(requestID, request)) {
+ return;
+ }
+ setModelsListCandidates(state, models);
+ } catch (error) {
+ if (!modelsListCandidatesTracker.isCurrent(requestID, request)) {
+ return;
+ }
+ console.error("Error loading group models list candidates:", error);
+ } finally {
+ if (modelsListCandidatesTracker.isCurrent(requestID, request)) {
+ loadingRef.value = false;
+ }
+ }
+};
+
+const moveCreateModelsListItem = (fromIndex: number, toIndex: number) => {
+ moveModelsListItem(createModelsListState, fromIndex, toIndex);
+};
+
+const moveEditModelsListItem = (fromIndex: number, toIndex: number) => {
+ moveModelsListItem(editModelsListState, fromIndex, toIndex);
+};
+
// 将 UI 格式的路由规则转换为 API 格式
const convertRoutingRulesToApiFormat = (
rules: ModelRoutingRule[],
@@ -3624,6 +3896,11 @@ const handleSort = (key: string, order: 'asc' | 'desc') => {
loadGroups();
};
+const openCreateModal = () => {
+ showCreateModal.value = true;
+ loadModelsListCandidates("create", 0, createForm.platform);
+};
+
const closeCreateModal = () => {
showCreateModal.value = false;
createModelRoutingRules.value.forEach((rule) => {
@@ -3654,6 +3931,8 @@ const closeCreateModal = () => {
createForm.supported_model_scopes = ["claude", "gemini_text", "gemini_image"];
createForm.mcp_xml_inject = true;
createForm.copy_accounts_from_group_ids = [];
+ createForm.rpm_limit = 0;
+ resetModelsListState(createModelsListState);
createModelRoutingRules.value = [];
};
@@ -3708,6 +3987,7 @@ const handleCreateGroup = async () => {
model_routing: convertRoutingRulesToApiFormat(
createModelRoutingRules.value,
),
+ models_list_config: buildModelsListConfig(createModelsListState),
supported_model_scopes: normalizeSupportedModelScopesForPlatform(
createForm.platform,
createForm.supported_model_scopes,
@@ -3794,10 +4074,12 @@ const handleEdit = async (group: AdminGroup) => {
editForm.mcp_xml_inject = group.mcp_xml_inject ?? true;
editForm.copy_accounts_from_group_ids = []; // 复制账号字段每次编辑时重置为空
editForm.rpm_limit = group.rpm_limit ?? 0;
+ resetModelsListState(editModelsListState, group.models_list_config);
// 加载模型路由规则(异步加载账号名称)
editModelRoutingRules.value = await convertApiFormatToRoutingRules(
group.model_routing,
);
+ loadModelsListCandidates("edit", group.id, group.platform);
showEditModal.value = true;
};
@@ -3811,6 +4093,7 @@ const closeEditModal = () => {
editModelRoutingRules.value = [];
editForm.copy_accounts_from_group_ids = [];
resetMessagesDispatchFormState(editForm);
+ resetModelsListState(editModelsListState);
};
const handleUpdateGroup = async () => {
@@ -3843,6 +4126,7 @@ const handleUpdateGroup = async () => {
model_routing: convertRoutingRulesToApiFormat(
editModelRoutingRules.value,
),
+ models_list_config: buildModelsListConfig(editModelsListState),
supported_model_scopes: normalizeSupportedModelScopesForPlatform(
editForm.platform,
editForm.supported_model_scopes,
@@ -3960,6 +4244,8 @@ watch(
createForm.require_oauth_only = false;
createForm.require_privacy_set = false;
}
+ resetModelsListState(createModelsListState);
+ loadModelsListCandidates("create", 0, newVal);
},
);
@@ -3976,6 +4262,10 @@ watch(
editForm.require_oauth_only = false;
editForm.require_privacy_set = false;
}
+ if (editingGroup.value) {
+ resetModelsListState(editModelsListState, editForm.platform === editingGroup.value.platform ? editingGroup.value.models_list_config : undefined);
+ loadModelsListCandidates("edit", editingGroup.value.id, newVal);
+ }
},
);
@@ -4049,6 +4339,7 @@ const saveSortOrder = async () => {
onMounted(() => {
loadGroups();
+ loadModelsListCandidates("create", 0, createForm.platform);
document.addEventListener("click", handleClickOutside);
});
diff --git a/frontend/src/views/admin/RiskControlView.vue b/frontend/src/views/admin/RiskControlView.vue
index 36a04756..b6d62767 100644
--- a/frontend/src/views/admin/RiskControlView.vue
+++ b/frontend/src/views/admin/RiskControlView.vue
@@ -53,7 +53,105 @@
-
+
+
+
+
+
{{ t('admin.riskControl.preBlockSyncStatus') }}
+
{{ t('admin.riskControl.preBlockSyncHint') }}
+
+
+ {{ modeLabel(status?.mode ?? configForm.mode) }}
+
+
+
+
+
+
+
{{ item.label }}
+
{{ item.value }}
+
{{ item.meta }}
+
+
+
+
+
+
+
+
+
{{ t('admin.riskControl.preBlockAPIKeyLoad') }}
+
+ {{ t('admin.riskControl.preBlockAPIKeyLoadHint') }}
+
+
+
+ {{ preBlockAPIKeyLoadSummaryText }}
+
+
+
+
+
+
+
+
+
+ #{{ item.index + 1 }}
+ {{ item.masked || '-' }}
+
+
+
+ {{ t('admin.riskControl.preBlockAPIKeyTotals', { total: formatNumber(item.total), success: formatNumber(item.success), errors: formatNumber(item.errors) }) }}
+
+
+
+
+
{{ t('admin.riskControl.preBlockKeyActiveShort') }}
+
{{ formatNumber(item.active) }}
+
+
+
{{ t('admin.riskControl.preBlockKeyTotalShort') }}
+
{{ formatNumber(item.total) }}
+
+
+
{{ t('admin.riskControl.preBlockKeyAvgShort') }}
+
{{ formatNumber(item.avg_latency_ms) }} ms
+
+
+
{{ t('admin.riskControl.preBlockKeyLastShort') }}
+
{{ formatNumber(item.last_latency_ms) }} ms
+
+
+
+
+
+
+
+ {{ t('admin.riskControl.preBlockAPIKeyLoadEmpty') }}
+
+
+
+
+
+
{{ t('admin.riskControl.workerStatus') }}
@@ -1013,6 +1111,7 @@ import Pagination from '@/components/common/Pagination.vue'
import ModelWhitelistSelector from '@/components/account/ModelWhitelistSelector.vue'
import { adminAPI } from '@/api/admin'
import type {
+ ContentModerationAPIKeyLoad,
ContentModerationAPIKeyStatus,
ContentModerationConfig,
ContentModerationLog,
@@ -1472,6 +1571,81 @@ const queueUsageStyle = computed(() => ({
width: queueUsagePercent.value,
}))
+const runtimeMode = computed(() => status.value?.mode ?? configForm.mode)
+
+const showPreBlockRuntimeCard = computed(() => runtimeMode.value === 'pre_block')
+
+const showWorkerRuntimeCard = computed(() => runtimeMode.value === 'observe')
+
+const preBlockMetricItems = computed(() => [
+ {
+ key: 'active',
+ label: t('admin.riskControl.preBlockActive'),
+ value: formatNumber(status.value?.pre_block_active ?? 0),
+ meta: t('admin.riskControl.preBlockActiveHint'),
+ class: 'bg-sky-50 dark:bg-sky-900/10',
+ valueClass: 'text-sky-700 dark:text-sky-300',
+ },
+ {
+ key: 'checked',
+ label: t('admin.riskControl.preBlockChecked'),
+ value: formatNumber(status.value?.pre_block_checked ?? 0),
+ meta: t('admin.riskControl.preBlockCheckedHint'),
+ class: 'bg-gray-50 dark:bg-dark-700/50',
+ valueClass: 'text-gray-900 dark:text-white',
+ },
+ {
+ key: 'allowed',
+ label: t('admin.riskControl.preBlockAllowed'),
+ value: formatNumber(status.value?.pre_block_allowed ?? 0),
+ meta: t('admin.riskControl.preBlockAllowedHint'),
+ class: 'bg-emerald-50 dark:bg-emerald-900/10',
+ valueClass: 'text-emerald-700 dark:text-emerald-300',
+ },
+ {
+ key: 'blocked',
+ label: t('admin.riskControl.preBlockBlocked'),
+ value: formatNumber(status.value?.pre_block_blocked ?? 0),
+ meta: t('admin.riskControl.preBlockBlockedHint'),
+ class: 'bg-rose-50 dark:bg-rose-900/10',
+ valueClass: 'text-rose-700 dark:text-rose-300',
+ },
+ {
+ key: 'errors',
+ label: t('admin.riskControl.preBlockErrors'),
+ value: formatNumber(status.value?.pre_block_errors ?? 0),
+ meta: t('admin.riskControl.preBlockErrorsHint'),
+ class: 'bg-amber-50 dark:bg-amber-900/10',
+ valueClass: 'text-amber-700 dark:text-amber-300',
+ },
+ {
+ key: 'latency',
+ label: t('admin.riskControl.preBlockAvgLatency'),
+ value: `${formatNumber(status.value?.pre_block_avg_latency_ms ?? 0)} ms`,
+ meta: t('admin.riskControl.preBlockAvgLatencyHint'),
+ class: 'bg-violet-50 dark:bg-violet-900/10',
+ valueClass: 'text-violet-700 dark:text-violet-300',
+ },
+])
+
+const preBlockAPIKeyLoads = computed(() => (
+ [...(status.value?.pre_block_api_key_loads ?? [])].sort((a, b) => a.index - b.index)
+))
+
+const preBlockAPIKeyMaxTotal = computed(() => Math.max(1, ...preBlockAPIKeyLoads.value.map((item) => item.total || 0)))
+
+const preBlockAPIKeyLoadSummaryText = computed(() => t('admin.riskControl.preBlockAPIKeyLoadSummary', {
+ active: formatNumber(status.value?.pre_block_api_key_active ?? 0),
+ available: formatNumber(status.value?.pre_block_api_key_available_count ?? 0),
+ total: formatNumber(status.value?.pre_block_api_key_total_calls ?? 0),
+ workerActive: formatNumber(status.value?.active_workers ?? 0),
+ workerTotal: formatNumber(status.value?.worker_count ?? configForm.worker_count),
+}))
+
+function preBlockAPIKeyLoadWidth(total: number): string {
+ return `${Math.min(100, Math.max(0, (total / preBlockAPIKeyMaxTotal.value) * 100)).toFixed(1)}%`
+}
+
const workerSlots = computed(() => {
const total = Math.max(0, status.value?.worker_count ?? configForm.worker_count)
const active = Math.max(0, status.value?.active_workers ?? 0)
diff --git a/frontend/src/views/admin/SettingsView.vue b/frontend/src/views/admin/SettingsView.vue
index f3d79368..959e6941 100644
--- a/frontend/src/views/admin/SettingsView.vue
+++ b/frontend/src/views/admin/SettingsView.vue
@@ -3948,6 +3948,19 @@
}}
+
+
+
+
+
+
+ {{ t("admin.settings.gatewayForwarding.openaiAllowClaudeCodeCodexPluginDesc") }}
+
+
+
+
@@ -4385,6 +4398,35 @@
+
+
+
+
+
+ {{ t('admin.settings.usageRecords.title') }}
+
+
+ {{ t('admin.settings.usageRecords.description') }}
+
+
+
+
+
+
+
+
+ {{ t('admin.settings.user_error_view.description') }}
+
+
+
+
+
+
@@ -7183,6 +7225,7 @@ const form = reactive({
rewrite_message_cache_control: false,
antigravity_user_agent_version: "",
openai_codex_user_agent: "",
+ openai_allow_claude_code_codex_plugin: false,
// 余额、订阅到期与账号限额通知
balance_low_notify_enabled: false,
balance_low_notify_threshold: 0,
@@ -7197,6 +7240,8 @@ const form = reactive({
available_channels_enabled: false,
// Affiliate (邀请返利) feature switch
affiliate_enabled: false,
+ // Allow user view error requests
+ allow_user_view_error_requests: false,
});
const authSourceDefaults = reactive(
@@ -8289,6 +8334,7 @@ async function saveSettings() {
form.antigravity_user_agent_version?.trim() || "",
openai_codex_user_agent:
form.openai_codex_user_agent?.trim() || "",
+ openai_allow_claude_code_codex_plugin: form.openai_allow_claude_code_codex_plugin,
// Payment configuration
payment_enabled: form.payment_enabled,
risk_control_enabled: form.risk_control_enabled,
@@ -8338,6 +8384,7 @@ async function saveSettings() {
available_channels_enabled: form.available_channels_enabled,
// Affiliate (邀请返利) feature switch
affiliate_enabled: form.affiliate_enabled,
+ allow_user_view_error_requests: form.allow_user_view_error_requests,
};
// 仅当 openai_fast_policy_settings 已成功从后端加载时才回写,
diff --git a/frontend/src/views/admin/UsageView.vue b/frontend/src/views/admin/UsageView.vue
index 495ca7ad..4fb2caac 100644
--- a/frontend/src/views/admin/UsageView.vue
+++ b/frontend/src/views/admin/UsageView.vue
@@ -64,7 +64,7 @@
-
+
@@ -144,6 +163,10 @@ import UsageStatsCards from '@/components/admin/usage/UsageStatsCards.vue'; impo
import UsageTable from '@/components/admin/usage/UsageTable.vue'; import UsageExportProgress from '@/components/admin/usage/UsageExportProgress.vue'
import UsageCleanupDialog from '@/components/admin/usage/UsageCleanupDialog.vue'
import UserBalanceHistoryModal from '@/components/admin/user/UserBalanceHistoryModal.vue'
+import OpsErrorLogTable from '@/views/admin/ops/components/OpsErrorLogTable.vue'
+import OpsErrorDetailModal from '@/views/admin/ops/components/OpsErrorDetailModal.vue'
+import { listErrorLogs } from '@/api/admin/ops'
+import type { OpsErrorLog } from '@/api/admin/ops'
import ModelDistributionChart from '@/components/charts/ModelDistributionChart.vue'; import GroupDistributionChart from '@/components/charts/GroupDistributionChart.vue'; import TokenUsageTrend from '@/components/charts/TokenUsageTrend.vue'
import EndpointDistributionChart from '@/components/charts/EndpointDistributionChart.vue'
import Icon from '@/components/icons/Icon.vue'
@@ -192,9 +215,13 @@ const breakdownFilters = computed(() => {
return f
})
+const modelNameOptions = computed(() =>
+ Array.from(new Set(requestedModelStats.value.map((m) => m.model).filter(Boolean))).sort()
+)
+
const handleUserClick = async (userId: number) => {
try {
- const user = await adminAPI.users.getById(userId)
+ const user = await adminAPI.users.getById(userId, true)
balanceHistoryUser.value = user
showBalanceHistoryModal.value = true
} catch {
@@ -306,13 +333,17 @@ const loadLogs = async () => {
if(!c.signal.aborted) { usageLogs.value = res.items; pagination.total = res.total }
} catch (error: any) { if(error?.name !== 'AbortError') console.error('Failed to load usage logs:', error) } finally { if(abortController === c) loading.value = false }
}
-const loadStats = async () => {
+const loadStats = async (force = false) => {
const seq = ++statsReqSeq
endpointStatsLoading.value = true
try {
const requestType = filters.value.request_type
const legacyStream = requestType ? requestTypeToLegacyStream(requestType) : filters.value.stream
- const s = await adminAPI.usage.getStats({ ...filters.value, stream: legacyStream === null ? undefined : legacyStream })
+ const s = await adminAPI.usage.getStats({
+ ...filters.value,
+ stream: legacyStream === null ? undefined : legacyStream,
+ ...(force ? { nocache: 1 } : {}),
+ })
if (seq !== statsReqSeq) return
usageStats.value = s
inboundEndpointStats.value = s.endpoints || []
@@ -329,10 +360,8 @@ const loadStats = async () => {
}
}
-const resetModelStatsCache = () => {
- requestedModelStats.value = []
- upstreamModelStats.value = []
- mappingModelStats.value = []
+// 失效模型统计缓存:仅标记需要重取,保留旧数据直到新数据到达(避免刷新时图表闪空)。
+const invalidateModelStatsCache = () => {
loadedModelSources.requested = false
loadedModelSources.upstream = false
loadedModelSources.mapping = false
@@ -421,18 +450,25 @@ const loadChartData = async () => {
}
const applyFilters = () => {
pagination.page = 1
- resetModelStatsCache()
+ invalidateModelStatsCache()
loadLogs()
loadStats()
loadModelStats(modelDistributionSource.value, true)
loadChartData()
+ errPage.value = 1
+ if (activeTab.value === 'errors') {
+ loadAdminErrors()
+ } else {
+ errRows.value = []
+ }
}
const refreshData = () => {
- resetModelStatsCache()
+ invalidateModelStatsCache()
loadLogs()
- loadStats()
+ loadStats(true)
loadModelStats(modelDistributionSource.value, true)
loadChartData()
+ if (activeTab.value === 'errors') loadAdminErrors()
}
const resetFilters = () => {
const range = getLast24HoursRangeDates()
@@ -586,6 +622,50 @@ const loadSavedColumns = () => {
}
}
+// Error tab state
+const activeTab = ref<'usage' | 'errors'>('usage')
+const errRows = ref([])
+const errLoading = ref(false)
+const errPage = ref(1)
+const errPageSize = ref(20)
+const errTotal = ref(0)
+const showErrorModal = ref(false)
+const selectedErrorId = ref(null)
+
+// 注意:'YYYY-MM-DDT00:00:00' 无时区后缀,按本地时区解析后再转 UTC——与页面其它日期处理语义一致,刻意如此,勿改成 'T00:00:00Z'
+const toRFC3339 = (d: string | undefined, endOfDay = false): string | undefined =>
+ d ? new Date(d + (endOfDay ? 'T23:59:59.999' : 'T00:00:00')).toISOString() : undefined
+
+const loadAdminErrors = async () => {
+ errLoading.value = true
+ try {
+ const resp = await listErrorLogs({
+ page: errPage.value,
+ page_size: errPageSize.value,
+ view: 'all',
+ start_time: toRFC3339(filters.value.start_date),
+ end_time: toRFC3339(filters.value.end_date, true),
+ user_id: filters.value.user_id ?? undefined,
+ api_key_id: filters.value.api_key_id ?? undefined,
+ account_id: filters.value.account_id ?? undefined,
+ group_id: filters.value.group_id ?? undefined,
+ model: filters.value.model || undefined,
+ })
+ errRows.value = resp.items
+ errTotal.value = resp.total
+ } catch (error) {
+ console.error('Failed to load admin errors:', error)
+ appStore.showError(t('usage.errors.failedToLoad'))
+ } finally {
+ errLoading.value = false
+ }
+}
+
+const onErrPage = (p: number) => { errPage.value = p; loadAdminErrors() }
+const onErrPageSize = (s: number) => { errPageSize.value = s; errPage.value = 1; loadAdminErrors() }
+const openError = (id: number) => { selectedErrorId.value = id; showErrorModal.value = true }
+const switchToErrorsTab = () => { activeTab.value = 'errors'; if (errRows.value.length === 0) loadAdminErrors() }
+
const showColumnDropdown = ref(false)
const columnDropdownRef = ref(null)
@@ -611,4 +691,6 @@ onUnmounted(() => { abortController?.abort(); exportAbortController?.abort(); do
watch(modelDistributionSource, (source) => {
void loadModelStats(source)
})
+
+defineExpose({ requestedModelStats, refreshData })
diff --git a/frontend/src/views/admin/__tests__/AccountsView.bulkEdit.spec.ts b/frontend/src/views/admin/__tests__/AccountsView.bulkEdit.spec.ts
index 112baf22..ed53a9e5 100644
--- a/frontend/src/views/admin/__tests__/AccountsView.bulkEdit.spec.ts
+++ b/frontend/src/views/admin/__tests__/AccountsView.bulkEdit.spec.ts
@@ -63,7 +63,14 @@ vi.mock('vue-i18n', async () => {
const DataTableStub = {
props: ['columns', 'data'],
- template: ''
+ template: `
+
+
{{ column.key }}
+
+
+
+
+ `
}
const AccountBulkActionsBarStub = {
@@ -149,4 +156,72 @@ describe('admin AccountsView bulk edit scope', () => {
expect(wrapper.get('[data-test="bulk-edit-modal"]').attributes('data-show')).toBe('true')
expect(wrapper.get('[data-test="bulk-edit-modal"]').attributes('data-target-mode')).toBe('filtered')
})
+
+ it('renders the created_at column by default', async () => {
+ listAccounts.mockResolvedValue({
+ items: [
+ {
+ id: 1,
+ name: 'test-account',
+ platform: 'anthropic',
+ type: 'oauth',
+ status: 'active',
+ schedulable: true,
+ created_at: '2026-03-07T10:00:00Z',
+ updated_at: '2026-03-07T10:00:00Z'
+ }
+ ],
+ total: 1,
+ page: 1,
+ page_size: 20,
+ pages: 1
+ })
+
+ const wrapper = mount(AccountsView, {
+ global: {
+ stubs: {
+ AppLayout: { template: '
' },
+ TablePageLayout: {
+ template: '
'
+ },
+ DataTable: DataTableStub,
+ Pagination: true,
+ ConfirmDialog: true,
+ AccountTableActions: { template: '
' },
+ AccountTableFilters: { template: '' },
+ AccountBulkActionsBar: AccountBulkActionsBarStub,
+ AccountActionMenu: true,
+ ImportDataModal: true,
+ ReAuthAccountModal: true,
+ AccountTestModal: true,
+ AccountStatsModal: true,
+ ScheduledTestsPanel: true,
+ SyncFromCrsModal: true,
+ TempUnschedStatusModal: true,
+ ErrorPassthroughRulesModal: true,
+ TLSFingerprintProfilesModal: true,
+ CreateAccountModal: true,
+ EditAccountModal: true,
+ BulkEditAccountModal: BulkEditAccountModalStub,
+ PlatformTypeBadge: true,
+ AccountCapacityCell: true,
+ AccountStatusIndicator: true,
+ AccountTodayStatsCell: true,
+ AccountGroupsCell: true,
+ AccountUsageCell: true,
+ Icon: true
+ }
+ }
+ })
+
+ await flushPromises()
+
+ const columnKeys = wrapper.findAll('[data-test="column-key"]').map(node => node.text())
+ expect(columnKeys).toContain('created_at')
+ const columns = wrapper.getComponent(DataTableStub).props('columns') as Array<{ key: string; label: string; sortable: boolean }>
+ expect(columns.find(column => column.key === 'created_at')).toMatchObject({
+ label: 'admin.accounts.columns.createdAt',
+ sortable: true
+ })
+ })
})
diff --git a/frontend/src/views/admin/__tests__/AccountsView.usageWindowsHint.spec.ts b/frontend/src/views/admin/__tests__/AccountsView.usageWindowsHint.spec.ts
new file mode 100644
index 00000000..81e7d87e
--- /dev/null
+++ b/frontend/src/views/admin/__tests__/AccountsView.usageWindowsHint.spec.ts
@@ -0,0 +1,164 @@
+import { beforeEach, describe, expect, it, vi } from 'vitest'
+import { flushPromises, mount } from '@vue/test-utils'
+
+import AccountsView from '../AccountsView.vue'
+
+const {
+ listAccounts,
+ listWithEtag,
+ getBatchTodayStats,
+ getAllProxies,
+ getAllGroups
+} = vi.hoisted(() => ({
+ listAccounts: vi.fn(),
+ listWithEtag: vi.fn(),
+ getBatchTodayStats: vi.fn(),
+ getAllProxies: vi.fn(),
+ getAllGroups: vi.fn()
+}))
+
+vi.mock('@/api/admin', () => ({
+ adminAPI: {
+ accounts: {
+ list: listAccounts,
+ listWithEtag,
+ getBatchTodayStats,
+ delete: vi.fn(),
+ batchClearError: vi.fn(),
+ batchRefresh: vi.fn(),
+ toggleSchedulable: vi.fn()
+ },
+ proxies: {
+ getAll: getAllProxies
+ },
+ groups: {
+ getAll: getAllGroups
+ }
+ }
+}))
+
+vi.mock('@/stores/app', () => ({
+ useAppStore: () => ({
+ showError: vi.fn(),
+ showSuccess: vi.fn(),
+ showInfo: vi.fn()
+ })
+}))
+
+vi.mock('@/stores/auth', () => ({
+ useAuthStore: () => ({
+ token: 'test-token'
+ })
+}))
+
+vi.mock('vue-i18n', async () => {
+ const actual = await vi.importActual('vue-i18n')
+ return {
+ ...actual,
+ useI18n: () => ({
+ t: (key: string) => key
+ })
+ }
+})
+
+// Render the per-column header slots so we can assert the usage-window header hint.
+const DataTableStub = {
+ props: ['columns', 'data'],
+ template: `
+
+ `
+}
+
+// Expose the content passed to HelpTooltip without dealing with its .
+const HelpTooltipStub = {
+ props: ['content', 'widthClass'],
+ template: '{{ content }}'
+}
+
+function mountView() {
+ return mount(AccountsView, {
+ global: {
+ stubs: {
+ AppLayout: { template: '
' },
+ TablePageLayout: {
+ template: '
'
+ },
+ DataTable: DataTableStub,
+ HelpTooltip: HelpTooltipStub,
+ Pagination: true,
+ ConfirmDialog: true,
+ AccountTableActions: { template: '
' },
+ AccountTableFilters: { template: '' },
+ AccountBulkActionsBar: true,
+ AccountActionMenu: true,
+ ImportDataModal: true,
+ ReAuthAccountModal: true,
+ AccountTestModal: true,
+ AccountStatsModal: true,
+ ScheduledTestsPanel: true,
+ SyncFromCrsModal: true,
+ TempUnschedStatusModal: true,
+ ErrorPassthroughRulesModal: true,
+ TLSFingerprintProfilesModal: true,
+ CreateAccountModal: true,
+ EditAccountModal: true,
+ BulkEditAccountModal: true,
+ PlatformTypeBadge: true,
+ AccountCapacityCell: true,
+ AccountStatusIndicator: true,
+ AccountTodayStatsCell: true,
+ AccountGroupsCell: true,
+ AccountUsageCell: true,
+ Icon: true
+ }
+ }
+ })
+}
+
+describe('admin AccountsView usage windows hint', () => {
+ beforeEach(() => {
+ localStorage.clear()
+
+ listAccounts.mockReset()
+ listWithEtag.mockReset()
+ getBatchTodayStats.mockReset()
+ getAllProxies.mockReset()
+ getAllGroups.mockReset()
+
+ listAccounts.mockResolvedValue({
+ items: [],
+ total: 0,
+ page: 1,
+ page_size: 20,
+ pages: 0
+ })
+ listWithEtag.mockResolvedValue({
+ notModified: true,
+ etag: null,
+ data: null
+ })
+ getBatchTodayStats.mockResolvedValue({ stats: {} })
+ getAllProxies.mockResolvedValue([])
+ getAllGroups.mockResolvedValue([])
+ })
+
+ it('renders an explanatory tooltip next to the usage windows column header', async () => {
+ const wrapper = mountView()
+ await flushPromises()
+
+ const header = wrapper.find('[data-test="usage-header"]')
+ expect(header.exists()).toBe(true)
+ // Column label is still shown alongside the help icon.
+ expect(header.text()).toContain('admin.accounts.columns.usageWindows')
+
+ const hint = wrapper.find('[data-test="usage-windows-hint"]')
+ expect(hint.exists()).toBe(true)
+ expect(hint.text()).toBe('admin.accounts.usageWindowsHint')
+ })
+})
diff --git a/frontend/src/views/admin/__tests__/RiskControlView.spec.ts b/frontend/src/views/admin/__tests__/RiskControlView.spec.ts
index 3c6aa0e9..5f1798a4 100644
--- a/frontend/src/views/admin/__tests__/RiskControlView.spec.ts
+++ b/frontend/src/views/admin/__tests__/RiskControlView.spec.ts
@@ -58,8 +58,12 @@ vi.mock('vue-i18n', async () => {
return {
...actual,
useI18n: () => ({
- t: (key: string, params?: Record) =>
- key.replace(/\{(\w+)\}/g, (_, token) => String(params?.[token] ?? `{${token}}`)),
+ t: (key: string, params?: Record) => {
+ if (key === 'admin.riskControl.preBlockAPIKeyLoadSummary') {
+ return `同步并发 ${params?.active} / 可用 Key ${params?.available},累计 ${params?.total} 次,worker:${params?.workerActive} / ${params?.workerTotal}`
+ }
+ return key.replace(/\{(\w+)\}/g, (_, token) => String(params?.[token] ?? `{${token}}`))
+ },
}),
}
})
@@ -118,6 +122,16 @@ const runtimeStatus = () => ({
dropped: 0,
processed: 0,
errors: 0,
+ pre_block_active: 0,
+ pre_block_checked: 0,
+ pre_block_allowed: 0,
+ pre_block_blocked: 0,
+ pre_block_errors: 0,
+ pre_block_avg_latency_ms: 0,
+ pre_block_api_key_active: 0,
+ pre_block_api_key_available_count: 0,
+ pre_block_api_key_total_calls: 0,
+ pre_block_api_key_loads: [],
api_key_statuses: [],
flagged_hash_count: 0,
last_cleanup_deleted_hit: 0,
@@ -261,4 +275,133 @@ describe('admin RiskControlView', () => {
}))
expect(showError).not.toHaveBeenCalled()
})
+
+ it('describes worker runtime as async audit and pre-block record processing', async () => {
+ getStatus.mockResolvedValue({
+ ...runtimeStatus(),
+ mode: 'observe',
+ processed: 12,
+ queue_length: 2,
+ })
+
+ const wrapper = mount(RiskControlView, {
+ global: {
+ stubs: {
+ AppLayout: AppLayoutStub,
+ BaseDialog: BaseDialogStub,
+ Icon: true,
+ Select: true,
+ Toggle: true,
+ Pagination: true,
+ ModelWhitelistSelector: ModelWhitelistSelectorStub,
+ },
+ },
+ })
+
+ await flushPromises()
+
+ expect(wrapper.text()).toContain('admin.riskControl.workerStatusHint')
+ expect(wrapper.text()).not.toContain('admin.riskControl.preBlockSyncStatus')
+ expect(wrapper.text()).toContain('admin.riskControl.records')
+ expect(wrapper.text()).toContain('12')
+ expect(wrapper.text()).toContain('2 / 32,768')
+ })
+
+ it('shows pre-block synchronous moderation metrics separately from worker queue', async () => {
+ getStatus.mockResolvedValue({
+ ...runtimeStatus(),
+ pre_block_active: 2,
+ pre_block_checked: 128,
+ pre_block_allowed: 120,
+ pre_block_blocked: 8,
+ pre_block_errors: 1,
+ pre_block_avg_latency_ms: 86,
+ pre_block_api_key_active: 2,
+ pre_block_api_key_available_count: 2,
+ pre_block_api_key_total_calls: 128,
+ active_workers: 3,
+ worker_count: 7,
+ pre_block_api_key_loads: [
+ {
+ index: 0,
+ key_hash: 'hash-one',
+ masked: 'sk-...one',
+ status: 'ok',
+ active: 1,
+ total: 72,
+ success: 70,
+ errors: 2,
+ avg_latency_ms: 84,
+ last_latency_ms: 80,
+ last_http_status: 200,
+ },
+ {
+ index: 1,
+ key_hash: 'hash-two',
+ masked: 'sk-...two',
+ status: 'ok',
+ active: 1,
+ total: 56,
+ success: 56,
+ errors: 0,
+ avg_latency_ms: 90,
+ last_latency_ms: 92,
+ last_http_status: 200,
+ },
+ ],
+ })
+
+ const wrapper = mount(RiskControlView, {
+ global: {
+ stubs: {
+ AppLayout: AppLayoutStub,
+ BaseDialog: BaseDialogStub,
+ Icon: true,
+ Select: true,
+ Toggle: true,
+ Pagination: true,
+ ModelWhitelistSelector: ModelWhitelistSelectorStub,
+ },
+ },
+ })
+
+ await flushPromises()
+
+ expect(wrapper.text()).toContain('admin.riskControl.preBlockSyncStatus')
+ expect(wrapper.text()).toContain('admin.riskControl.preBlockSyncHint')
+ expect(wrapper.text()).not.toContain('admin.riskControl.workerStatus')
+ expect(wrapper.text()).toContain('admin.riskControl.records')
+ expect(wrapper.text()).toContain('128')
+ expect(wrapper.text()).toContain('120')
+ expect(wrapper.text()).toContain('8')
+ expect(wrapper.text()).toContain('86 ms')
+ expect(wrapper.text()).toContain('admin.riskControl.preBlockAPIKeyLoad')
+ expect(wrapper.text()).toContain('sk-...one')
+ expect(wrapper.text()).toContain('sk-...two')
+ expect(wrapper.text()).toContain('72')
+ expect(wrapper.text()).toContain('56')
+ expect(wrapper.text()).toContain('同步并发 2 / 可用 Key 2,累计 128 次,worker:3 / 7')
+
+ const runtimeCards = wrapper.get('[data-test="pre-block-runtime-cards"]')
+ const syncCard = wrapper.get('[data-test="pre-block-sync-card"]')
+ const apiKeyLoadCard = wrapper.get('[data-test="pre-block-api-key-load-card"]')
+
+ expect(runtimeCards.classes()).toEqual(expect.arrayContaining([
+ 'grid',
+ 'grid-cols-1',
+ 'xl:grid-cols-[minmax(0,520px)_minmax(0,1fr)]',
+ ]))
+ expect(syncCard.element.parentElement).toBe(runtimeCards.element)
+ expect(apiKeyLoadCard.element.parentElement).toBe(runtimeCards.element)
+ expect(syncCard.classes()).toContain('card')
+ expect(apiKeyLoadCard.classes()).toContain('card')
+ expect(syncCard.get('h2').text()).toBe('admin.riskControl.preBlockSyncStatus')
+ expect(syncCard.text()).toContain('admin.riskControl.preBlockSyncHint')
+ expect(apiKeyLoadCard.get('h2').text()).toBe('admin.riskControl.preBlockAPIKeyLoad')
+ expect(apiKeyLoadCard.text()).toContain('admin.riskControl.preBlockAPIKeyLoadHint')
+ expect(wrapper.get('[data-test="pre-block-api-key-load-list"]').classes()).toEqual(expect.arrayContaining([
+ 'max-h-[280px]',
+ 'overflow-y-auto',
+ ]))
+ })
})
diff --git a/frontend/src/views/admin/__tests__/UsageView.spec.ts b/frontend/src/views/admin/__tests__/UsageView.spec.ts
index 1a5d285a..1da0d09c 100644
--- a/frontend/src/views/admin/__tests__/UsageView.spec.ts
+++ b/frontend/src/views/admin/__tests__/UsageView.spec.ts
@@ -3,7 +3,7 @@ import { flushPromises, mount } from '@vue/test-utils'
import UsageView from '../UsageView.vue'
-const { list, getStats, getSnapshotV2, getById } = vi.hoisted(() => {
+const { list, getStats, getSnapshotV2, getById, getModelStats, listErrorLogs } = vi.hoisted(() => {
vi.stubGlobal('localStorage', {
getItem: vi.fn(() => null),
setItem: vi.fn(),
@@ -15,6 +15,8 @@ const { list, getStats, getSnapshotV2, getById } = vi.hoisted(() => {
getStats: vi.fn(),
getSnapshotV2: vi.fn(),
getById: vi.fn(),
+ getModelStats: vi.fn(),
+ listErrorLogs: vi.fn(),
}
})
@@ -40,6 +42,7 @@ vi.mock('@/api/admin', () => ({
},
dashboard: {
getSnapshotV2,
+ getModelStats,
},
users: {
getById,
@@ -53,6 +56,10 @@ vi.mock('@/api/admin/usage', () => ({
},
}))
+vi.mock('@/api/admin/ops', () => ({
+ listErrorLogs,
+}))
+
vi.mock('@/stores/app', () => ({
useAppStore: () => ({
showError: vi.fn(),
@@ -84,6 +91,10 @@ vi.mock('vue-router', () => ({
const AppLayoutStub = { template: '
' }
const UsageFiltersStub = { template: '
' }
+const UsageTableStub = {
+ emits: ['userClick'],
+ template: 'user
',
+}
const ModelDistributionChartStub = {
props: ['metric'],
emits: ['update:metric'],
@@ -112,6 +123,7 @@ describe('admin UsageView distribution metric toggles', () => {
getStats.mockReset()
getSnapshotV2.mockReset()
getById.mockReset()
+ getModelStats.mockReset()
list.mockResolvedValue({
items: [],
@@ -133,12 +145,44 @@ describe('admin UsageView distribution metric toggles', () => {
models: [],
groups: [],
})
+ getModelStats.mockResolvedValue({ models: [] })
})
afterEach(() => {
vi.useRealTimers()
})
+ it('keeps previous model stats visible during refresh until new data arrives', async () => {
+ // 首次加载返回 A
+ getModelStats.mockResolvedValueOnce({ models: [{ model: 'A', total_tokens: 10 }] })
+
+ const wrapper = mount(UsageView, {
+ global: { stubs: {
+ AppLayout: AppLayoutStub, UsageStatsCards: true, UsageFilters: UsageFiltersStub,
+ UsageTable: true, UsageExportProgress: true, UsageCleanupDialog: true,
+ UserBalanceHistoryModal: true, AuditLogModal: true, Pagination: true, Select: true,
+ DateRangePicker: true, Icon: true, TokenUsageTrend: true,
+ ModelDistributionChart: ModelDistributionChartStub, GroupDistributionChart: GroupDistributionChartStub,
+ EndpointDistributionChart: true,
+ } },
+ })
+ vi.advanceTimersByTime(120)
+ await flushPromises()
+ expect((wrapper.vm as any).requestedModelStats).toEqual([{ model: 'A', total_tokens: 10 }])
+
+ // 刷新:让第二次 getModelStats 处于 pending,断言旧数据 A 仍在(不被清空成 [])
+ let resolveSecond: (v: any) => void = () => {}
+ getModelStats.mockReturnValueOnce(new Promise((res) => { resolveSecond = res }))
+ ;(wrapper.vm as any).refreshData()
+ await flushPromises()
+ expect((wrapper.vm as any).requestedModelStats).toEqual([{ model: 'A', total_tokens: 10 }])
+
+ // 新数据到达后替换为 B
+ resolveSecond({ models: [{ model: 'B', total_tokens: 20 }] })
+ await flushPromises()
+ expect((wrapper.vm as any).requestedModelStats).toEqual([{ model: 'B', total_tokens: 20 }])
+ })
+
it('keeps model and group metric toggles independent without refetching chart data', async () => {
const wrapper = mount(UsageView, {
global: {
@@ -194,3 +238,117 @@ describe('admin UsageView distribution metric toggles', () => {
expect(getSnapshotV2).toHaveBeenCalledTimes(1)
})
})
+
+describe('admin UsageView handleUserClick', () => {
+ beforeEach(() => {
+ vi.useFakeTimers()
+ list.mockReset()
+ getStats.mockReset()
+ getSnapshotV2.mockReset()
+ getById.mockReset()
+
+ list.mockResolvedValue({ items: [], total: 0, pages: 0 })
+ getStats.mockResolvedValue({
+ total_requests: 0, total_input_tokens: 0, total_output_tokens: 0,
+ total_cache_tokens: 0, total_tokens: 0, total_cost: 0, total_actual_cost: 0, average_duration_ms: 0,
+ })
+ getSnapshotV2.mockResolvedValue({ trend: [], models: [], groups: [] })
+ })
+
+ afterEach(() => {
+ vi.useRealTimers()
+ })
+
+ it('opens user via include_deleted when clicking a usage row user', async () => {
+ getById.mockResolvedValue({ id: 2, email: 'd@test.com', deleted_at: '2026-05-28T00:00:00Z' })
+
+ const wrapper = mount(UsageView, {
+ global: {
+ stubs: {
+ AppLayout: AppLayoutStub,
+ UsageStatsCards: true,
+ UsageFilters: UsageFiltersStub,
+ UsageTable: UsageTableStub,
+ UsageExportProgress: true,
+ UsageCleanupDialog: true,
+ UserBalanceHistoryModal: true,
+ AuditLogModal: true,
+ Pagination: true,
+ Select: true,
+ DateRangePicker: true,
+ Icon: true,
+ TokenUsageTrend: true,
+ ModelDistributionChart: true,
+ GroupDistributionChart: true,
+ EndpointDistributionChart: true,
+ },
+ },
+ })
+
+ vi.advanceTimersByTime(120)
+ await flushPromises()
+
+ await wrapper.find('[data-test="usage-table"] .user-click').trigger('click')
+ await flushPromises()
+
+ expect(getById).toHaveBeenCalledWith(2, true)
+ })
+})
+
+describe('admin UsageView errors tab filter forwarding', () => {
+ beforeEach(() => {
+ vi.useFakeTimers()
+ list.mockReset()
+ getStats.mockReset()
+ getSnapshotV2.mockReset()
+ getModelStats.mockReset()
+ listErrorLogs.mockReset()
+
+ list.mockResolvedValue({ items: [], total: 0, pages: 0 })
+ getStats.mockResolvedValue({
+ total_requests: 0, total_input_tokens: 0, total_output_tokens: 0,
+ total_cache_tokens: 0, total_tokens: 0, total_cost: 0, total_actual_cost: 0, average_duration_ms: 0,
+ })
+ getSnapshotV2.mockResolvedValue({ trend: [], models: [], groups: [] })
+ getModelStats.mockResolvedValue({ models: [] })
+ listErrorLogs.mockResolvedValue({ items: [], total: 0, pages: 0 })
+ })
+
+ afterEach(() => {
+ vi.useRealTimers()
+ })
+
+ it('forwards model/account_id/group_id to listErrorLogs on the errors tab', async () => {
+ const wrapper = mount(UsageView, {
+ global: { stubs: {
+ AppLayout: AppLayoutStub, UsageStatsCards: true, UsageFilters: UsageFiltersStub,
+ UsageTable: true, UsageExportProgress: true, UsageCleanupDialog: true,
+ UserBalanceHistoryModal: true, AuditLogModal: true, Pagination: true, Select: true,
+ DateRangePicker: true, Icon: true, TokenUsageTrend: true,
+ ModelDistributionChart: true, GroupDistributionChart: true, EndpointDistributionChart: true,
+ OpsErrorLogTable: true, OpsErrorDetailModal: true,
+ } },
+ })
+ vi.advanceTimersByTime(120)
+ await flushPromises()
+
+ // 模拟用户在过滤器里选择了模型/账户/分组
+ const vm = wrapper.vm as any
+ vm.filters.model = 'gpt-5.3-codex'
+ vm.filters.account_id = 7
+ vm.filters.group_id = 3
+ await flushPromises()
+
+ // 切换到「错误请求」标签(第二个 .tab 按钮)触发 loadAdminErrors
+ const tabs = wrapper.findAll('button.tab')
+ await tabs[1].trigger('click')
+ await flushPromises()
+
+ expect(listErrorLogs).toHaveBeenCalledWith(expect.objectContaining({
+ view: 'all',
+ model: 'gpt-5.3-codex',
+ account_id: 7,
+ group_id: 3,
+ }))
+ })
+})
diff --git a/frontend/src/views/admin/__tests__/groupsModelsList.spec.ts b/frontend/src/views/admin/__tests__/groupsModelsList.spec.ts
new file mode 100644
index 00000000..ae50c861
--- /dev/null
+++ b/frontend/src/views/admin/__tests__/groupsModelsList.spec.ts
@@ -0,0 +1,125 @@
+import { describe, expect, it } from "vitest";
+
+import {
+ buildModelsListConfig,
+ createModelsListState,
+ hydrateModelsListState,
+ invertModelsListSelection,
+ moveModelsListItem,
+ selectAllModelsListItems,
+ setModelsListCandidates,
+ toggleModelsListItem,
+} from "../groupsModelsList";
+
+describe("groupsModelsList", () => {
+ it("selects all default candidates for a new disabled config", () => {
+ const state = createModelsListState();
+
+ setModelsListCandidates(state, ["gpt-5.5", "gpt-5.4"]);
+
+ expect(state.enabled).toBe(false);
+ expect(state.items).toEqual([
+ { id: "gpt-5.5", selected: true },
+ { id: "gpt-5.4", selected: true },
+ ]);
+ });
+
+ it("keeps saved selections and marks new candidates as unselected when editing", () => {
+ const state = createModelsListState({
+ enabled: true,
+ models: ["gpt-5.5", "gpt-5.4"],
+ });
+
+ setModelsListCandidates(state, ["gpt-5.4", "legacy-gpt", "gpt-5.5"]);
+
+ expect(state.enabled).toBe(true);
+ expect(state.items).toEqual([
+ { id: "gpt-5.5", selected: true },
+ { id: "gpt-5.4", selected: true },
+ { id: "legacy-gpt", selected: false },
+ ]);
+ });
+
+ it("preserves explicitly unselected saved candidates when candidates refresh", () => {
+ const state = createModelsListState({
+ enabled: true,
+ models: ["gpt-5.5"],
+ });
+
+ setModelsListCandidates(state, ["gpt-5.5", "gpt-5.4"]);
+
+ expect(state.items).toEqual([
+ { id: "gpt-5.5", selected: true },
+ { id: "gpt-5.4", selected: false },
+ ]);
+ });
+
+ it("builds config with selected models in current display order", () => {
+ const state = hydrateModelsListState({
+ enabled: true,
+ models: ["gpt-5.5", "gpt-5.4", "legacy-gpt"],
+ }, ["gpt-5.5", "gpt-5.4", "legacy-gpt"]);
+
+ toggleModelsListItem(state, "legacy-gpt");
+ moveModelsListItem(state, 1, 0);
+
+ expect(buildModelsListConfig(state)).toEqual({
+ enabled: true,
+ models: ["gpt-5.4", "gpt-5.5"],
+ });
+ });
+
+ it("keeps selected models in payload even when disabled so reopening can restore choices", () => {
+ const state = hydrateModelsListState({
+ enabled: false,
+ models: ["gpt-5.5"],
+ }, ["gpt-5.5", "gpt-5.4"]);
+
+ expect(buildModelsListConfig(state)).toEqual({
+ enabled: false,
+ models: ["gpt-5.5"],
+ });
+ });
+
+ it("preserves saved models when candidates have not loaded yet", () => {
+ const state = createModelsListState({
+ enabled: true,
+ models: ["gpt-5.5", "gpt-5.4"],
+ });
+
+ expect(buildModelsListConfig(state)).toEqual({
+ enabled: true,
+ models: ["gpt-5.5", "gpt-5.4"],
+ });
+ });
+
+ it("selects all candidate models from the toolbar action", () => {
+ const state = hydrateModelsListState({
+ enabled: true,
+ models: ["gpt-5.5"],
+ }, ["gpt-5.5", "gpt-5.4", "gpt-5.4-mini"]);
+
+ selectAllModelsListItems(state);
+
+ expect(state.items).toEqual([
+ { id: "gpt-5.5", selected: true },
+ { id: "gpt-5.4", selected: true },
+ { id: "gpt-5.4-mini", selected: true },
+ ]);
+ });
+
+ it("inverts selected models from the toolbar action", () => {
+ const state = hydrateModelsListState({
+ enabled: true,
+ models: ["gpt-5.5"],
+ }, ["gpt-5.5", "gpt-5.4", "gpt-5.4-mini"]);
+
+ invertModelsListSelection(state);
+
+ expect(state.items).toEqual([
+ { id: "gpt-5.5", selected: false },
+ { id: "gpt-5.4", selected: true },
+ { id: "gpt-5.4-mini", selected: true },
+ ]);
+ });
+});
diff --git a/frontend/src/views/admin/__tests__/groupsModelsListCandidates.spec.ts b/frontend/src/views/admin/__tests__/groupsModelsListCandidates.spec.ts
new file mode 100644
index 00000000..ec292c63
--- /dev/null
+++ b/frontend/src/views/admin/__tests__/groupsModelsListCandidates.spec.ts
@@ -0,0 +1,65 @@
+import { describe, expect, it } from "vitest";
+
+import {
+ createModelsListCandidatesTracker,
+} from "../groupsModelsListCandidates";
+
+describe("groupsModelsListCandidates", () => {
+ it("rejects stale candidate responses after a newer platform request starts", () => {
+ const tracker = createModelsListCandidatesTracker();
+ const first = {
+ mode: "create" as const,
+ groupID: 0,
+ platform: "openai" as const,
+ };
+ const second = {
+ mode: "create" as const,
+ groupID: 0,
+ platform: "anthropic" as const,
+ };
+
+ const firstID = tracker.next(first);
+ const secondID = tracker.next(second);
+
+ expect(tracker.isCurrent(firstID, first)).toBe(false);
+ expect(tracker.isCurrent(secondID, second)).toBe(true);
+ });
+
+ it("rejects responses for a previous edit group even with the same platform", () => {
+ const tracker = createModelsListCandidatesTracker();
+ const first = {
+ mode: "edit" as const,
+ groupID: 10,
+ platform: "openai" as const,
+ };
+ const second = {
+ mode: "edit" as const,
+ groupID: 11,
+ platform: "openai" as const,
+ };
+
+ const firstID = tracker.next(first);
+ tracker.next(second);
+
+ expect(tracker.isCurrent(firstID, first)).toBe(false);
+ });
+
+ it("tracks create and edit requests independently", () => {
+ const tracker = createModelsListCandidatesTracker();
+ const editRequest = {
+ mode: "edit" as const,
+ groupID: 10,
+ platform: "openai" as const,
+ };
+ const createRequest = {
+ mode: "create" as const,
+ groupID: 0,
+ platform: "anthropic" as const,
+ };
+
+ const editID = tracker.next(editRequest);
+ tracker.next(createRequest);
+
+ expect(tracker.isCurrent(editID, editRequest)).toBe(true);
+ });
+});
diff --git a/frontend/src/views/admin/__tests__/groupsModelsListLayout.spec.ts b/frontend/src/views/admin/__tests__/groupsModelsListLayout.spec.ts
new file mode 100644
index 00000000..6ac3d769
--- /dev/null
+++ b/frontend/src/views/admin/__tests__/groupsModelsListLayout.spec.ts
@@ -0,0 +1,19 @@
+import { readFileSync } from "node:fs";
+import { fileURLToPath } from "node:url";
+import { dirname, resolve } from "node:path";
+
+import { describe, expect, it } from "vitest";
+
+const currentDir = dirname(fileURLToPath(import.meta.url));
+const groupsViewSource = readFileSync(
+ resolve(currentDir, "../GroupsView.vue"),
+ "utf8",
+);
+
+describe("groups models list layout", () => {
+ it("keeps the toolbar outside of the scrolling list content", () => {
+ expect(groupsViewSource).toContain("overflow-hidden rounded-lg border");
+ expect(groupsViewSource).toContain("max-h-64 space-y-2 overflow-y-auto p-2");
+ expect(groupsViewSource).not.toContain("sticky top-0");
+ });
+});
diff --git a/frontend/src/views/admin/groupsModelsList.ts b/frontend/src/views/admin/groupsModelsList.ts
new file mode 100644
index 00000000..790268fe
--- /dev/null
+++ b/frontend/src/views/admin/groupsModelsList.ts
@@ -0,0 +1,121 @@
+export interface ModelsListConfig {
+ enabled: boolean
+ models: string[]
+}
+
+export interface ModelsListItem {
+ id: string
+ selected: boolean
+}
+
+export interface ModelsListState {
+ enabled: boolean
+ savedModels: string[]
+ items: ModelsListItem[]
+}
+
+export const createModelsListState = (
+ config?: Partial | null,
+): ModelsListState => ({
+ enabled: config?.enabled ?? false,
+ savedModels: normalizeModels(config?.models ?? []),
+ items: [],
+})
+
+export const hydrateModelsListState = (
+ config: Partial | null | undefined,
+ candidates: string[],
+): ModelsListState => {
+ const state = createModelsListState(config)
+ setModelsListCandidates(state, candidates)
+ return state
+}
+
+export const setModelsListCandidates = (
+ state: ModelsListState,
+ candidates: string[],
+) => {
+ const normalizedCandidates = normalizeModels(candidates)
+ const currentSelected = new Set(
+ state.items.filter(item => item.selected).map(item => item.id),
+ )
+ const currentKnown = new Set(state.items.map(item => item.id))
+ const savedSelected = new Set(state.savedModels)
+ const hasExistingItems = state.items.length > 0
+ const selectionOrder = normalizeModels([
+ ...state.items.map(item => item.id),
+ ...state.savedModels,
+ ...normalizedCandidates,
+ ])
+
+ state.items = selectionOrder.map(id => {
+ const selected = hasExistingItems
+ ? currentSelected.has(id)
+ : state.savedModels.length > 0
+ ? savedSelected.has(id)
+ : normalizedCandidates.includes(id)
+
+ return {
+ id,
+ selected: selected && (currentKnown.has(id) || savedSelected.has(id) || state.savedModels.length === 0),
+ }
+ })
+}
+
+export const toggleModelsListItem = (state: ModelsListState, modelID: string) => {
+ const item = state.items.find(item => item.id === modelID)
+ if (item) {
+ item.selected = !item.selected
+ }
+}
+
+export const selectAllModelsListItems = (state: ModelsListState) => {
+ state.items.forEach(item => {
+ item.selected = true
+ })
+}
+
+export const invertModelsListSelection = (state: ModelsListState) => {
+ state.items.forEach(item => {
+ item.selected = !item.selected
+ })
+}
+
+export const moveModelsListItem = (
+ state: ModelsListState,
+ fromIndex: number,
+ toIndex: number,
+) => {
+ if (
+ fromIndex === toIndex ||
+ fromIndex < 0 ||
+ toIndex < 0 ||
+ fromIndex >= state.items.length ||
+ toIndex >= state.items.length
+ ) {
+ return
+ }
+ const [item] = state.items.splice(fromIndex, 1)
+ state.items.splice(toIndex, 0, item)
+}
+
+export const buildModelsListConfig = (state: ModelsListState): ModelsListConfig => ({
+ enabled: state.enabled,
+ models: state.items.length > 0
+ ? state.items.filter(item => item.selected).map(item => item.id)
+ : [...state.savedModels],
+})
+
+const normalizeModels = (models: string[]): string[] => {
+ const seen = new Set()
+ const out: string[] = []
+ for (const raw of models) {
+ const model = raw.trim()
+ if (!model || seen.has(model)) {
+ continue
+ }
+ seen.add(model)
+ out.push(model)
+ }
+ return out
+}
diff --git a/frontend/src/views/admin/groupsModelsListCandidates.ts b/frontend/src/views/admin/groupsModelsListCandidates.ts
new file mode 100644
index 00000000..2c722af8
--- /dev/null
+++ b/frontend/src/views/admin/groupsModelsListCandidates.ts
@@ -0,0 +1,41 @@
+import type { GroupPlatform } from "@/types";
+
+export type ModelsListCandidatesMode = "create" | "edit";
+
+export interface ModelsListCandidatesRequest {
+ mode: ModelsListCandidatesMode;
+ groupID: number;
+ platform: GroupPlatform;
+}
+
+export interface ModelsListCandidatesTracker {
+ next(request: ModelsListCandidatesRequest): number;
+ isCurrent(requestID: number, request: ModelsListCandidatesRequest): boolean;
+}
+
+export const createModelsListCandidatesTracker = (): ModelsListCandidatesTracker => {
+ let currentRequestID = 0;
+ const currentByMode: Partial> = {};
+
+ return {
+ next(request) {
+ currentRequestID += 1;
+ currentByMode[request.mode] = {
+ id: currentRequestID,
+ request: { ...request },
+ };
+ return currentRequestID;
+ },
+ isCurrent(requestID, request) {
+ const current = currentByMode[request.mode];
+ return (
+ current?.id === requestID &&
+ current.request.groupID === request.groupID &&
+ current.request.platform === request.platform
+ );
+ },
+ };
+};
diff --git a/frontend/src/views/admin/ops/components/OpsErrorDetailModal.vue b/frontend/src/views/admin/ops/components/OpsErrorDetailModal.vue
index d29607e5..c346c547 100644
--- a/frontend/src/views/admin/ops/components/OpsErrorDetailModal.vue
+++ b/frontend/src/views/admin/ops/components/OpsErrorDetailModal.vue
@@ -106,6 +106,31 @@
{{ detail.message || '—' }}
+
+
+
{{ t('admin.ops.errorDetail.apiKeyPrefix') }}
+
+ {{ detail.api_key_prefix }}
+
+
+
+
+
{{ t('admin.ops.errorDetail.attemptedKeyPrefix') }}
+
+ {{ detail.attempted_key_prefix }}
+
+
+
+
+
{{ t('admin.ops.errorDetail.deletedKeyOwner') }}
+
+ {{ detail.deleted_key_owner_email }}
+ ({{ detail.deleted_key_name }})
+
+ {{ t('admin.ops.errorDetail.keyDeletedBadge') }}
+
+
+
diff --git a/frontend/src/views/admin/ops/components/OpsErrorLogTable.vue b/frontend/src/views/admin/ops/components/OpsErrorLogTable.vue
index d779f7b1..becf6327 100644
--- a/frontend/src/views/admin/ops/components/OpsErrorLogTable.vue
+++ b/frontend/src/views/admin/ops/components/OpsErrorLogTable.vue
@@ -32,6 +32,12 @@
{{ t('admin.ops.errorLog.user') }}
|
+
+ {{ t('admin.ops.errorLog.apiKey') }}
+ |
+
+ {{ t('admin.ops.errorLog.account') }}
+ |
{{ t('admin.ops.errorLog.status') }}
|
@@ -45,7 +51,7 @@
- |
+ |
{{ t('admin.ops.errorLog.noErrors') }}
|
@@ -127,24 +133,40 @@
-
-
+
-
-
-
- {{ log.account_name || '-' }}
-
-
- -
-
-
-
-
- {{ log.user_email || '-' }}
-
-
- -
-
+
+
+ {{ log.user_email || '-' }}
+
+
+ -
+ |
+
+
+
+
+
+ {{ log.api_key_name || ('#' + log.api_key_id) }}
+
+
+ {{ t('admin.ops.errorLog.keyDeletedBadge') }}
+
+
+ -
+ |
+
+
+
+
+
+ {{ log.account_name || '-' }}
+
+
+ -
|
diff --git a/frontend/src/views/admin/ops/components/OpsSettingsDialog.vue b/frontend/src/views/admin/ops/components/OpsSettingsDialog.vue
index 5dba5b1d..bfb7a65f 100644
--- a/frontend/src/views/admin/ops/components/OpsSettingsDialog.vue
+++ b/frontend/src/views/admin/ops/components/OpsSettingsDialog.vue
@@ -50,6 +50,10 @@ async function loadAllSettings() {
runtimeSettings.value = runtime
emailConfig.value = email
advancedSettings.value = advanced
+ // 兼容旧 payload:后端未返回该字段时补默认值,保证表单可绑定
+ if (advancedSettings.value && !advancedSettings.value.openai_account_quota_auto_pause) {
+ advancedSettings.value.openai_account_quota_auto_pause = { default_threshold_5h: 0, default_threshold_7d: 0 }
+ }
// 如果后端返回了阈值,使用后端的值;否则保持默认值
if (thresholds && Object.keys(thresholds).length > 0) {
metricThresholds.value = {
@@ -119,6 +123,28 @@ function removeRecipient(target: 'alert' | 'report', email: string) {
if (idx >= 0) list.splice(idx, 1)
}
+// OpenAI 账号配额自动暂停:后端按 0~1 分数存储,UI 按百分比(0~100)展示
+const quotaAutoPause5hPercent = computed({
+ get() {
+ const v = advancedSettings.value?.openai_account_quota_auto_pause?.default_threshold_5h
+ return v && v > 0 ? Math.round(v * 1000) / 10 : null
+ },
+ set(val) {
+ if (!advancedSettings.value?.openai_account_quota_auto_pause) return
+ advancedSettings.value.openai_account_quota_auto_pause.default_threshold_5h = val != null && val > 0 ? val / 100 : 0
+ }
+})
+const quotaAutoPause7dPercent = computed({
+ get() {
+ const v = advancedSettings.value?.openai_account_quota_auto_pause?.default_threshold_7d
+ return v && v > 0 ? Math.round(v * 1000) / 10 : null
+ },
+ set(val) {
+ if (!advancedSettings.value?.openai_account_quota_auto_pause) return
+ advancedSettings.value.openai_account_quota_auto_pause.default_threshold_7d = val != null && val > 0 ? val / 100 : 0
+ }
+})
+
// 验证
const validation = computed(() => {
const errors: string[] = []
@@ -145,6 +171,11 @@ const validation = computed(() => {
if (hourly_metrics_retention_days < 0 || hourly_metrics_retention_days > 365) {
errors.push(t('admin.ops.settings.validation.retentionDaysRange'))
}
+
+ const { default_threshold_5h, default_threshold_7d } = advancedSettings.value.openai_account_quota_auto_pause
+ if (default_threshold_5h < 0 || default_threshold_5h > 1 || default_threshold_7d < 0 || default_threshold_7d > 1) {
+ errors.push(t('admin.ops.settings.validation.openaiQuotaAutoPauseRange'))
+ }
}
// 验证指标阈值
@@ -473,6 +504,40 @@ async function saveAllSettings() {
+
+
+
{{ t('admin.ops.settings.openaiQuotaAutoPause') }}
+
{{ t('admin.ops.settings.openaiQuotaAutoPauseHint') }}
+
+
+
{{ t('admin.ops.settings.openaiQuotaAutoPauseThresholdHint') }}
+
+
{{ t('admin.ops.settings.errorFiltering') }}
diff --git a/frontend/src/views/admin/ops/components/__tests__/OpsErrorLogTable.spec.ts b/frontend/src/views/admin/ops/components/__tests__/OpsErrorLogTable.spec.ts
new file mode 100644
index 00000000..3f23028f
--- /dev/null
+++ b/frontend/src/views/admin/ops/components/__tests__/OpsErrorLogTable.spec.ts
@@ -0,0 +1,93 @@
+import { describe, it, expect, vi } from 'vitest'
+import { mount } from '@vue/test-utils'
+import OpsErrorLogTable from '../OpsErrorLogTable.vue'
+import zhLocale from '@/i18n/locales/zh'
+import enLocale from '@/i18n/locales/en'
+import type { OpsErrorLog } from '@/api/admin/ops'
+
+vi.mock('vue-i18n', async (importOriginal) => {
+ const actual = await importOriginal
()
+ return {
+ ...actual,
+ useI18n: () => ({ t: (key: string) => key }),
+ }
+})
+
+const TooltipStub = { template: '
' }
+const PaginationStub = { template: '' }
+
+function mountTable(row: Partial) {
+ const base = {
+ id: 1,
+ created_at: '2026-06-05T23:59:50Z',
+ phase: 'upstream',
+ type: '',
+ error_owner: 'provider',
+ error_source: 'upstream_http',
+ severity: 'error',
+ status_code: 529,
+ platform: 'anthropic',
+ model: 'claude-opus-4-8',
+ resolved: false,
+ client_request_id: '',
+ request_id: 'req-1',
+ message: 'boom',
+ user_email: '',
+ account_name: '',
+ group_name: '',
+ ...row,
+ } as OpsErrorLog
+
+ return mount(OpsErrorLogTable, {
+ props: { rows: [base], total: 1, loading: false, page: 1, pageSize: 20 },
+ global: { stubs: { 'el-tooltip': TooltipStub, Pagination: PaginationStub } },
+ })
+}
+
+describe('OpsErrorLogTable user/api-key/account columns', () => {
+ // 回归:上游错误行(phase=upstream, owner=provider)以前在单一「用户」列里只显示账号、
+ // 丢失用户;现在用户/API Key/账号各占独立列,三者同时可见。
+ it('renders user, api key and account in separate columns for an upstream row', () => {
+ const wrapper = mountTable({
+ user_id: 2,
+ user_email: 'alice@test.com',
+ api_key_id: 5,
+ api_key_name: 'my-key',
+ account_id: 9,
+ account_name: 'acct-A',
+ })
+
+ const text = wrapper.text()
+ expect(text).toContain('alice@test.com') // 用户列(上游行也显示用户)
+ expect(text).toContain('my-key') // API Key 列
+ expect(text).toContain('acct-A') // 账号列
+ })
+
+ it('shows the deleted badge for a soft-deleted api key', () => {
+ const wrapper = mountTable({
+ api_key_id: 5,
+ api_key_name: 'old-key',
+ api_key_deleted: true,
+ })
+
+ expect(wrapper.text()).toContain('old-key')
+ expect(wrapper.text()).toContain('admin.ops.errorLog.keyDeletedBadge')
+ })
+})
+
+// 防回归:组件用 admin.ops.errorLog.* 命名空间。若 i18n 键写错命名空间(如误放到
+// errorDetail),真实 vue-i18n 会回退返回 key 本身 → 界面显示原始路径字符串。
+// 这里用真实 locale 校验键确实可解析(返回译文而非 key)。
+// 防回归:组件用 admin.ops.errorLog.* 命名空间。若键写错命名空间(如误放到
+// errorDetail),界面会显示原始路径字符串而非译文。vitest 的 vue-i18n 为 runtime-only
+// (无消息编译器,t() 对任何键都回退返回 key),故直接校验 locale 对象的命名空间含这些键。
+describe('OpsErrorLogTable i18n keys exist in the errorLog namespace', () => {
+ const locales: Record = { zh: zhLocale, en: enLocale }
+ for (const [name, msgs] of Object.entries(locales)) {
+ it(`has apiKey & keyDeletedBadge for ${name}`, () => {
+ const errorLog = msgs?.admin?.ops?.errorLog
+ expect(errorLog?.apiKey).toBeTruthy()
+ expect(errorLog?.keyDeletedBadge).toBeTruthy()
+ })
+ }
+})
diff --git a/frontend/src/views/user/UsageView.vue b/frontend/src/views/user/UsageView.vue
index d6807aaf..01053603 100644
--- a/frontend/src/views/user/UsageView.vue
+++ b/frontend/src/views/user/UsageView.vue
@@ -149,11 +149,27 @@
-
+
+
+ {{ t('usage.tabs.usage') }}
+
+
+ {{ t('usage.tabs.errors') }}
+
+
+
+
+
+
+
{{
- row.input_tokens.toLocaleString()
+ (row.input_tokens ?? 0).toLocaleString()
}}
{{
- row.output_tokens.toLocaleString()
+ (row.output_tokens ?? 0).toLocaleString()
}}
@@ -280,7 +296,7 @@
- ${{ row.actual_cost.toFixed(6) }}
+ ${{ (row.actual_cost ?? 0).toFixed(6) }}
+
+
+
+
+
+
{
+const formatDuration = (ms: number | null | undefined): string => {
+ if (ms == null) return '-'
if (ms < 1000) return `${ms.toFixed(0)}ms`
return `${(ms / 1000).toFixed(2)}s`
}
@@ -926,8 +966,8 @@ const exportToCSV = async () => {
log.cache_read_tokens,
log.cache_creation_tokens,
log.rate_multiplier,
- log.actual_cost.toFixed(8),
- log.total_cost.toFixed(8),
+ (log.actual_cost ?? 0).toFixed(8),
+ (log.total_cost ?? 0).toFixed(8),
log.first_token_ms ?? '',
log.duration_ms
].map(escapeCSVValue)
@@ -988,6 +1028,52 @@ const hideTokenTooltip = () => {
tokenTooltipData.value = null
}
+// ── Error Requests Tab ──────────────────────────────────────────────────────
+const activeTab = ref<'usage' | 'errors'>('usage')
+const errorViewEnabled = computed(() => appStore.cachedPublicSettings?.allow_user_view_error_requests ?? false)
+
+const errorRows = ref([])
+const errorLoading = ref(false)
+const errorPage = ref(1)
+const errorPageSize = ref(20)
+const errorTotal = ref(0)
+const errorFilter = ref<{ model: string; category: string; api_key_id: number | null }>({ model: '', category: '', api_key_id: null })
+
+const loadErrors = async () => {
+ errorLoading.value = true
+ try {
+ const resp = await usageAPI.listMyErrorRequests({
+ page: errorPage.value,
+ page_size: errorPageSize.value,
+ start_date: startDate.value,
+ end_date: endDate.value,
+ model: errorFilter.value.model || undefined,
+ category: errorFilter.value.category || undefined,
+ api_key_id: errorFilter.value.api_key_id ?? undefined,
+ })
+ errorRows.value = resp.items
+ errorTotal.value = resp.total
+ } catch (error) {
+ console.error('[UsageView] loadErrors failed:', error)
+ appStore.showError(t('usage.errors.failedToLoad'))
+ } finally {
+ errorLoading.value = false
+ }
+}
+
+const onErrorFilter = (f: { model: string; category: string; api_key_id: number | null }) => {
+ errorFilter.value = f
+ errorPage.value = 1
+ loadErrors()
+}
+const onErrorPage = (p: number) => { errorPage.value = p; loadErrors() }
+const onErrorPageSize = (s: number) => { errorPageSize.value = s; errorPage.value = 1; loadErrors() }
+
+const switchToErrors = () => {
+ activeTab.value = 'errors'
+ if (errorRows.value.length === 0) loadErrors()
+}
+
onMounted(() => {
loadApiKeys()
loadUsageLogs()