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:
zizi 2026-06-05 16:35:10 +08:00
commit 63b9c880b5
333 changed files with 26206 additions and 4931 deletions

View File

@ -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:

View File

@ -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

View File

@ -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: |

View File

@ -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

View File

@ -1,4 +1,4 @@
FROM golang:1.26.3-alpine
FROM golang:1.26.4-alpine
WORKDIR /app

View File

@ -1 +1 @@
0.1.131
0.1.133

View File

@ -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{

View File

@ -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{

View File

@ -77,6 +77,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
nil, // backupSvc
nil, // paymentOrderExpiry
nil, // channelMonitorRunner
nil, // quotaFlusher
)
require.NotPanics(t, func() {

View File

@ -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(')')

View File

@ -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
)

View File

@ -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) {

View File

@ -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)
}

View File

@ -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.

View File

@ -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

View File

@ -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()

View File

@ -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").

View File

@ -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

View File

@ -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")
}

View File

@ -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 {

View File

@ -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",

View File

@ -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)
}
}

View 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"`
}

View File

@ -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) {

View 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)
}

View File

@ -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))

View File

@ -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

View File

@ -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,
})

View File

@ -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.

View File

@ -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")
}

View File

@ -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
}

View 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)
}

View File

@ -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,
}
}

View File

@ -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")
}

View 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
}

View 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")
}

View File

@ -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)
}
}

View File

@ -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)
})
}

View File

@ -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

View 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"
}

View 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)
})
}
}

View File

@ -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,

View 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")
}

View File

@ -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 {

View File

@ -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"`

View File

@ -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.

View File

@ -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},

View File

@ -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

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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"),

View File

@ -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 {

View File

@ -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 {

View File

@ -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,

View 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
}
}

View File

@ -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

View File

@ -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)

View File

@ -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)
}

View File

@ -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))
}

View File

@ -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

View 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)
}
})
}
}

View File

@ -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")
}

View File

@ -98,6 +98,8 @@ func (h *SettingHandler) GetPublicSettings(c *gin.Context) {
AffiliateEnabled: settings.AffiliateEnabled,
RiskControlEnabled: settings.RiskControlEnabled,
AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests,
})
}

View File

@ -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) {

View File

@ -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})

View File

@ -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})

View File

@ -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)

View File

@ -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)

View File

@ -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(),

View 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)
}
}
})
}
}

View File

@ -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"},
}

View File

@ -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",

View File

@ -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 {

View File

@ -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)
}

View File

@ -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{

View File

@ -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] = &copyCall
stored = &copyCall
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 {

View File

@ -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")
}
}

View File

@ -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"`)
}

View File

@ -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"

View 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
}

View File

@ -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")
}

View File

@ -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")
}

View File

@ -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 {

View File

@ -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"))
}

View File

@ -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,

View File

@ -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.

View File

@ -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",

View 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
}

View 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)
}
})
}
}

View File

@ -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
}

View File

@ -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,

View File

@ -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)
}

View File

@ -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
}

View File

@ -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")

View File

@ -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":

View 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())
}

View File

@ -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),
),
)
}

View File

@ -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) {

View File

@ -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