Merge remote-tracking branch 'origin/main' into upgrade/upstream-v0.1.133-20260605
# Conflicts: # backend/internal/service/setting_service_public_test.go
This commit is contained in:
commit
63b9c880b5
4
.github/workflows/backend-ci.yml
vendored
4
.github/workflows/backend-ci.yml
vendored
@ -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:
|
||||
|
||||
2
.github/workflows/release.yml
vendored
2
.github/workflows/release.yml
vendored
@ -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
|
||||
|
||||
2
.github/workflows/security-scan.yml
vendored
2
.github/workflows/security-scan.yml
vendored
@ -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: |
|
||||
|
||||
@ -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
|
||||
|
||||
@ -1,4 +1,4 @@
|
||||
FROM golang:1.26.3-alpine
|
||||
FROM golang:1.26.4-alpine
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
|
||||
@ -1 +1 @@
|
||||
0.1.131
|
||||
0.1.133
|
||||
|
||||
@ -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{
|
||||
|
||||
@ -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{
|
||||
|
||||
@ -77,6 +77,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
|
||||
nil, // backupSvc
|
||||
nil, // paymentOrderExpiry
|
||||
nil, // channelMonitorRunner
|
||||
nil, // quotaFlusher
|
||||
)
|
||||
|
||||
require.NotPanics(t, func() {
|
||||
|
||||
@ -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(')')
|
||||
|
||||
@ -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
|
||||
)
|
||||
|
||||
@ -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) {
|
||||
|
||||
@ -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)
|
||||
}
|
||||
|
||||
@ -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.
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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").
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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")
|
||||
}
|
||||
|
||||
@ -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 {
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
7
backend/internal/domain/models_list_config.go
Normal file
7
backend/internal/domain/models_list_config.go
Normal file
@ -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"`
|
||||
}
|
||||
@ -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) {
|
||||
|
||||
52
backend/internal/handler/admin/account_handler_list_test.go
Normal file
52
backend/internal/handler/admin/account_handler_list_test.go
Normal file
@ -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)
|
||||
}
|
||||
@ -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))
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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,
|
||||
})
|
||||
|
||||
@ -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.
|
||||
|
||||
@ -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")
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
|
||||
144
backend/internal/handler/admin/system_handler_test.go
Normal file
144
backend/internal/handler/admin/system_handler_test.go
Normal file
@ -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)
|
||||
}
|
||||
@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@ -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")
|
||||
}
|
||||
62
backend/internal/handler/admin/usage_query_cache.go
Normal file
62
backend/internal/handler/admin/usage_query_cache.go
Normal file
@ -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
|
||||
}
|
||||
28
backend/internal/handler/admin/usage_query_cache_test.go
Normal file
28
backend/internal/handler/admin/usage_query_cache_test.go
Normal file
@ -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")
|
||||
}
|
||||
@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@ -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)
|
||||
})
|
||||
}
|
||||
@ -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
|
||||
|
||||
27
backend/internal/handler/concurrency_error_response.go
Normal file
27
backend/internal/handler/concurrency_error_response.go
Normal file
@ -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"
|
||||
}
|
||||
63
backend/internal/handler/concurrency_error_response_test.go
Normal file
63
backend/internal/handler/concurrency_error_response_test.go
Normal file
@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@ -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,
|
||||
|
||||
20
backend/internal/handler/dto/mappers_deleted_user_test.go
Normal file
20
backend/internal/handler/dto/mappers_deleted_user_test.go
Normal file
@ -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")
|
||||
}
|
||||
@ -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 {
|
||||
|
||||
@ -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"`
|
||||
|
||||
@ -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.
|
||||
|
||||
@ -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},
|
||||
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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"),
|
||||
|
||||
@ -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 {
|
||||
|
||||
@ -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 {
|
||||
|
||||
@ -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,
|
||||
|
||||
247
backend/internal/handler/openai_embeddings.go
Normal file
247
backend/internal/handler/openai_embeddings.go
Normal file
@ -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
|
||||
}
|
||||
}
|
||||
@ -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
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
}
|
||||
@ -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))
|
||||
}
|
||||
|
||||
@ -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
|
||||
|
||||
118
backend/internal/handler/ops_error_logger_attribution_test.go
Normal file
118
backend/internal/handler/ops_error_logger_attribution_test.go
Normal file
@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@ -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")
|
||||
}
|
||||
|
||||
@ -98,6 +98,8 @@ func (h *SettingHandler) GetPublicSettings(c *gin.Context) {
|
||||
AffiliateEnabled: settings.AffiliateEnabled,
|
||||
|
||||
RiskControlEnabled: settings.RiskControlEnabled,
|
||||
|
||||
AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@ -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) {
|
||||
|
||||
@ -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})
|
||||
|
||||
@ -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})
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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(),
|
||||
|
||||
131
backend/internal/payment/provider/easypay_query_test.go
Normal file
131
backend/internal/payment/provider/easypay_query_test.go
Normal file
@ -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)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@ -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"},
|
||||
}
|
||||
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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 {
|
||||
|
||||
@ -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)
|
||||
}
|
||||
|
||||
@ -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{
|
||||
|
||||
@ -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 {
|
||||
|
||||
@ -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")
|
||||
}
|
||||
}
|
||||
@ -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"`)
|
||||
}
|
||||
@ -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"
|
||||
|
||||
199
backend/internal/pkg/apicompat/responses_stream_event_wire.go
Normal file
199
backend/internal/pkg/apicompat/responses_stream_event_wire.go
Normal file
@ -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
|
||||
}
|
||||
@ -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")
|
||||
}
|
||||
@ -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")
|
||||
}
|
||||
@ -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 {
|
||||
|
||||
@ -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_<id>.
|
||||
|
||||
// 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"))
|
||||
}
|
||||
@ -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,
|
||||
|
||||
@ -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.
|
||||
|
||||
@ -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",
|
||||
|
||||
78
backend/internal/pkg/openai/allowed_client.go
Normal file
78
backend/internal/pkg/openai/allowed_client.go
Normal file
@ -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
|
||||
}
|
||||
95
backend/internal/pkg/openai/allowed_client_test.go
Normal file
95
backend/internal/pkg/openai/allowed_client_test.go
Normal file
@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@ -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
|
||||
}
|
||||
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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)
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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":
|
||||
|
||||
40
backend/internal/repository/content_moderation_repo_test.go
Normal file
40
backend/internal/repository/content_moderation_repo_test.go
Normal file
@ -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())
|
||||
}
|
||||
@ -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),
|
||||
),
|
||||
)
|
||||
}
|
||||
|
||||
@ -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) {
|
||||
|
||||
@ -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。
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Loading…
x
Reference in New Issue
Block a user