diff --git a/.gitignore b/.gitignore index 8c7e67f7e2..4ea6ff311b 100644 --- a/.gitignore +++ b/.gitignore @@ -143,3 +143,6 @@ docs/* frontend/coverage/ aicodex output/ + +# Vitest / Vite cache at repo root +.vite/ diff --git a/backend/cmd/server/VERSION b/backend/cmd/server/VERSION index 8ca4b7176c..d26f7fe8bb 100644 --- a/backend/cmd/server/VERSION +++ b/backend/cmd/server/VERSION @@ -1 +1 @@ -0.1.171 +0.1.172 diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index eb4fdfecf8..8969bf4411 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -149,7 +149,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { openAIOAuthService := service.ProvideOpenAIOAuthService(proxyRepository, openAIOAuthClient, privacyClientFactory) openAITokenProvider := service.ProvideOpenAITokenProvider(accountRepository, geminiTokenCache, openAIOAuthService, oAuthRefreshAPI) grokOAuthClient := repository.NewGrokOAuthClient() - grokOAuthService := service.NewGrokOAuthService(proxyRepository, grokOAuthClient) + grokOAuthService := service.ProvideGrokOAuthService(proxyRepository, grokOAuthClient, configConfig, redisClient) grokTokenProvider := service.ProvideGrokTokenProvider(accountRepository, geminiTokenCache, grokOAuthService, oAuthRefreshAPI, tempUnschedCache) openAIGatewayService := service.NewOpenAIGatewayService(accountRepository, usageLogRepository, usageBillingRepository, userRepository, userSubscriptionRepository, userGroupRateRepository, gatewayCache, configConfig, schedulerSnapshotService, concurrencyService, billingService, rateLimitService, billingCacheService, httpUpstream, deferredService, openAITokenProvider, grokTokenProvider, modelPricingResolver, channelService, balanceNotifyService, settingService, serviceUserPlatformQuotaRepository) geminiOAuthClient := repository.NewGeminiOAuthClient(configConfig) @@ -196,11 +196,11 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) { claudeUsageFetcher := repository.NewClaudeUsageFetcher(httpUpstream) antigravityQuotaFetcher := service.NewAntigravityQuotaFetcher(proxyRepository) grokQuotaFetcher := service.NewGrokQuotaFetcher() - grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream, configConfig, usageLogRepository) + grokQuotaService := service.ProvideGrokQuotaService(accountRepository, proxyRepository, grokTokenProvider, httpUpstream, configConfig, usageLogRepository, settingService) openAIQuotaService := service.ProvideOpenAIQuotaService(accountRepository, proxyRepository, openAITokenProvider, privacyClientFactory, openAIGatewayService) usageCache := service.NewUsageCache() accountUsageService := service.ProvideAccountUsageService(accountRepository, usageLogRepository, claudeUsageFetcher, geminiQuotaService, antigravityQuotaFetcher, grokQuotaFetcher, grokQuotaService, openAIQuotaService, usageCache, identityCache, tlsFingerprintProfileService, openAIGatewayService) - accountTestService := service.ProvideAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService, openAIGatewayService) + accountTestService := service.ProvideAccountTestService(accountRepository, geminiTokenProvider, claudeTokenProvider, grokTokenProvider, antigravityGatewayService, httpUpstream, configConfig, tlsFingerprintProfileService, openAIGatewayService, settingService) crsSyncService := service.NewCRSSyncService(accountRepository, proxyRepository, oAuthService, openAIOAuthService, geminiOAuthService, configConfig) accountHandler := admin.ProvideAccountHandler(adminService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, rateLimitService, accountUsageService, accountTestService, concurrencyService, crsSyncService, sessionLimitCache, rpmCache, compositeTokenCacheInvalidator, grokQuotaService) adminAnnouncementHandler := admin.NewAnnouncementHandler(announcementService) diff --git a/backend/ent/group.go b/backend/ent/group.go index 81a3e349c0..3b1d2a4a51 100644 --- a/backend/ent/group.go +++ b/backend/ent/group.go @@ -85,8 +85,18 @@ type Group struct { VideoPrice720p *float64 `json:"video_price_720p,omitempty"` // VideoPrice1080p holds the value of the "video_price_1080p" field. VideoPrice1080p *float64 `json:"video_price_1080p,omitempty"` + // 按模型族和分辨率覆盖视频每秒价格 + VideoModelPrices map[string]map[string]float64 `json:"video_model_prices,omitempty"` // Codex alpha/search 网页搜索单次价格(USD/次);nil 表示使用默认价 0.01(官方 $10/1000 次) WebSearchPricePerCall *float64 `json:"web_search_price_per_call,omitempty"` + // 搜索工具价格 per 1000 calls(web_search 等) + SearchPricePer1k *float64 `json:"search_price_per_1k,omitempty"` + // Voice realtime 每分钟价格(USD) + AudioRealtimePricePerMin *float64 `json:"audio_realtime_price_per_min,omitempty"` + // TTS 每百万字符价格(USD) + AudioTtsPricePerMillionChars *float64 `json:"audio_tts_price_per_million_chars,omitempty"` + // STT 每小时价格(USD) + AudioSttPricePerHour *float64 `json:"audio_stt_price_per_hour,omitempty"` // 是否仅允许 Claude Code 客户端 ClaudeCodeOnly bool `json:"claude_code_only,omitempty"` // 非 Claude Code 请求降级使用的分组 ID @@ -235,11 +245,11 @@ 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, group.FieldModelsListConfig, group.FieldReasoningEffortMappings: + case group.FieldVideoModelPrices, group.FieldModelRouting, group.FieldSupportedModelScopes, group.FieldMessagesDispatchModelConfig, group.FieldModelsListConfig, group.FieldReasoningEffortMappings: values[i] = new([]byte) case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldVideoRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldAllowLive, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet, group.FieldProfitControlEnabled: values[i] = new(sql.NullBool) - case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldBatchImageDiscountMultiplier, group.FieldBatchImageHoldMultiplier, group.FieldVideoRateMultiplier, group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p, group.FieldWebSearchPricePerCall, group.FieldProfitMinMargin, group.FieldProfitSafetyBuffer: + case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldBatchImageDiscountMultiplier, group.FieldBatchImageHoldMultiplier, group.FieldVideoRateMultiplier, group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p, group.FieldWebSearchPricePerCall, group.FieldSearchPricePer1k, group.FieldAudioRealtimePricePerMin, group.FieldAudioTtsPricePerMillionChars, group.FieldAudioSttPricePerHour, group.FieldProfitMinMargin, group.FieldProfitSafetyBuffer: values[i] = new(sql.NullFloat64) case group.FieldID, group.FieldDefaultValidityDays, group.FieldFallbackGroupID, group.FieldFallbackGroupIDOnInvalidRequest, group.FieldSortOrder, group.FieldRpmLimit: values[i] = new(sql.NullInt64) @@ -478,6 +488,14 @@ func (_m *Group) assignValues(columns []string, values []any) error { _m.VideoPrice1080p = new(float64) *_m.VideoPrice1080p = value.Float64 } + case group.FieldVideoModelPrices: + if value, ok := values[i].(*[]byte); !ok { + return fmt.Errorf("unexpected type %T for field video_model_prices", values[i]) + } else if value != nil && len(*value) > 0 { + if err := json.Unmarshal(*value, &_m.VideoModelPrices); err != nil { + return fmt.Errorf("unmarshal field video_model_prices: %w", err) + } + } case group.FieldWebSearchPricePerCall: if value, ok := values[i].(*sql.NullFloat64); !ok { return fmt.Errorf("unexpected type %T for field web_search_price_per_call", values[i]) @@ -485,6 +503,34 @@ func (_m *Group) assignValues(columns []string, values []any) error { _m.WebSearchPricePerCall = new(float64) *_m.WebSearchPricePerCall = value.Float64 } + case group.FieldSearchPricePer1k: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field search_price_per_1k", values[i]) + } else if value.Valid { + _m.SearchPricePer1k = new(float64) + *_m.SearchPricePer1k = value.Float64 + } + case group.FieldAudioRealtimePricePerMin: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field audio_realtime_price_per_min", values[i]) + } else if value.Valid { + _m.AudioRealtimePricePerMin = new(float64) + *_m.AudioRealtimePricePerMin = value.Float64 + } + case group.FieldAudioTtsPricePerMillionChars: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field audio_tts_price_per_million_chars", values[i]) + } else if value.Valid { + _m.AudioTtsPricePerMillionChars = new(float64) + *_m.AudioTtsPricePerMillionChars = value.Float64 + } + case group.FieldAudioSttPricePerHour: + if value, ok := values[i].(*sql.NullFloat64); !ok { + return fmt.Errorf("unexpected type %T for field audio_stt_price_per_hour", values[i]) + } else if value.Valid { + _m.AudioSttPricePerHour = new(float64) + *_m.AudioSttPricePerHour = value.Float64 + } case group.FieldClaudeCodeOnly: if value, ok := values[i].(*sql.NullBool); !ok { return fmt.Errorf("unexpected type %T for field claude_code_only", values[i]) @@ -822,11 +868,34 @@ func (_m *Group) String() string { builder.WriteString(fmt.Sprintf("%v", *v)) } builder.WriteString(", ") + builder.WriteString("video_model_prices=") + builder.WriteString(fmt.Sprintf("%v", _m.VideoModelPrices)) + builder.WriteString(", ") if v := _m.WebSearchPricePerCall; v != nil { builder.WriteString("web_search_price_per_call=") builder.WriteString(fmt.Sprintf("%v", *v)) } builder.WriteString(", ") + if v := _m.SearchPricePer1k; v != nil { + builder.WriteString("search_price_per_1k=") + builder.WriteString(fmt.Sprintf("%v", *v)) + } + builder.WriteString(", ") + if v := _m.AudioRealtimePricePerMin; v != nil { + builder.WriteString("audio_realtime_price_per_min=") + builder.WriteString(fmt.Sprintf("%v", *v)) + } + builder.WriteString(", ") + if v := _m.AudioTtsPricePerMillionChars; v != nil { + builder.WriteString("audio_tts_price_per_million_chars=") + builder.WriteString(fmt.Sprintf("%v", *v)) + } + builder.WriteString(", ") + if v := _m.AudioSttPricePerHour; v != nil { + builder.WriteString("audio_stt_price_per_hour=") + builder.WriteString(fmt.Sprintf("%v", *v)) + } + builder.WriteString(", ") builder.WriteString("claude_code_only=") builder.WriteString(fmt.Sprintf("%v", _m.ClaudeCodeOnly)) builder.WriteString(", ") diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go index 35d6de1336..76a07b96a0 100644 --- a/backend/ent/group/group.go +++ b/backend/ent/group/group.go @@ -82,8 +82,18 @@ const ( FieldVideoPrice720p = "video_price_720p" // FieldVideoPrice1080p holds the string denoting the video_price_1080p field in the database. FieldVideoPrice1080p = "video_price_1080p" + // FieldVideoModelPrices holds the string denoting the video_model_prices field in the database. + FieldVideoModelPrices = "video_model_prices" // FieldWebSearchPricePerCall holds the string denoting the web_search_price_per_call field in the database. FieldWebSearchPricePerCall = "web_search_price_per_call" + // FieldSearchPricePer1k holds the string denoting the search_price_per_1k field in the database. + FieldSearchPricePer1k = "search_price_per_1k" + // FieldAudioRealtimePricePerMin holds the string denoting the audio_realtime_price_per_min field in the database. + FieldAudioRealtimePricePerMin = "audio_realtime_price_per_min" + // FieldAudioTtsPricePerMillionChars holds the string denoting the audio_tts_price_per_million_chars field in the database. + FieldAudioTtsPricePerMillionChars = "audio_tts_price_per_million_chars" + // FieldAudioSttPricePerHour holds the string denoting the audio_stt_price_per_hour field in the database. + FieldAudioSttPricePerHour = "audio_stt_price_per_hour" // FieldClaudeCodeOnly holds the string denoting the claude_code_only field in the database. FieldClaudeCodeOnly = "claude_code_only" // FieldFallbackGroupID holds the string denoting the fallback_group_id field in the database. @@ -234,7 +244,12 @@ var Columns = []string{ FieldVideoPrice480p, FieldVideoPrice720p, FieldVideoPrice1080p, + FieldVideoModelPrices, FieldWebSearchPricePerCall, + FieldSearchPricePer1k, + FieldAudioRealtimePricePerMin, + FieldAudioTtsPricePerMillionChars, + FieldAudioSttPricePerHour, FieldClaudeCodeOnly, FieldFallbackGroupID, FieldFallbackGroupIDOnInvalidRequest, @@ -341,6 +356,14 @@ var ( DefaultVideoRateIndependent bool // DefaultVideoRateMultiplier holds the default value on creation for the "video_rate_multiplier" field. DefaultVideoRateMultiplier float64 + // SearchPricePer1kValidator is a validator for the "search_price_per_1k" field. It is called by the builders before save. + SearchPricePer1kValidator func(float64) error + // AudioRealtimePricePerMinValidator is a validator for the "audio_realtime_price_per_min" field. It is called by the builders before save. + AudioRealtimePricePerMinValidator func(float64) error + // AudioTtsPricePerMillionCharsValidator is a validator for the "audio_tts_price_per_million_chars" field. It is called by the builders before save. + AudioTtsPricePerMillionCharsValidator func(float64) error + // AudioSttPricePerHourValidator is a validator for the "audio_stt_price_per_hour" field. It is called by the builders before save. + AudioSttPricePerHourValidator func(float64) error // DefaultClaudeCodeOnly holds the default value on creation for the "claude_code_only" field. DefaultClaudeCodeOnly bool // DefaultModelRoutingEnabled holds the default value on creation for the "model_routing_enabled" field. @@ -561,6 +584,26 @@ func ByWebSearchPricePerCall(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldWebSearchPricePerCall, opts...).ToFunc() } +// BySearchPricePer1k orders the results by the search_price_per_1k field. +func BySearchPricePer1k(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldSearchPricePer1k, opts...).ToFunc() +} + +// ByAudioRealtimePricePerMin orders the results by the audio_realtime_price_per_min field. +func ByAudioRealtimePricePerMin(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldAudioRealtimePricePerMin, opts...).ToFunc() +} + +// ByAudioTtsPricePerMillionChars orders the results by the audio_tts_price_per_million_chars field. +func ByAudioTtsPricePerMillionChars(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldAudioTtsPricePerMillionChars, opts...).ToFunc() +} + +// ByAudioSttPricePerHour orders the results by the audio_stt_price_per_hour field. +func ByAudioSttPricePerHour(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldAudioSttPricePerHour, opts...).ToFunc() +} + // ByClaudeCodeOnly orders the results by the claude_code_only field. func ByClaudeCodeOnly(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldClaudeCodeOnly, opts...).ToFunc() diff --git a/backend/ent/group/where.go b/backend/ent/group/where.go index 64e52ba80a..630ec582e8 100644 --- a/backend/ent/group/where.go +++ b/backend/ent/group/where.go @@ -225,6 +225,26 @@ func WebSearchPricePerCall(v float64) predicate.Group { return predicate.Group(sql.FieldEQ(FieldWebSearchPricePerCall, v)) } +// SearchPricePer1k applies equality check predicate on the "search_price_per_1k" field. It's identical to SearchPricePer1kEQ. +func SearchPricePer1k(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldSearchPricePer1k, v)) +} + +// AudioRealtimePricePerMin applies equality check predicate on the "audio_realtime_price_per_min" field. It's identical to AudioRealtimePricePerMinEQ. +func AudioRealtimePricePerMin(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAudioRealtimePricePerMin, v)) +} + +// AudioTtsPricePerMillionChars applies equality check predicate on the "audio_tts_price_per_million_chars" field. It's identical to AudioTtsPricePerMillionCharsEQ. +func AudioTtsPricePerMillionChars(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAudioTtsPricePerMillionChars, v)) +} + +// AudioSttPricePerHour applies equality check predicate on the "audio_stt_price_per_hour" field. It's identical to AudioSttPricePerHourEQ. +func AudioSttPricePerHour(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAudioSttPricePerHour, v)) +} + // ClaudeCodeOnly applies equality check predicate on the "claude_code_only" field. It's identical to ClaudeCodeOnlyEQ. func ClaudeCodeOnly(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldClaudeCodeOnly, v)) @@ -1765,6 +1785,16 @@ func VideoPrice1080pNotNil() predicate.Group { return predicate.Group(sql.FieldNotNull(FieldVideoPrice1080p)) } +// VideoModelPricesIsNil applies the IsNil predicate on the "video_model_prices" field. +func VideoModelPricesIsNil() predicate.Group { + return predicate.Group(sql.FieldIsNull(FieldVideoModelPrices)) +} + +// VideoModelPricesNotNil applies the NotNil predicate on the "video_model_prices" field. +func VideoModelPricesNotNil() predicate.Group { + return predicate.Group(sql.FieldNotNull(FieldVideoModelPrices)) +} + // WebSearchPricePerCallEQ applies the EQ predicate on the "web_search_price_per_call" field. func WebSearchPricePerCallEQ(v float64) predicate.Group { return predicate.Group(sql.FieldEQ(FieldWebSearchPricePerCall, v)) @@ -1815,6 +1845,206 @@ func WebSearchPricePerCallNotNil() predicate.Group { return predicate.Group(sql.FieldNotNull(FieldWebSearchPricePerCall)) } +// SearchPricePer1kEQ applies the EQ predicate on the "search_price_per_1k" field. +func SearchPricePer1kEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldSearchPricePer1k, v)) +} + +// SearchPricePer1kNEQ applies the NEQ predicate on the "search_price_per_1k" field. +func SearchPricePer1kNEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldSearchPricePer1k, v)) +} + +// SearchPricePer1kIn applies the In predicate on the "search_price_per_1k" field. +func SearchPricePer1kIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldIn(FieldSearchPricePer1k, vs...)) +} + +// SearchPricePer1kNotIn applies the NotIn predicate on the "search_price_per_1k" field. +func SearchPricePer1kNotIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldNotIn(FieldSearchPricePer1k, vs...)) +} + +// SearchPricePer1kGT applies the GT predicate on the "search_price_per_1k" field. +func SearchPricePer1kGT(v float64) predicate.Group { + return predicate.Group(sql.FieldGT(FieldSearchPricePer1k, v)) +} + +// SearchPricePer1kGTE applies the GTE predicate on the "search_price_per_1k" field. +func SearchPricePer1kGTE(v float64) predicate.Group { + return predicate.Group(sql.FieldGTE(FieldSearchPricePer1k, v)) +} + +// SearchPricePer1kLT applies the LT predicate on the "search_price_per_1k" field. +func SearchPricePer1kLT(v float64) predicate.Group { + return predicate.Group(sql.FieldLT(FieldSearchPricePer1k, v)) +} + +// SearchPricePer1kLTE applies the LTE predicate on the "search_price_per_1k" field. +func SearchPricePer1kLTE(v float64) predicate.Group { + return predicate.Group(sql.FieldLTE(FieldSearchPricePer1k, v)) +} + +// SearchPricePer1kIsNil applies the IsNil predicate on the "search_price_per_1k" field. +func SearchPricePer1kIsNil() predicate.Group { + return predicate.Group(sql.FieldIsNull(FieldSearchPricePer1k)) +} + +// SearchPricePer1kNotNil applies the NotNil predicate on the "search_price_per_1k" field. +func SearchPricePer1kNotNil() predicate.Group { + return predicate.Group(sql.FieldNotNull(FieldSearchPricePer1k)) +} + +// AudioRealtimePricePerMinEQ applies the EQ predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAudioRealtimePricePerMin, v)) +} + +// AudioRealtimePricePerMinNEQ applies the NEQ predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinNEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldAudioRealtimePricePerMin, v)) +} + +// AudioRealtimePricePerMinIn applies the In predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldIn(FieldAudioRealtimePricePerMin, vs...)) +} + +// AudioRealtimePricePerMinNotIn applies the NotIn predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinNotIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldNotIn(FieldAudioRealtimePricePerMin, vs...)) +} + +// AudioRealtimePricePerMinGT applies the GT predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinGT(v float64) predicate.Group { + return predicate.Group(sql.FieldGT(FieldAudioRealtimePricePerMin, v)) +} + +// AudioRealtimePricePerMinGTE applies the GTE predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinGTE(v float64) predicate.Group { + return predicate.Group(sql.FieldGTE(FieldAudioRealtimePricePerMin, v)) +} + +// AudioRealtimePricePerMinLT applies the LT predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinLT(v float64) predicate.Group { + return predicate.Group(sql.FieldLT(FieldAudioRealtimePricePerMin, v)) +} + +// AudioRealtimePricePerMinLTE applies the LTE predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinLTE(v float64) predicate.Group { + return predicate.Group(sql.FieldLTE(FieldAudioRealtimePricePerMin, v)) +} + +// AudioRealtimePricePerMinIsNil applies the IsNil predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinIsNil() predicate.Group { + return predicate.Group(sql.FieldIsNull(FieldAudioRealtimePricePerMin)) +} + +// AudioRealtimePricePerMinNotNil applies the NotNil predicate on the "audio_realtime_price_per_min" field. +func AudioRealtimePricePerMinNotNil() predicate.Group { + return predicate.Group(sql.FieldNotNull(FieldAudioRealtimePricePerMin)) +} + +// AudioTtsPricePerMillionCharsEQ applies the EQ predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAudioTtsPricePerMillionChars, v)) +} + +// AudioTtsPricePerMillionCharsNEQ applies the NEQ predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsNEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldAudioTtsPricePerMillionChars, v)) +} + +// AudioTtsPricePerMillionCharsIn applies the In predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldIn(FieldAudioTtsPricePerMillionChars, vs...)) +} + +// AudioTtsPricePerMillionCharsNotIn applies the NotIn predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsNotIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldNotIn(FieldAudioTtsPricePerMillionChars, vs...)) +} + +// AudioTtsPricePerMillionCharsGT applies the GT predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsGT(v float64) predicate.Group { + return predicate.Group(sql.FieldGT(FieldAudioTtsPricePerMillionChars, v)) +} + +// AudioTtsPricePerMillionCharsGTE applies the GTE predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsGTE(v float64) predicate.Group { + return predicate.Group(sql.FieldGTE(FieldAudioTtsPricePerMillionChars, v)) +} + +// AudioTtsPricePerMillionCharsLT applies the LT predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsLT(v float64) predicate.Group { + return predicate.Group(sql.FieldLT(FieldAudioTtsPricePerMillionChars, v)) +} + +// AudioTtsPricePerMillionCharsLTE applies the LTE predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsLTE(v float64) predicate.Group { + return predicate.Group(sql.FieldLTE(FieldAudioTtsPricePerMillionChars, v)) +} + +// AudioTtsPricePerMillionCharsIsNil applies the IsNil predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsIsNil() predicate.Group { + return predicate.Group(sql.FieldIsNull(FieldAudioTtsPricePerMillionChars)) +} + +// AudioTtsPricePerMillionCharsNotNil applies the NotNil predicate on the "audio_tts_price_per_million_chars" field. +func AudioTtsPricePerMillionCharsNotNil() predicate.Group { + return predicate.Group(sql.FieldNotNull(FieldAudioTtsPricePerMillionChars)) +} + +// AudioSttPricePerHourEQ applies the EQ predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAudioSttPricePerHour, v)) +} + +// AudioSttPricePerHourNEQ applies the NEQ predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourNEQ(v float64) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldAudioSttPricePerHour, v)) +} + +// AudioSttPricePerHourIn applies the In predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldIn(FieldAudioSttPricePerHour, vs...)) +} + +// AudioSttPricePerHourNotIn applies the NotIn predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourNotIn(vs ...float64) predicate.Group { + return predicate.Group(sql.FieldNotIn(FieldAudioSttPricePerHour, vs...)) +} + +// AudioSttPricePerHourGT applies the GT predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourGT(v float64) predicate.Group { + return predicate.Group(sql.FieldGT(FieldAudioSttPricePerHour, v)) +} + +// AudioSttPricePerHourGTE applies the GTE predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourGTE(v float64) predicate.Group { + return predicate.Group(sql.FieldGTE(FieldAudioSttPricePerHour, v)) +} + +// AudioSttPricePerHourLT applies the LT predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourLT(v float64) predicate.Group { + return predicate.Group(sql.FieldLT(FieldAudioSttPricePerHour, v)) +} + +// AudioSttPricePerHourLTE applies the LTE predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourLTE(v float64) predicate.Group { + return predicate.Group(sql.FieldLTE(FieldAudioSttPricePerHour, v)) +} + +// AudioSttPricePerHourIsNil applies the IsNil predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourIsNil() predicate.Group { + return predicate.Group(sql.FieldIsNull(FieldAudioSttPricePerHour)) +} + +// AudioSttPricePerHourNotNil applies the NotNil predicate on the "audio_stt_price_per_hour" field. +func AudioSttPricePerHourNotNil() predicate.Group { + return predicate.Group(sql.FieldNotNull(FieldAudioSttPricePerHour)) +} + // ClaudeCodeOnlyEQ applies the EQ predicate on the "claude_code_only" field. func ClaudeCodeOnlyEQ(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldClaudeCodeOnly, v)) diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go index 54e34a32f6..3246b92a50 100644 --- a/backend/ent/group_create.go +++ b/backend/ent/group_create.go @@ -483,6 +483,12 @@ func (_c *GroupCreate) SetNillableVideoPrice1080p(v *float64) *GroupCreate { return _c } +// SetVideoModelPrices sets the "video_model_prices" field. +func (_c *GroupCreate) SetVideoModelPrices(v map[string]map[string]float64) *GroupCreate { + _c.mutation.SetVideoModelPrices(v) + return _c +} + // SetWebSearchPricePerCall sets the "web_search_price_per_call" field. func (_c *GroupCreate) SetWebSearchPricePerCall(v float64) *GroupCreate { _c.mutation.SetWebSearchPricePerCall(v) @@ -497,6 +503,62 @@ func (_c *GroupCreate) SetNillableWebSearchPricePerCall(v *float64) *GroupCreate return _c } +// SetSearchPricePer1k sets the "search_price_per_1k" field. +func (_c *GroupCreate) SetSearchPricePer1k(v float64) *GroupCreate { + _c.mutation.SetSearchPricePer1k(v) + return _c +} + +// SetNillableSearchPricePer1k sets the "search_price_per_1k" field if the given value is not nil. +func (_c *GroupCreate) SetNillableSearchPricePer1k(v *float64) *GroupCreate { + if v != nil { + _c.SetSearchPricePer1k(*v) + } + return _c +} + +// SetAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field. +func (_c *GroupCreate) SetAudioRealtimePricePerMin(v float64) *GroupCreate { + _c.mutation.SetAudioRealtimePricePerMin(v) + return _c +} + +// SetNillableAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field if the given value is not nil. +func (_c *GroupCreate) SetNillableAudioRealtimePricePerMin(v *float64) *GroupCreate { + if v != nil { + _c.SetAudioRealtimePricePerMin(*v) + } + return _c +} + +// SetAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field. +func (_c *GroupCreate) SetAudioTtsPricePerMillionChars(v float64) *GroupCreate { + _c.mutation.SetAudioTtsPricePerMillionChars(v) + return _c +} + +// SetNillableAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field if the given value is not nil. +func (_c *GroupCreate) SetNillableAudioTtsPricePerMillionChars(v *float64) *GroupCreate { + if v != nil { + _c.SetAudioTtsPricePerMillionChars(*v) + } + return _c +} + +// SetAudioSttPricePerHour sets the "audio_stt_price_per_hour" field. +func (_c *GroupCreate) SetAudioSttPricePerHour(v float64) *GroupCreate { + _c.mutation.SetAudioSttPricePerHour(v) + return _c +} + +// SetNillableAudioSttPricePerHour sets the "audio_stt_price_per_hour" field if the given value is not nil. +func (_c *GroupCreate) SetNillableAudioSttPricePerHour(v *float64) *GroupCreate { + if v != nil { + _c.SetAudioSttPricePerHour(*v) + } + return _c +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_c *GroupCreate) SetClaudeCodeOnly(v bool) *GroupCreate { _c.mutation.SetClaudeCodeOnly(v) @@ -1155,6 +1217,26 @@ func (_c *GroupCreate) check() error { if _, ok := _c.mutation.VideoRateMultiplier(); !ok { return &ValidationError{Name: "video_rate_multiplier", err: errors.New(`ent: missing required field "Group.video_rate_multiplier"`)} } + if v, ok := _c.mutation.SearchPricePer1k(); ok { + if err := group.SearchPricePer1kValidator(v); err != nil { + return &ValidationError{Name: "search_price_per_1k", err: fmt.Errorf(`ent: validator failed for field "Group.search_price_per_1k": %w`, err)} + } + } + if v, ok := _c.mutation.AudioRealtimePricePerMin(); ok { + if err := group.AudioRealtimePricePerMinValidator(v); err != nil { + return &ValidationError{Name: "audio_realtime_price_per_min", err: fmt.Errorf(`ent: validator failed for field "Group.audio_realtime_price_per_min": %w`, err)} + } + } + if v, ok := _c.mutation.AudioTtsPricePerMillionChars(); ok { + if err := group.AudioTtsPricePerMillionCharsValidator(v); err != nil { + return &ValidationError{Name: "audio_tts_price_per_million_chars", err: fmt.Errorf(`ent: validator failed for field "Group.audio_tts_price_per_million_chars": %w`, err)} + } + } + if v, ok := _c.mutation.AudioSttPricePerHour(); ok { + if err := group.AudioSttPricePerHourValidator(v); err != nil { + return &ValidationError{Name: "audio_stt_price_per_hour", err: fmt.Errorf(`ent: validator failed for field "Group.audio_stt_price_per_hour": %w`, err)} + } + } if _, ok := _c.mutation.ClaudeCodeOnly(); !ok { return &ValidationError{Name: "claude_code_only", err: errors.New(`ent: missing required field "Group.claude_code_only"`)} } @@ -1378,10 +1460,30 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldVideoPrice1080p, field.TypeFloat64, value) _node.VideoPrice1080p = &value } + if value, ok := _c.mutation.VideoModelPrices(); ok { + _spec.SetField(group.FieldVideoModelPrices, field.TypeJSON, value) + _node.VideoModelPrices = value + } if value, ok := _c.mutation.WebSearchPricePerCall(); ok { _spec.SetField(group.FieldWebSearchPricePerCall, field.TypeFloat64, value) _node.WebSearchPricePerCall = &value } + if value, ok := _c.mutation.SearchPricePer1k(); ok { + _spec.SetField(group.FieldSearchPricePer1k, field.TypeFloat64, value) + _node.SearchPricePer1k = &value + } + if value, ok := _c.mutation.AudioRealtimePricePerMin(); ok { + _spec.SetField(group.FieldAudioRealtimePricePerMin, field.TypeFloat64, value) + _node.AudioRealtimePricePerMin = &value + } + if value, ok := _c.mutation.AudioTtsPricePerMillionChars(); ok { + _spec.SetField(group.FieldAudioTtsPricePerMillionChars, field.TypeFloat64, value) + _node.AudioTtsPricePerMillionChars = &value + } + if value, ok := _c.mutation.AudioSttPricePerHour(); ok { + _spec.SetField(group.FieldAudioSttPricePerHour, field.TypeFloat64, value) + _node.AudioSttPricePerHour = &value + } if value, ok := _c.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) _node.ClaudeCodeOnly = value @@ -2156,6 +2258,24 @@ func (u *GroupUpsert) ClearVideoPrice1080p() *GroupUpsert { return u } +// SetVideoModelPrices sets the "video_model_prices" field. +func (u *GroupUpsert) SetVideoModelPrices(v map[string]map[string]float64) *GroupUpsert { + u.Set(group.FieldVideoModelPrices, v) + return u +} + +// UpdateVideoModelPrices sets the "video_model_prices" field to the value that was provided on create. +func (u *GroupUpsert) UpdateVideoModelPrices() *GroupUpsert { + u.SetExcluded(group.FieldVideoModelPrices) + return u +} + +// ClearVideoModelPrices clears the value of the "video_model_prices" field. +func (u *GroupUpsert) ClearVideoModelPrices() *GroupUpsert { + u.SetNull(group.FieldVideoModelPrices) + return u +} + // SetWebSearchPricePerCall sets the "web_search_price_per_call" field. func (u *GroupUpsert) SetWebSearchPricePerCall(v float64) *GroupUpsert { u.Set(group.FieldWebSearchPricePerCall, v) @@ -2180,6 +2300,102 @@ func (u *GroupUpsert) ClearWebSearchPricePerCall() *GroupUpsert { return u } +// SetSearchPricePer1k sets the "search_price_per_1k" field. +func (u *GroupUpsert) SetSearchPricePer1k(v float64) *GroupUpsert { + u.Set(group.FieldSearchPricePer1k, v) + return u +} + +// UpdateSearchPricePer1k sets the "search_price_per_1k" field to the value that was provided on create. +func (u *GroupUpsert) UpdateSearchPricePer1k() *GroupUpsert { + u.SetExcluded(group.FieldSearchPricePer1k) + return u +} + +// AddSearchPricePer1k adds v to the "search_price_per_1k" field. +func (u *GroupUpsert) AddSearchPricePer1k(v float64) *GroupUpsert { + u.Add(group.FieldSearchPricePer1k, v) + return u +} + +// ClearSearchPricePer1k clears the value of the "search_price_per_1k" field. +func (u *GroupUpsert) ClearSearchPricePer1k() *GroupUpsert { + u.SetNull(group.FieldSearchPricePer1k) + return u +} + +// SetAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field. +func (u *GroupUpsert) SetAudioRealtimePricePerMin(v float64) *GroupUpsert { + u.Set(group.FieldAudioRealtimePricePerMin, v) + return u +} + +// UpdateAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field to the value that was provided on create. +func (u *GroupUpsert) UpdateAudioRealtimePricePerMin() *GroupUpsert { + u.SetExcluded(group.FieldAudioRealtimePricePerMin) + return u +} + +// AddAudioRealtimePricePerMin adds v to the "audio_realtime_price_per_min" field. +func (u *GroupUpsert) AddAudioRealtimePricePerMin(v float64) *GroupUpsert { + u.Add(group.FieldAudioRealtimePricePerMin, v) + return u +} + +// ClearAudioRealtimePricePerMin clears the value of the "audio_realtime_price_per_min" field. +func (u *GroupUpsert) ClearAudioRealtimePricePerMin() *GroupUpsert { + u.SetNull(group.FieldAudioRealtimePricePerMin) + return u +} + +// SetAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field. +func (u *GroupUpsert) SetAudioTtsPricePerMillionChars(v float64) *GroupUpsert { + u.Set(group.FieldAudioTtsPricePerMillionChars, v) + return u +} + +// UpdateAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field to the value that was provided on create. +func (u *GroupUpsert) UpdateAudioTtsPricePerMillionChars() *GroupUpsert { + u.SetExcluded(group.FieldAudioTtsPricePerMillionChars) + return u +} + +// AddAudioTtsPricePerMillionChars adds v to the "audio_tts_price_per_million_chars" field. +func (u *GroupUpsert) AddAudioTtsPricePerMillionChars(v float64) *GroupUpsert { + u.Add(group.FieldAudioTtsPricePerMillionChars, v) + return u +} + +// ClearAudioTtsPricePerMillionChars clears the value of the "audio_tts_price_per_million_chars" field. +func (u *GroupUpsert) ClearAudioTtsPricePerMillionChars() *GroupUpsert { + u.SetNull(group.FieldAudioTtsPricePerMillionChars) + return u +} + +// SetAudioSttPricePerHour sets the "audio_stt_price_per_hour" field. +func (u *GroupUpsert) SetAudioSttPricePerHour(v float64) *GroupUpsert { + u.Set(group.FieldAudioSttPricePerHour, v) + return u +} + +// UpdateAudioSttPricePerHour sets the "audio_stt_price_per_hour" field to the value that was provided on create. +func (u *GroupUpsert) UpdateAudioSttPricePerHour() *GroupUpsert { + u.SetExcluded(group.FieldAudioSttPricePerHour) + return u +} + +// AddAudioSttPricePerHour adds v to the "audio_stt_price_per_hour" field. +func (u *GroupUpsert) AddAudioSttPricePerHour(v float64) *GroupUpsert { + u.Add(group.FieldAudioSttPricePerHour, v) + return u +} + +// ClearAudioSttPricePerHour clears the value of the "audio_stt_price_per_hour" field. +func (u *GroupUpsert) ClearAudioSttPricePerHour() *GroupUpsert { + u.SetNull(group.FieldAudioSttPricePerHour) + return u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsert) SetClaudeCodeOnly(v bool) *GroupUpsert { u.Set(group.FieldClaudeCodeOnly, v) @@ -3157,6 +3373,27 @@ func (u *GroupUpsertOne) ClearVideoPrice1080p() *GroupUpsertOne { }) } +// SetVideoModelPrices sets the "video_model_prices" field. +func (u *GroupUpsertOne) SetVideoModelPrices(v map[string]map[string]float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetVideoModelPrices(v) + }) +} + +// UpdateVideoModelPrices sets the "video_model_prices" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateVideoModelPrices() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateVideoModelPrices() + }) +} + +// ClearVideoModelPrices clears the value of the "video_model_prices" field. +func (u *GroupUpsertOne) ClearVideoModelPrices() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.ClearVideoModelPrices() + }) +} + // SetWebSearchPricePerCall sets the "web_search_price_per_call" field. func (u *GroupUpsertOne) SetWebSearchPricePerCall(v float64) *GroupUpsertOne { return u.Update(func(s *GroupUpsert) { @@ -3185,6 +3422,118 @@ func (u *GroupUpsertOne) ClearWebSearchPricePerCall() *GroupUpsertOne { }) } +// SetSearchPricePer1k sets the "search_price_per_1k" field. +func (u *GroupUpsertOne) SetSearchPricePer1k(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetSearchPricePer1k(v) + }) +} + +// AddSearchPricePer1k adds v to the "search_price_per_1k" field. +func (u *GroupUpsertOne) AddSearchPricePer1k(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.AddSearchPricePer1k(v) + }) +} + +// UpdateSearchPricePer1k sets the "search_price_per_1k" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateSearchPricePer1k() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateSearchPricePer1k() + }) +} + +// ClearSearchPricePer1k clears the value of the "search_price_per_1k" field. +func (u *GroupUpsertOne) ClearSearchPricePer1k() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.ClearSearchPricePer1k() + }) +} + +// SetAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field. +func (u *GroupUpsertOne) SetAudioRealtimePricePerMin(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetAudioRealtimePricePerMin(v) + }) +} + +// AddAudioRealtimePricePerMin adds v to the "audio_realtime_price_per_min" field. +func (u *GroupUpsertOne) AddAudioRealtimePricePerMin(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.AddAudioRealtimePricePerMin(v) + }) +} + +// UpdateAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateAudioRealtimePricePerMin() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateAudioRealtimePricePerMin() + }) +} + +// ClearAudioRealtimePricePerMin clears the value of the "audio_realtime_price_per_min" field. +func (u *GroupUpsertOne) ClearAudioRealtimePricePerMin() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.ClearAudioRealtimePricePerMin() + }) +} + +// SetAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field. +func (u *GroupUpsertOne) SetAudioTtsPricePerMillionChars(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetAudioTtsPricePerMillionChars(v) + }) +} + +// AddAudioTtsPricePerMillionChars adds v to the "audio_tts_price_per_million_chars" field. +func (u *GroupUpsertOne) AddAudioTtsPricePerMillionChars(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.AddAudioTtsPricePerMillionChars(v) + }) +} + +// UpdateAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateAudioTtsPricePerMillionChars() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateAudioTtsPricePerMillionChars() + }) +} + +// ClearAudioTtsPricePerMillionChars clears the value of the "audio_tts_price_per_million_chars" field. +func (u *GroupUpsertOne) ClearAudioTtsPricePerMillionChars() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.ClearAudioTtsPricePerMillionChars() + }) +} + +// SetAudioSttPricePerHour sets the "audio_stt_price_per_hour" field. +func (u *GroupUpsertOne) SetAudioSttPricePerHour(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetAudioSttPricePerHour(v) + }) +} + +// AddAudioSttPricePerHour adds v to the "audio_stt_price_per_hour" field. +func (u *GroupUpsertOne) AddAudioSttPricePerHour(v float64) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.AddAudioSttPricePerHour(v) + }) +} + +// UpdateAudioSttPricePerHour sets the "audio_stt_price_per_hour" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateAudioSttPricePerHour() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateAudioSttPricePerHour() + }) +} + +// ClearAudioSttPricePerHour clears the value of the "audio_stt_price_per_hour" field. +func (u *GroupUpsertOne) ClearAudioSttPricePerHour() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.ClearAudioSttPricePerHour() + }) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsertOne) SetClaudeCodeOnly(v bool) *GroupUpsertOne { return u.Update(func(s *GroupUpsert) { @@ -4379,6 +4728,27 @@ func (u *GroupUpsertBulk) ClearVideoPrice1080p() *GroupUpsertBulk { }) } +// SetVideoModelPrices sets the "video_model_prices" field. +func (u *GroupUpsertBulk) SetVideoModelPrices(v map[string]map[string]float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetVideoModelPrices(v) + }) +} + +// UpdateVideoModelPrices sets the "video_model_prices" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateVideoModelPrices() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateVideoModelPrices() + }) +} + +// ClearVideoModelPrices clears the value of the "video_model_prices" field. +func (u *GroupUpsertBulk) ClearVideoModelPrices() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.ClearVideoModelPrices() + }) +} + // SetWebSearchPricePerCall sets the "web_search_price_per_call" field. func (u *GroupUpsertBulk) SetWebSearchPricePerCall(v float64) *GroupUpsertBulk { return u.Update(func(s *GroupUpsert) { @@ -4407,6 +4777,118 @@ func (u *GroupUpsertBulk) ClearWebSearchPricePerCall() *GroupUpsertBulk { }) } +// SetSearchPricePer1k sets the "search_price_per_1k" field. +func (u *GroupUpsertBulk) SetSearchPricePer1k(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetSearchPricePer1k(v) + }) +} + +// AddSearchPricePer1k adds v to the "search_price_per_1k" field. +func (u *GroupUpsertBulk) AddSearchPricePer1k(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.AddSearchPricePer1k(v) + }) +} + +// UpdateSearchPricePer1k sets the "search_price_per_1k" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateSearchPricePer1k() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateSearchPricePer1k() + }) +} + +// ClearSearchPricePer1k clears the value of the "search_price_per_1k" field. +func (u *GroupUpsertBulk) ClearSearchPricePer1k() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.ClearSearchPricePer1k() + }) +} + +// SetAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field. +func (u *GroupUpsertBulk) SetAudioRealtimePricePerMin(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetAudioRealtimePricePerMin(v) + }) +} + +// AddAudioRealtimePricePerMin adds v to the "audio_realtime_price_per_min" field. +func (u *GroupUpsertBulk) AddAudioRealtimePricePerMin(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.AddAudioRealtimePricePerMin(v) + }) +} + +// UpdateAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateAudioRealtimePricePerMin() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateAudioRealtimePricePerMin() + }) +} + +// ClearAudioRealtimePricePerMin clears the value of the "audio_realtime_price_per_min" field. +func (u *GroupUpsertBulk) ClearAudioRealtimePricePerMin() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.ClearAudioRealtimePricePerMin() + }) +} + +// SetAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field. +func (u *GroupUpsertBulk) SetAudioTtsPricePerMillionChars(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetAudioTtsPricePerMillionChars(v) + }) +} + +// AddAudioTtsPricePerMillionChars adds v to the "audio_tts_price_per_million_chars" field. +func (u *GroupUpsertBulk) AddAudioTtsPricePerMillionChars(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.AddAudioTtsPricePerMillionChars(v) + }) +} + +// UpdateAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateAudioTtsPricePerMillionChars() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateAudioTtsPricePerMillionChars() + }) +} + +// ClearAudioTtsPricePerMillionChars clears the value of the "audio_tts_price_per_million_chars" field. +func (u *GroupUpsertBulk) ClearAudioTtsPricePerMillionChars() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.ClearAudioTtsPricePerMillionChars() + }) +} + +// SetAudioSttPricePerHour sets the "audio_stt_price_per_hour" field. +func (u *GroupUpsertBulk) SetAudioSttPricePerHour(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetAudioSttPricePerHour(v) + }) +} + +// AddAudioSttPricePerHour adds v to the "audio_stt_price_per_hour" field. +func (u *GroupUpsertBulk) AddAudioSttPricePerHour(v float64) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.AddAudioSttPricePerHour(v) + }) +} + +// UpdateAudioSttPricePerHour sets the "audio_stt_price_per_hour" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateAudioSttPricePerHour() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateAudioSttPricePerHour() + }) +} + +// ClearAudioSttPricePerHour clears the value of the "audio_stt_price_per_hour" field. +func (u *GroupUpsertBulk) ClearAudioSttPricePerHour() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.ClearAudioSttPricePerHour() + }) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (u *GroupUpsertBulk) SetClaudeCodeOnly(v bool) *GroupUpsertBulk { return u.Update(func(s *GroupUpsert) { diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go index c2fdea3f51..de5474e069 100644 --- a/backend/ent/group_update.go +++ b/backend/ent/group_update.go @@ -640,6 +640,18 @@ func (_u *GroupUpdate) ClearVideoPrice1080p() *GroupUpdate { return _u } +// SetVideoModelPrices sets the "video_model_prices" field. +func (_u *GroupUpdate) SetVideoModelPrices(v map[string]map[string]float64) *GroupUpdate { + _u.mutation.SetVideoModelPrices(v) + return _u +} + +// ClearVideoModelPrices clears the value of the "video_model_prices" field. +func (_u *GroupUpdate) ClearVideoModelPrices() *GroupUpdate { + _u.mutation.ClearVideoModelPrices() + return _u +} + // SetWebSearchPricePerCall sets the "web_search_price_per_call" field. func (_u *GroupUpdate) SetWebSearchPricePerCall(v float64) *GroupUpdate { _u.mutation.ResetWebSearchPricePerCall() @@ -667,6 +679,114 @@ func (_u *GroupUpdate) ClearWebSearchPricePerCall() *GroupUpdate { return _u } +// SetSearchPricePer1k sets the "search_price_per_1k" field. +func (_u *GroupUpdate) SetSearchPricePer1k(v float64) *GroupUpdate { + _u.mutation.ResetSearchPricePer1k() + _u.mutation.SetSearchPricePer1k(v) + return _u +} + +// SetNillableSearchPricePer1k sets the "search_price_per_1k" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableSearchPricePer1k(v *float64) *GroupUpdate { + if v != nil { + _u.SetSearchPricePer1k(*v) + } + return _u +} + +// AddSearchPricePer1k adds value to the "search_price_per_1k" field. +func (_u *GroupUpdate) AddSearchPricePer1k(v float64) *GroupUpdate { + _u.mutation.AddSearchPricePer1k(v) + return _u +} + +// ClearSearchPricePer1k clears the value of the "search_price_per_1k" field. +func (_u *GroupUpdate) ClearSearchPricePer1k() *GroupUpdate { + _u.mutation.ClearSearchPricePer1k() + return _u +} + +// SetAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field. +func (_u *GroupUpdate) SetAudioRealtimePricePerMin(v float64) *GroupUpdate { + _u.mutation.ResetAudioRealtimePricePerMin() + _u.mutation.SetAudioRealtimePricePerMin(v) + return _u +} + +// SetNillableAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableAudioRealtimePricePerMin(v *float64) *GroupUpdate { + if v != nil { + _u.SetAudioRealtimePricePerMin(*v) + } + return _u +} + +// AddAudioRealtimePricePerMin adds value to the "audio_realtime_price_per_min" field. +func (_u *GroupUpdate) AddAudioRealtimePricePerMin(v float64) *GroupUpdate { + _u.mutation.AddAudioRealtimePricePerMin(v) + return _u +} + +// ClearAudioRealtimePricePerMin clears the value of the "audio_realtime_price_per_min" field. +func (_u *GroupUpdate) ClearAudioRealtimePricePerMin() *GroupUpdate { + _u.mutation.ClearAudioRealtimePricePerMin() + return _u +} + +// SetAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field. +func (_u *GroupUpdate) SetAudioTtsPricePerMillionChars(v float64) *GroupUpdate { + _u.mutation.ResetAudioTtsPricePerMillionChars() + _u.mutation.SetAudioTtsPricePerMillionChars(v) + return _u +} + +// SetNillableAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableAudioTtsPricePerMillionChars(v *float64) *GroupUpdate { + if v != nil { + _u.SetAudioTtsPricePerMillionChars(*v) + } + return _u +} + +// AddAudioTtsPricePerMillionChars adds value to the "audio_tts_price_per_million_chars" field. +func (_u *GroupUpdate) AddAudioTtsPricePerMillionChars(v float64) *GroupUpdate { + _u.mutation.AddAudioTtsPricePerMillionChars(v) + return _u +} + +// ClearAudioTtsPricePerMillionChars clears the value of the "audio_tts_price_per_million_chars" field. +func (_u *GroupUpdate) ClearAudioTtsPricePerMillionChars() *GroupUpdate { + _u.mutation.ClearAudioTtsPricePerMillionChars() + return _u +} + +// SetAudioSttPricePerHour sets the "audio_stt_price_per_hour" field. +func (_u *GroupUpdate) SetAudioSttPricePerHour(v float64) *GroupUpdate { + _u.mutation.ResetAudioSttPricePerHour() + _u.mutation.SetAudioSttPricePerHour(v) + return _u +} + +// SetNillableAudioSttPricePerHour sets the "audio_stt_price_per_hour" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableAudioSttPricePerHour(v *float64) *GroupUpdate { + if v != nil { + _u.SetAudioSttPricePerHour(*v) + } + return _u +} + +// AddAudioSttPricePerHour adds value to the "audio_stt_price_per_hour" field. +func (_u *GroupUpdate) AddAudioSttPricePerHour(v float64) *GroupUpdate { + _u.mutation.AddAudioSttPricePerHour(v) + return _u +} + +// ClearAudioSttPricePerHour clears the value of the "audio_stt_price_per_hour" field. +func (_u *GroupUpdate) ClearAudioSttPricePerHour() *GroupUpdate { + _u.mutation.ClearAudioSttPricePerHour() + return _u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_u *GroupUpdate) SetClaudeCodeOnly(v bool) *GroupUpdate { _u.mutation.SetClaudeCodeOnly(v) @@ -1304,6 +1424,26 @@ func (_u *GroupUpdate) check() error { return &ValidationError{Name: "subscription_type", err: fmt.Errorf(`ent: validator failed for field "Group.subscription_type": %w`, err)} } } + if v, ok := _u.mutation.SearchPricePer1k(); ok { + if err := group.SearchPricePer1kValidator(v); err != nil { + return &ValidationError{Name: "search_price_per_1k", err: fmt.Errorf(`ent: validator failed for field "Group.search_price_per_1k": %w`, err)} + } + } + if v, ok := _u.mutation.AudioRealtimePricePerMin(); ok { + if err := group.AudioRealtimePricePerMinValidator(v); err != nil { + return &ValidationError{Name: "audio_realtime_price_per_min", err: fmt.Errorf(`ent: validator failed for field "Group.audio_realtime_price_per_min": %w`, err)} + } + } + if v, ok := _u.mutation.AudioTtsPricePerMillionChars(); ok { + if err := group.AudioTtsPricePerMillionCharsValidator(v); err != nil { + return &ValidationError{Name: "audio_tts_price_per_million_chars", err: fmt.Errorf(`ent: validator failed for field "Group.audio_tts_price_per_million_chars": %w`, err)} + } + } + if v, ok := _u.mutation.AudioSttPricePerHour(); ok { + if err := group.AudioSttPricePerHourValidator(v); err != nil { + return &ValidationError{Name: "audio_stt_price_per_hour", err: fmt.Errorf(`ent: validator failed for field "Group.audio_stt_price_per_hour": %w`, err)} + } + } if v, ok := _u.mutation.DefaultMappedModel(); ok { if err := group.DefaultMappedModelValidator(v); err != nil { return &ValidationError{Name: "default_mapped_model", err: fmt.Errorf(`ent: validator failed for field "Group.default_mapped_model": %w`, err)} @@ -1506,6 +1646,12 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if _u.mutation.VideoPrice1080pCleared() { _spec.ClearField(group.FieldVideoPrice1080p, field.TypeFloat64) } + if value, ok := _u.mutation.VideoModelPrices(); ok { + _spec.SetField(group.FieldVideoModelPrices, field.TypeJSON, value) + } + if _u.mutation.VideoModelPricesCleared() { + _spec.ClearField(group.FieldVideoModelPrices, field.TypeJSON) + } if value, ok := _u.mutation.WebSearchPricePerCall(); ok { _spec.SetField(group.FieldWebSearchPricePerCall, field.TypeFloat64, value) } @@ -1515,6 +1661,42 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if _u.mutation.WebSearchPricePerCallCleared() { _spec.ClearField(group.FieldWebSearchPricePerCall, field.TypeFloat64) } + if value, ok := _u.mutation.SearchPricePer1k(); ok { + _spec.SetField(group.FieldSearchPricePer1k, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedSearchPricePer1k(); ok { + _spec.AddField(group.FieldSearchPricePer1k, field.TypeFloat64, value) + } + if _u.mutation.SearchPricePer1kCleared() { + _spec.ClearField(group.FieldSearchPricePer1k, field.TypeFloat64) + } + if value, ok := _u.mutation.AudioRealtimePricePerMin(); ok { + _spec.SetField(group.FieldAudioRealtimePricePerMin, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedAudioRealtimePricePerMin(); ok { + _spec.AddField(group.FieldAudioRealtimePricePerMin, field.TypeFloat64, value) + } + if _u.mutation.AudioRealtimePricePerMinCleared() { + _spec.ClearField(group.FieldAudioRealtimePricePerMin, field.TypeFloat64) + } + if value, ok := _u.mutation.AudioTtsPricePerMillionChars(); ok { + _spec.SetField(group.FieldAudioTtsPricePerMillionChars, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedAudioTtsPricePerMillionChars(); ok { + _spec.AddField(group.FieldAudioTtsPricePerMillionChars, field.TypeFloat64, value) + } + if _u.mutation.AudioTtsPricePerMillionCharsCleared() { + _spec.ClearField(group.FieldAudioTtsPricePerMillionChars, field.TypeFloat64) + } + if value, ok := _u.mutation.AudioSttPricePerHour(); ok { + _spec.SetField(group.FieldAudioSttPricePerHour, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedAudioSttPricePerHour(); ok { + _spec.AddField(group.FieldAudioSttPricePerHour, field.TypeFloat64, value) + } + if _u.mutation.AudioSttPricePerHourCleared() { + _spec.ClearField(group.FieldAudioSttPricePerHour, field.TypeFloat64) + } if value, ok := _u.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) } @@ -2533,6 +2715,18 @@ func (_u *GroupUpdateOne) ClearVideoPrice1080p() *GroupUpdateOne { return _u } +// SetVideoModelPrices sets the "video_model_prices" field. +func (_u *GroupUpdateOne) SetVideoModelPrices(v map[string]map[string]float64) *GroupUpdateOne { + _u.mutation.SetVideoModelPrices(v) + return _u +} + +// ClearVideoModelPrices clears the value of the "video_model_prices" field. +func (_u *GroupUpdateOne) ClearVideoModelPrices() *GroupUpdateOne { + _u.mutation.ClearVideoModelPrices() + return _u +} + // SetWebSearchPricePerCall sets the "web_search_price_per_call" field. func (_u *GroupUpdateOne) SetWebSearchPricePerCall(v float64) *GroupUpdateOne { _u.mutation.ResetWebSearchPricePerCall() @@ -2560,6 +2754,114 @@ func (_u *GroupUpdateOne) ClearWebSearchPricePerCall() *GroupUpdateOne { return _u } +// SetSearchPricePer1k sets the "search_price_per_1k" field. +func (_u *GroupUpdateOne) SetSearchPricePer1k(v float64) *GroupUpdateOne { + _u.mutation.ResetSearchPricePer1k() + _u.mutation.SetSearchPricePer1k(v) + return _u +} + +// SetNillableSearchPricePer1k sets the "search_price_per_1k" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableSearchPricePer1k(v *float64) *GroupUpdateOne { + if v != nil { + _u.SetSearchPricePer1k(*v) + } + return _u +} + +// AddSearchPricePer1k adds value to the "search_price_per_1k" field. +func (_u *GroupUpdateOne) AddSearchPricePer1k(v float64) *GroupUpdateOne { + _u.mutation.AddSearchPricePer1k(v) + return _u +} + +// ClearSearchPricePer1k clears the value of the "search_price_per_1k" field. +func (_u *GroupUpdateOne) ClearSearchPricePer1k() *GroupUpdateOne { + _u.mutation.ClearSearchPricePer1k() + return _u +} + +// SetAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field. +func (_u *GroupUpdateOne) SetAudioRealtimePricePerMin(v float64) *GroupUpdateOne { + _u.mutation.ResetAudioRealtimePricePerMin() + _u.mutation.SetAudioRealtimePricePerMin(v) + return _u +} + +// SetNillableAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableAudioRealtimePricePerMin(v *float64) *GroupUpdateOne { + if v != nil { + _u.SetAudioRealtimePricePerMin(*v) + } + return _u +} + +// AddAudioRealtimePricePerMin adds value to the "audio_realtime_price_per_min" field. +func (_u *GroupUpdateOne) AddAudioRealtimePricePerMin(v float64) *GroupUpdateOne { + _u.mutation.AddAudioRealtimePricePerMin(v) + return _u +} + +// ClearAudioRealtimePricePerMin clears the value of the "audio_realtime_price_per_min" field. +func (_u *GroupUpdateOne) ClearAudioRealtimePricePerMin() *GroupUpdateOne { + _u.mutation.ClearAudioRealtimePricePerMin() + return _u +} + +// SetAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field. +func (_u *GroupUpdateOne) SetAudioTtsPricePerMillionChars(v float64) *GroupUpdateOne { + _u.mutation.ResetAudioTtsPricePerMillionChars() + _u.mutation.SetAudioTtsPricePerMillionChars(v) + return _u +} + +// SetNillableAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableAudioTtsPricePerMillionChars(v *float64) *GroupUpdateOne { + if v != nil { + _u.SetAudioTtsPricePerMillionChars(*v) + } + return _u +} + +// AddAudioTtsPricePerMillionChars adds value to the "audio_tts_price_per_million_chars" field. +func (_u *GroupUpdateOne) AddAudioTtsPricePerMillionChars(v float64) *GroupUpdateOne { + _u.mutation.AddAudioTtsPricePerMillionChars(v) + return _u +} + +// ClearAudioTtsPricePerMillionChars clears the value of the "audio_tts_price_per_million_chars" field. +func (_u *GroupUpdateOne) ClearAudioTtsPricePerMillionChars() *GroupUpdateOne { + _u.mutation.ClearAudioTtsPricePerMillionChars() + return _u +} + +// SetAudioSttPricePerHour sets the "audio_stt_price_per_hour" field. +func (_u *GroupUpdateOne) SetAudioSttPricePerHour(v float64) *GroupUpdateOne { + _u.mutation.ResetAudioSttPricePerHour() + _u.mutation.SetAudioSttPricePerHour(v) + return _u +} + +// SetNillableAudioSttPricePerHour sets the "audio_stt_price_per_hour" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableAudioSttPricePerHour(v *float64) *GroupUpdateOne { + if v != nil { + _u.SetAudioSttPricePerHour(*v) + } + return _u +} + +// AddAudioSttPricePerHour adds value to the "audio_stt_price_per_hour" field. +func (_u *GroupUpdateOne) AddAudioSttPricePerHour(v float64) *GroupUpdateOne { + _u.mutation.AddAudioSttPricePerHour(v) + return _u +} + +// ClearAudioSttPricePerHour clears the value of the "audio_stt_price_per_hour" field. +func (_u *GroupUpdateOne) ClearAudioSttPricePerHour() *GroupUpdateOne { + _u.mutation.ClearAudioSttPricePerHour() + return _u +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (_u *GroupUpdateOne) SetClaudeCodeOnly(v bool) *GroupUpdateOne { _u.mutation.SetClaudeCodeOnly(v) @@ -3210,6 +3512,26 @@ func (_u *GroupUpdateOne) check() error { return &ValidationError{Name: "subscription_type", err: fmt.Errorf(`ent: validator failed for field "Group.subscription_type": %w`, err)} } } + if v, ok := _u.mutation.SearchPricePer1k(); ok { + if err := group.SearchPricePer1kValidator(v); err != nil { + return &ValidationError{Name: "search_price_per_1k", err: fmt.Errorf(`ent: validator failed for field "Group.search_price_per_1k": %w`, err)} + } + } + if v, ok := _u.mutation.AudioRealtimePricePerMin(); ok { + if err := group.AudioRealtimePricePerMinValidator(v); err != nil { + return &ValidationError{Name: "audio_realtime_price_per_min", err: fmt.Errorf(`ent: validator failed for field "Group.audio_realtime_price_per_min": %w`, err)} + } + } + if v, ok := _u.mutation.AudioTtsPricePerMillionChars(); ok { + if err := group.AudioTtsPricePerMillionCharsValidator(v); err != nil { + return &ValidationError{Name: "audio_tts_price_per_million_chars", err: fmt.Errorf(`ent: validator failed for field "Group.audio_tts_price_per_million_chars": %w`, err)} + } + } + if v, ok := _u.mutation.AudioSttPricePerHour(); ok { + if err := group.AudioSttPricePerHourValidator(v); err != nil { + return &ValidationError{Name: "audio_stt_price_per_hour", err: fmt.Errorf(`ent: validator failed for field "Group.audio_stt_price_per_hour": %w`, err)} + } + } if v, ok := _u.mutation.DefaultMappedModel(); ok { if err := group.DefaultMappedModelValidator(v); err != nil { return &ValidationError{Name: "default_mapped_model", err: fmt.Errorf(`ent: validator failed for field "Group.default_mapped_model": %w`, err)} @@ -3429,6 +3751,12 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if _u.mutation.VideoPrice1080pCleared() { _spec.ClearField(group.FieldVideoPrice1080p, field.TypeFloat64) } + if value, ok := _u.mutation.VideoModelPrices(); ok { + _spec.SetField(group.FieldVideoModelPrices, field.TypeJSON, value) + } + if _u.mutation.VideoModelPricesCleared() { + _spec.ClearField(group.FieldVideoModelPrices, field.TypeJSON) + } if value, ok := _u.mutation.WebSearchPricePerCall(); ok { _spec.SetField(group.FieldWebSearchPricePerCall, field.TypeFloat64, value) } @@ -3438,6 +3766,42 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if _u.mutation.WebSearchPricePerCallCleared() { _spec.ClearField(group.FieldWebSearchPricePerCall, field.TypeFloat64) } + if value, ok := _u.mutation.SearchPricePer1k(); ok { + _spec.SetField(group.FieldSearchPricePer1k, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedSearchPricePer1k(); ok { + _spec.AddField(group.FieldSearchPricePer1k, field.TypeFloat64, value) + } + if _u.mutation.SearchPricePer1kCleared() { + _spec.ClearField(group.FieldSearchPricePer1k, field.TypeFloat64) + } + if value, ok := _u.mutation.AudioRealtimePricePerMin(); ok { + _spec.SetField(group.FieldAudioRealtimePricePerMin, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedAudioRealtimePricePerMin(); ok { + _spec.AddField(group.FieldAudioRealtimePricePerMin, field.TypeFloat64, value) + } + if _u.mutation.AudioRealtimePricePerMinCleared() { + _spec.ClearField(group.FieldAudioRealtimePricePerMin, field.TypeFloat64) + } + if value, ok := _u.mutation.AudioTtsPricePerMillionChars(); ok { + _spec.SetField(group.FieldAudioTtsPricePerMillionChars, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedAudioTtsPricePerMillionChars(); ok { + _spec.AddField(group.FieldAudioTtsPricePerMillionChars, field.TypeFloat64, value) + } + if _u.mutation.AudioTtsPricePerMillionCharsCleared() { + _spec.ClearField(group.FieldAudioTtsPricePerMillionChars, field.TypeFloat64) + } + if value, ok := _u.mutation.AudioSttPricePerHour(); ok { + _spec.SetField(group.FieldAudioSttPricePerHour, field.TypeFloat64, value) + } + if value, ok := _u.mutation.AddedAudioSttPricePerHour(); ok { + _spec.AddField(group.FieldAudioSttPricePerHour, field.TypeFloat64, value) + } + if _u.mutation.AudioSttPricePerHourCleared() { + _spec.ClearField(group.FieldAudioSttPricePerHour, field.TypeFloat64) + } if value, ok := _u.mutation.ClaudeCodeOnly(); ok { _spec.SetField(group.FieldClaudeCodeOnly, field.TypeBool, value) } diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index 0475bca3bf..491f038bc7 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -928,7 +928,12 @@ var ( {Name: "video_price_480p", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "video_price_720p", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "video_price_1080p", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, + {Name: "video_model_prices", Type: field.TypeJSON, Nullable: true, SchemaType: map[string]string{"postgres": "jsonb"}}, {Name: "web_search_price_per_call", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, + {Name: "search_price_per_1k", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, + {Name: "audio_realtime_price_per_min", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, + {Name: "audio_tts_price_per_million_chars", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, + {Name: "audio_stt_price_per_hour", Type: field.TypeFloat64, Nullable: true, SchemaType: map[string]string{"postgres": "decimal(20,8)"}}, {Name: "claude_code_only", Type: field.TypeBool, Default: false}, {Name: "fallback_group_id", Type: field.TypeInt64, Nullable: true}, {Name: "fallback_group_id_on_invalid_request", Type: field.TypeInt64, Nullable: true}, @@ -985,7 +990,7 @@ var ( { Name: "group_sort_order", Unique: false, - Columns: []*schema.Column{GroupsColumns[42]}, + Columns: []*schema.Column{GroupsColumns[47]}, }, { Name: "idx_groups_duplicate_operation_id_active", @@ -1621,6 +1626,8 @@ var ( {Name: "model", Type: field.TypeString, Size: 100}, {Name: "requested_model", Type: field.TypeString, Nullable: true, Size: 100}, {Name: "upstream_model", Type: field.TypeString, Nullable: true, Size: 100}, + {Name: "upstream_response_model", Type: field.TypeString, Nullable: true, Size: 200}, + {Name: "upstream_model_mismatch", Type: field.TypeBool, Nullable: true}, {Name: "channel_id", Type: field.TypeInt64, Nullable: true}, {Name: "model_mapping_chain", Type: field.TypeString, Nullable: true, Size: 500}, {Name: "billing_tier", Type: field.TypeString, Nullable: true, Size: 50}, @@ -1671,31 +1678,31 @@ var ( ForeignKeys: []*schema.ForeignKey{ { Symbol: "usage_logs_api_keys_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[41]}, + Columns: []*schema.Column{UsageLogsColumns[43]}, RefColumns: []*schema.Column{APIKeysColumns[0]}, OnDelete: schema.NoAction, }, { Symbol: "usage_logs_accounts_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[42]}, + Columns: []*schema.Column{UsageLogsColumns[44]}, RefColumns: []*schema.Column{AccountsColumns[0]}, OnDelete: schema.NoAction, }, { Symbol: "usage_logs_groups_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[43]}, + Columns: []*schema.Column{UsageLogsColumns[45]}, RefColumns: []*schema.Column{GroupsColumns[0]}, OnDelete: schema.SetNull, }, { Symbol: "usage_logs_users_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[44]}, + Columns: []*schema.Column{UsageLogsColumns[46]}, RefColumns: []*schema.Column{UsersColumns[0]}, OnDelete: schema.NoAction, }, { Symbol: "usage_logs_user_subscriptions_usage_logs", - Columns: []*schema.Column{UsageLogsColumns[45]}, + Columns: []*schema.Column{UsageLogsColumns[47]}, RefColumns: []*schema.Column{UserSubscriptionsColumns[0]}, OnDelete: schema.SetNull, }, @@ -1704,32 +1711,32 @@ var ( { Name: "usagelog_user_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[44]}, + Columns: []*schema.Column{UsageLogsColumns[46]}, }, { Name: "usagelog_api_key_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[41]}, + Columns: []*schema.Column{UsageLogsColumns[43]}, }, { Name: "usagelog_account_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[42]}, + Columns: []*schema.Column{UsageLogsColumns[44]}, }, { Name: "usagelog_group_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[43]}, + Columns: []*schema.Column{UsageLogsColumns[45]}, }, { Name: "usagelog_subscription_id", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[45]}, + Columns: []*schema.Column{UsageLogsColumns[47]}, }, { Name: "usagelog_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[40]}, + Columns: []*schema.Column{UsageLogsColumns[42]}, }, { Name: "usagelog_model", @@ -1749,17 +1756,17 @@ var ( { Name: "usagelog_user_id_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[44], UsageLogsColumns[40]}, + Columns: []*schema.Column{UsageLogsColumns[46], UsageLogsColumns[42]}, }, { Name: "usagelog_api_key_id_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[41], UsageLogsColumns[40]}, + Columns: []*schema.Column{UsageLogsColumns[43], UsageLogsColumns[42]}, }, { Name: "usagelog_group_id_created_at", Unique: false, - Columns: []*schema.Column{UsageLogsColumns[43], UsageLogsColumns[40]}, + Columns: []*schema.Column{UsageLogsColumns[45], UsageLogsColumns[42]}, }, }, } diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index 61ff4f9e6d..99e8a2cd95 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -21896,8 +21896,17 @@ type GroupMutation struct { addvideo_price_720p *float64 video_price_1080p *float64 addvideo_price_1080p *float64 + video_model_prices *map[string]map[string]float64 web_search_price_per_call *float64 addweb_search_price_per_call *float64 + search_price_per_1k *float64 + addsearch_price_per_1k *float64 + audio_realtime_price_per_min *float64 + addaudio_realtime_price_per_min *float64 + audio_tts_price_per_million_chars *float64 + addaudio_tts_price_per_million_chars *float64 + audio_stt_price_per_hour *float64 + addaudio_stt_price_per_hour *float64 claude_code_only *bool fallback_group_id *int64 addfallback_group_id *int64 @@ -23722,6 +23731,55 @@ func (m *GroupMutation) ResetVideoPrice1080p() { delete(m.clearedFields, group.FieldVideoPrice1080p) } +// SetVideoModelPrices sets the "video_model_prices" field. +func (m *GroupMutation) SetVideoModelPrices(value map[string]map[string]float64) { + m.video_model_prices = &value +} + +// VideoModelPrices returns the value of the "video_model_prices" field in the mutation. +func (m *GroupMutation) VideoModelPrices() (r map[string]map[string]float64, exists bool) { + v := m.video_model_prices + if v == nil { + return + } + return *v, true +} + +// OldVideoModelPrices returns the old "video_model_prices" 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) OldVideoModelPrices(ctx context.Context) (v map[string]map[string]float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldVideoModelPrices is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldVideoModelPrices requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldVideoModelPrices: %w", err) + } + return oldValue.VideoModelPrices, nil +} + +// ClearVideoModelPrices clears the value of the "video_model_prices" field. +func (m *GroupMutation) ClearVideoModelPrices() { + m.video_model_prices = nil + m.clearedFields[group.FieldVideoModelPrices] = struct{}{} +} + +// VideoModelPricesCleared returns if the "video_model_prices" field was cleared in this mutation. +func (m *GroupMutation) VideoModelPricesCleared() bool { + _, ok := m.clearedFields[group.FieldVideoModelPrices] + return ok +} + +// ResetVideoModelPrices resets all changes to the "video_model_prices" field. +func (m *GroupMutation) ResetVideoModelPrices() { + m.video_model_prices = nil + delete(m.clearedFields, group.FieldVideoModelPrices) +} + // SetWebSearchPricePerCall sets the "web_search_price_per_call" field. func (m *GroupMutation) SetWebSearchPricePerCall(f float64) { m.web_search_price_per_call = &f @@ -23792,6 +23850,286 @@ func (m *GroupMutation) ResetWebSearchPricePerCall() { delete(m.clearedFields, group.FieldWebSearchPricePerCall) } +// SetSearchPricePer1k sets the "search_price_per_1k" field. +func (m *GroupMutation) SetSearchPricePer1k(f float64) { + m.search_price_per_1k = &f + m.addsearch_price_per_1k = nil +} + +// SearchPricePer1k returns the value of the "search_price_per_1k" field in the mutation. +func (m *GroupMutation) SearchPricePer1k() (r float64, exists bool) { + v := m.search_price_per_1k + if v == nil { + return + } + return *v, true +} + +// OldSearchPricePer1k returns the old "search_price_per_1k" 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) OldSearchPricePer1k(ctx context.Context) (v *float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldSearchPricePer1k is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldSearchPricePer1k requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldSearchPricePer1k: %w", err) + } + return oldValue.SearchPricePer1k, nil +} + +// AddSearchPricePer1k adds f to the "search_price_per_1k" field. +func (m *GroupMutation) AddSearchPricePer1k(f float64) { + if m.addsearch_price_per_1k != nil { + *m.addsearch_price_per_1k += f + } else { + m.addsearch_price_per_1k = &f + } +} + +// AddedSearchPricePer1k returns the value that was added to the "search_price_per_1k" field in this mutation. +func (m *GroupMutation) AddedSearchPricePer1k() (r float64, exists bool) { + v := m.addsearch_price_per_1k + if v == nil { + return + } + return *v, true +} + +// ClearSearchPricePer1k clears the value of the "search_price_per_1k" field. +func (m *GroupMutation) ClearSearchPricePer1k() { + m.search_price_per_1k = nil + m.addsearch_price_per_1k = nil + m.clearedFields[group.FieldSearchPricePer1k] = struct{}{} +} + +// SearchPricePer1kCleared returns if the "search_price_per_1k" field was cleared in this mutation. +func (m *GroupMutation) SearchPricePer1kCleared() bool { + _, ok := m.clearedFields[group.FieldSearchPricePer1k] + return ok +} + +// ResetSearchPricePer1k resets all changes to the "search_price_per_1k" field. +func (m *GroupMutation) ResetSearchPricePer1k() { + m.search_price_per_1k = nil + m.addsearch_price_per_1k = nil + delete(m.clearedFields, group.FieldSearchPricePer1k) +} + +// SetAudioRealtimePricePerMin sets the "audio_realtime_price_per_min" field. +func (m *GroupMutation) SetAudioRealtimePricePerMin(f float64) { + m.audio_realtime_price_per_min = &f + m.addaudio_realtime_price_per_min = nil +} + +// AudioRealtimePricePerMin returns the value of the "audio_realtime_price_per_min" field in the mutation. +func (m *GroupMutation) AudioRealtimePricePerMin() (r float64, exists bool) { + v := m.audio_realtime_price_per_min + if v == nil { + return + } + return *v, true +} + +// OldAudioRealtimePricePerMin returns the old "audio_realtime_price_per_min" 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) OldAudioRealtimePricePerMin(ctx context.Context) (v *float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldAudioRealtimePricePerMin is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldAudioRealtimePricePerMin requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldAudioRealtimePricePerMin: %w", err) + } + return oldValue.AudioRealtimePricePerMin, nil +} + +// AddAudioRealtimePricePerMin adds f to the "audio_realtime_price_per_min" field. +func (m *GroupMutation) AddAudioRealtimePricePerMin(f float64) { + if m.addaudio_realtime_price_per_min != nil { + *m.addaudio_realtime_price_per_min += f + } else { + m.addaudio_realtime_price_per_min = &f + } +} + +// AddedAudioRealtimePricePerMin returns the value that was added to the "audio_realtime_price_per_min" field in this mutation. +func (m *GroupMutation) AddedAudioRealtimePricePerMin() (r float64, exists bool) { + v := m.addaudio_realtime_price_per_min + if v == nil { + return + } + return *v, true +} + +// ClearAudioRealtimePricePerMin clears the value of the "audio_realtime_price_per_min" field. +func (m *GroupMutation) ClearAudioRealtimePricePerMin() { + m.audio_realtime_price_per_min = nil + m.addaudio_realtime_price_per_min = nil + m.clearedFields[group.FieldAudioRealtimePricePerMin] = struct{}{} +} + +// AudioRealtimePricePerMinCleared returns if the "audio_realtime_price_per_min" field was cleared in this mutation. +func (m *GroupMutation) AudioRealtimePricePerMinCleared() bool { + _, ok := m.clearedFields[group.FieldAudioRealtimePricePerMin] + return ok +} + +// ResetAudioRealtimePricePerMin resets all changes to the "audio_realtime_price_per_min" field. +func (m *GroupMutation) ResetAudioRealtimePricePerMin() { + m.audio_realtime_price_per_min = nil + m.addaudio_realtime_price_per_min = nil + delete(m.clearedFields, group.FieldAudioRealtimePricePerMin) +} + +// SetAudioTtsPricePerMillionChars sets the "audio_tts_price_per_million_chars" field. +func (m *GroupMutation) SetAudioTtsPricePerMillionChars(f float64) { + m.audio_tts_price_per_million_chars = &f + m.addaudio_tts_price_per_million_chars = nil +} + +// AudioTtsPricePerMillionChars returns the value of the "audio_tts_price_per_million_chars" field in the mutation. +func (m *GroupMutation) AudioTtsPricePerMillionChars() (r float64, exists bool) { + v := m.audio_tts_price_per_million_chars + if v == nil { + return + } + return *v, true +} + +// OldAudioTtsPricePerMillionChars returns the old "audio_tts_price_per_million_chars" 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) OldAudioTtsPricePerMillionChars(ctx context.Context) (v *float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldAudioTtsPricePerMillionChars is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldAudioTtsPricePerMillionChars requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldAudioTtsPricePerMillionChars: %w", err) + } + return oldValue.AudioTtsPricePerMillionChars, nil +} + +// AddAudioTtsPricePerMillionChars adds f to the "audio_tts_price_per_million_chars" field. +func (m *GroupMutation) AddAudioTtsPricePerMillionChars(f float64) { + if m.addaudio_tts_price_per_million_chars != nil { + *m.addaudio_tts_price_per_million_chars += f + } else { + m.addaudio_tts_price_per_million_chars = &f + } +} + +// AddedAudioTtsPricePerMillionChars returns the value that was added to the "audio_tts_price_per_million_chars" field in this mutation. +func (m *GroupMutation) AddedAudioTtsPricePerMillionChars() (r float64, exists bool) { + v := m.addaudio_tts_price_per_million_chars + if v == nil { + return + } + return *v, true +} + +// ClearAudioTtsPricePerMillionChars clears the value of the "audio_tts_price_per_million_chars" field. +func (m *GroupMutation) ClearAudioTtsPricePerMillionChars() { + m.audio_tts_price_per_million_chars = nil + m.addaudio_tts_price_per_million_chars = nil + m.clearedFields[group.FieldAudioTtsPricePerMillionChars] = struct{}{} +} + +// AudioTtsPricePerMillionCharsCleared returns if the "audio_tts_price_per_million_chars" field was cleared in this mutation. +func (m *GroupMutation) AudioTtsPricePerMillionCharsCleared() bool { + _, ok := m.clearedFields[group.FieldAudioTtsPricePerMillionChars] + return ok +} + +// ResetAudioTtsPricePerMillionChars resets all changes to the "audio_tts_price_per_million_chars" field. +func (m *GroupMutation) ResetAudioTtsPricePerMillionChars() { + m.audio_tts_price_per_million_chars = nil + m.addaudio_tts_price_per_million_chars = nil + delete(m.clearedFields, group.FieldAudioTtsPricePerMillionChars) +} + +// SetAudioSttPricePerHour sets the "audio_stt_price_per_hour" field. +func (m *GroupMutation) SetAudioSttPricePerHour(f float64) { + m.audio_stt_price_per_hour = &f + m.addaudio_stt_price_per_hour = nil +} + +// AudioSttPricePerHour returns the value of the "audio_stt_price_per_hour" field in the mutation. +func (m *GroupMutation) AudioSttPricePerHour() (r float64, exists bool) { + v := m.audio_stt_price_per_hour + if v == nil { + return + } + return *v, true +} + +// OldAudioSttPricePerHour returns the old "audio_stt_price_per_hour" 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) OldAudioSttPricePerHour(ctx context.Context) (v *float64, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldAudioSttPricePerHour is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldAudioSttPricePerHour requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldAudioSttPricePerHour: %w", err) + } + return oldValue.AudioSttPricePerHour, nil +} + +// AddAudioSttPricePerHour adds f to the "audio_stt_price_per_hour" field. +func (m *GroupMutation) AddAudioSttPricePerHour(f float64) { + if m.addaudio_stt_price_per_hour != nil { + *m.addaudio_stt_price_per_hour += f + } else { + m.addaudio_stt_price_per_hour = &f + } +} + +// AddedAudioSttPricePerHour returns the value that was added to the "audio_stt_price_per_hour" field in this mutation. +func (m *GroupMutation) AddedAudioSttPricePerHour() (r float64, exists bool) { + v := m.addaudio_stt_price_per_hour + if v == nil { + return + } + return *v, true +} + +// ClearAudioSttPricePerHour clears the value of the "audio_stt_price_per_hour" field. +func (m *GroupMutation) ClearAudioSttPricePerHour() { + m.audio_stt_price_per_hour = nil + m.addaudio_stt_price_per_hour = nil + m.clearedFields[group.FieldAudioSttPricePerHour] = struct{}{} +} + +// AudioSttPricePerHourCleared returns if the "audio_stt_price_per_hour" field was cleared in this mutation. +func (m *GroupMutation) AudioSttPricePerHourCleared() bool { + _, ok := m.clearedFields[group.FieldAudioSttPricePerHour] + return ok +} + +// ResetAudioSttPricePerHour resets all changes to the "audio_stt_price_per_hour" field. +func (m *GroupMutation) ResetAudioSttPricePerHour() { + m.audio_stt_price_per_hour = nil + m.addaudio_stt_price_per_hour = nil + delete(m.clearedFields, group.FieldAudioSttPricePerHour) +} + // SetClaudeCodeOnly sets the "claude_code_only" field. func (m *GroupMutation) SetClaudeCodeOnly(b bool) { m.claude_code_only = &b @@ -25097,7 +25435,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, 55) + fields := make([]string, 0, 60) if m.created_at != nil { fields = append(fields, group.FieldCreatedAt) } @@ -25197,9 +25535,24 @@ func (m *GroupMutation) Fields() []string { if m.video_price_1080p != nil { fields = append(fields, group.FieldVideoPrice1080p) } + if m.video_model_prices != nil { + fields = append(fields, group.FieldVideoModelPrices) + } if m.web_search_price_per_call != nil { fields = append(fields, group.FieldWebSearchPricePerCall) } + if m.search_price_per_1k != nil { + fields = append(fields, group.FieldSearchPricePer1k) + } + if m.audio_realtime_price_per_min != nil { + fields = append(fields, group.FieldAudioRealtimePricePerMin) + } + if m.audio_tts_price_per_million_chars != nil { + fields = append(fields, group.FieldAudioTtsPricePerMillionChars) + } + if m.audio_stt_price_per_hour != nil { + fields = append(fields, group.FieldAudioSttPricePerHour) + } if m.claude_code_only != nil { fields = append(fields, group.FieldClaudeCodeOnly) } @@ -25337,8 +25690,18 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.VideoPrice720p() case group.FieldVideoPrice1080p: return m.VideoPrice1080p() + case group.FieldVideoModelPrices: + return m.VideoModelPrices() case group.FieldWebSearchPricePerCall: return m.WebSearchPricePerCall() + case group.FieldSearchPricePer1k: + return m.SearchPricePer1k() + case group.FieldAudioRealtimePricePerMin: + return m.AudioRealtimePricePerMin() + case group.FieldAudioTtsPricePerMillionChars: + return m.AudioTtsPricePerMillionChars() + case group.FieldAudioSttPricePerHour: + return m.AudioSttPricePerHour() case group.FieldClaudeCodeOnly: return m.ClaudeCodeOnly() case group.FieldFallbackGroupID: @@ -25456,8 +25819,18 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldVideoPrice720p(ctx) case group.FieldVideoPrice1080p: return m.OldVideoPrice1080p(ctx) + case group.FieldVideoModelPrices: + return m.OldVideoModelPrices(ctx) case group.FieldWebSearchPricePerCall: return m.OldWebSearchPricePerCall(ctx) + case group.FieldSearchPricePer1k: + return m.OldSearchPricePer1k(ctx) + case group.FieldAudioRealtimePricePerMin: + return m.OldAudioRealtimePricePerMin(ctx) + case group.FieldAudioTtsPricePerMillionChars: + return m.OldAudioTtsPricePerMillionChars(ctx) + case group.FieldAudioSttPricePerHour: + return m.OldAudioSttPricePerHour(ctx) case group.FieldClaudeCodeOnly: return m.OldClaudeCodeOnly(ctx) case group.FieldFallbackGroupID: @@ -25740,6 +26113,13 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetVideoPrice1080p(v) return nil + case group.FieldVideoModelPrices: + v, ok := value.(map[string]map[string]float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetVideoModelPrices(v) + return nil case group.FieldWebSearchPricePerCall: v, ok := value.(float64) if !ok { @@ -25747,6 +26127,34 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetWebSearchPricePerCall(v) return nil + case group.FieldSearchPricePer1k: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetSearchPricePer1k(v) + return nil + case group.FieldAudioRealtimePricePerMin: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetAudioRealtimePricePerMin(v) + return nil + case group.FieldAudioTtsPricePerMillionChars: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetAudioTtsPricePerMillionChars(v) + return nil + case group.FieldAudioSttPricePerHour: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetAudioSttPricePerHour(v) + return nil case group.FieldClaudeCodeOnly: v, ok := value.(bool) if !ok { @@ -25953,6 +26361,18 @@ func (m *GroupMutation) AddedFields() []string { if m.addweb_search_price_per_call != nil { fields = append(fields, group.FieldWebSearchPricePerCall) } + if m.addsearch_price_per_1k != nil { + fields = append(fields, group.FieldSearchPricePer1k) + } + if m.addaudio_realtime_price_per_min != nil { + fields = append(fields, group.FieldAudioRealtimePricePerMin) + } + if m.addaudio_tts_price_per_million_chars != nil { + fields = append(fields, group.FieldAudioTtsPricePerMillionChars) + } + if m.addaudio_stt_price_per_hour != nil { + fields = append(fields, group.FieldAudioSttPricePerHour) + } if m.addfallback_group_id != nil { fields = append(fields, group.FieldFallbackGroupID) } @@ -26013,6 +26433,14 @@ func (m *GroupMutation) AddedField(name string) (ent.Value, bool) { return m.AddedVideoPrice1080p() case group.FieldWebSearchPricePerCall: return m.AddedWebSearchPricePerCall() + case group.FieldSearchPricePer1k: + return m.AddedSearchPricePer1k() + case group.FieldAudioRealtimePricePerMin: + return m.AddedAudioRealtimePricePerMin() + case group.FieldAudioTtsPricePerMillionChars: + return m.AddedAudioTtsPricePerMillionChars() + case group.FieldAudioSttPricePerHour: + return m.AddedAudioSttPricePerHour() case group.FieldFallbackGroupID: return m.AddedFallbackGroupID() case group.FieldFallbackGroupIDOnInvalidRequest: @@ -26153,6 +26581,34 @@ func (m *GroupMutation) AddField(name string, value ent.Value) error { } m.AddWebSearchPricePerCall(v) return nil + case group.FieldSearchPricePer1k: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddSearchPricePer1k(v) + return nil + case group.FieldAudioRealtimePricePerMin: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddAudioRealtimePricePerMin(v) + return nil + case group.FieldAudioTtsPricePerMillionChars: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddAudioTtsPricePerMillionChars(v) + return nil + case group.FieldAudioSttPricePerHour: + v, ok := value.(float64) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.AddAudioSttPricePerHour(v) + return nil case group.FieldFallbackGroupID: v, ok := value.(int64) if !ok { @@ -26239,9 +26695,24 @@ func (m *GroupMutation) ClearedFields() []string { if m.FieldCleared(group.FieldVideoPrice1080p) { fields = append(fields, group.FieldVideoPrice1080p) } + if m.FieldCleared(group.FieldVideoModelPrices) { + fields = append(fields, group.FieldVideoModelPrices) + } if m.FieldCleared(group.FieldWebSearchPricePerCall) { fields = append(fields, group.FieldWebSearchPricePerCall) } + if m.FieldCleared(group.FieldSearchPricePer1k) { + fields = append(fields, group.FieldSearchPricePer1k) + } + if m.FieldCleared(group.FieldAudioRealtimePricePerMin) { + fields = append(fields, group.FieldAudioRealtimePricePerMin) + } + if m.FieldCleared(group.FieldAudioTtsPricePerMillionChars) { + fields = append(fields, group.FieldAudioTtsPricePerMillionChars) + } + if m.FieldCleared(group.FieldAudioSttPricePerHour) { + fields = append(fields, group.FieldAudioSttPricePerHour) + } if m.FieldCleared(group.FieldFallbackGroupID) { fields = append(fields, group.FieldFallbackGroupID) } @@ -26301,9 +26772,24 @@ func (m *GroupMutation) ClearField(name string) error { case group.FieldVideoPrice1080p: m.ClearVideoPrice1080p() return nil + case group.FieldVideoModelPrices: + m.ClearVideoModelPrices() + return nil case group.FieldWebSearchPricePerCall: m.ClearWebSearchPricePerCall() return nil + case group.FieldSearchPricePer1k: + m.ClearSearchPricePer1k() + return nil + case group.FieldAudioRealtimePricePerMin: + m.ClearAudioRealtimePricePerMin() + return nil + case group.FieldAudioTtsPricePerMillionChars: + m.ClearAudioTtsPricePerMillionChars() + return nil + case group.FieldAudioSttPricePerHour: + m.ClearAudioSttPricePerHour() + return nil case group.FieldFallbackGroupID: m.ClearFallbackGroupID() return nil @@ -26420,9 +26906,24 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldVideoPrice1080p: m.ResetVideoPrice1080p() return nil + case group.FieldVideoModelPrices: + m.ResetVideoModelPrices() + return nil case group.FieldWebSearchPricePerCall: m.ResetWebSearchPricePerCall() return nil + case group.FieldSearchPricePer1k: + m.ResetSearchPricePer1k() + return nil + case group.FieldAudioRealtimePricePerMin: + m.ResetAudioRealtimePricePerMin() + return nil + case group.FieldAudioTtsPricePerMillionChars: + m.ResetAudioTtsPricePerMillionChars() + return nil + case group.FieldAudioSttPricePerHour: + m.ResetAudioSttPricePerHour() + return nil case group.FieldClaudeCodeOnly: m.ResetClaudeCodeOnly() return nil @@ -43356,6 +43857,8 @@ type UsageLogMutation struct { model *string requested_model *string upstream_model *string + upstream_response_model *string + upstream_model_mismatch *bool channel_id *int64 addchannel_id *int64 model_mapping_chain *string @@ -43805,6 +44308,104 @@ func (m *UsageLogMutation) ResetUpstreamModel() { delete(m.clearedFields, usagelog.FieldUpstreamModel) } +// SetUpstreamResponseModel sets the "upstream_response_model" field. +func (m *UsageLogMutation) SetUpstreamResponseModel(s string) { + m.upstream_response_model = &s +} + +// UpstreamResponseModel returns the value of the "upstream_response_model" field in the mutation. +func (m *UsageLogMutation) UpstreamResponseModel() (r string, exists bool) { + v := m.upstream_response_model + if v == nil { + return + } + return *v, true +} + +// OldUpstreamResponseModel returns the old "upstream_response_model" field's value of the UsageLog entity. +// If the UsageLog 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 *UsageLogMutation) OldUpstreamResponseModel(ctx context.Context) (v *string, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldUpstreamResponseModel is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldUpstreamResponseModel requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldUpstreamResponseModel: %w", err) + } + return oldValue.UpstreamResponseModel, nil +} + +// ClearUpstreamResponseModel clears the value of the "upstream_response_model" field. +func (m *UsageLogMutation) ClearUpstreamResponseModel() { + m.upstream_response_model = nil + m.clearedFields[usagelog.FieldUpstreamResponseModel] = struct{}{} +} + +// UpstreamResponseModelCleared returns if the "upstream_response_model" field was cleared in this mutation. +func (m *UsageLogMutation) UpstreamResponseModelCleared() bool { + _, ok := m.clearedFields[usagelog.FieldUpstreamResponseModel] + return ok +} + +// ResetUpstreamResponseModel resets all changes to the "upstream_response_model" field. +func (m *UsageLogMutation) ResetUpstreamResponseModel() { + m.upstream_response_model = nil + delete(m.clearedFields, usagelog.FieldUpstreamResponseModel) +} + +// SetUpstreamModelMismatch sets the "upstream_model_mismatch" field. +func (m *UsageLogMutation) SetUpstreamModelMismatch(b bool) { + m.upstream_model_mismatch = &b +} + +// UpstreamModelMismatch returns the value of the "upstream_model_mismatch" field in the mutation. +func (m *UsageLogMutation) UpstreamModelMismatch() (r bool, exists bool) { + v := m.upstream_model_mismatch + if v == nil { + return + } + return *v, true +} + +// OldUpstreamModelMismatch returns the old "upstream_model_mismatch" field's value of the UsageLog entity. +// If the UsageLog 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 *UsageLogMutation) OldUpstreamModelMismatch(ctx context.Context) (v *bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldUpstreamModelMismatch is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldUpstreamModelMismatch requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldUpstreamModelMismatch: %w", err) + } + return oldValue.UpstreamModelMismatch, nil +} + +// ClearUpstreamModelMismatch clears the value of the "upstream_model_mismatch" field. +func (m *UsageLogMutation) ClearUpstreamModelMismatch() { + m.upstream_model_mismatch = nil + m.clearedFields[usagelog.FieldUpstreamModelMismatch] = struct{}{} +} + +// UpstreamModelMismatchCleared returns if the "upstream_model_mismatch" field was cleared in this mutation. +func (m *UsageLogMutation) UpstreamModelMismatchCleared() bool { + _, ok := m.clearedFields[usagelog.FieldUpstreamModelMismatch] + return ok +} + +// ResetUpstreamModelMismatch resets all changes to the "upstream_model_mismatch" field. +func (m *UsageLogMutation) ResetUpstreamModelMismatch() { + m.upstream_model_mismatch = nil + delete(m.clearedFields, usagelog.FieldUpstreamModelMismatch) +} + // SetChannelID sets the "channel_id" field. func (m *UsageLogMutation) SetChannelID(i int64) { m.channel_id = &i @@ -46001,7 +46602,7 @@ func (m *UsageLogMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *UsageLogMutation) Fields() []string { - fields := make([]string, 0, 45) + fields := make([]string, 0, 47) if m.user != nil { fields = append(fields, usagelog.FieldUserID) } @@ -46023,6 +46624,12 @@ func (m *UsageLogMutation) Fields() []string { if m.upstream_model != nil { fields = append(fields, usagelog.FieldUpstreamModel) } + if m.upstream_response_model != nil { + fields = append(fields, usagelog.FieldUpstreamResponseModel) + } + if m.upstream_model_mismatch != nil { + fields = append(fields, usagelog.FieldUpstreamModelMismatch) + } if m.channel_id != nil { fields = append(fields, usagelog.FieldChannelID) } @@ -46159,6 +46766,10 @@ func (m *UsageLogMutation) Field(name string) (ent.Value, bool) { return m.RequestedModel() case usagelog.FieldUpstreamModel: return m.UpstreamModel() + case usagelog.FieldUpstreamResponseModel: + return m.UpstreamResponseModel() + case usagelog.FieldUpstreamModelMismatch: + return m.UpstreamModelMismatch() case usagelog.FieldChannelID: return m.ChannelID() case usagelog.FieldModelMappingChain: @@ -46258,6 +46869,10 @@ func (m *UsageLogMutation) OldField(ctx context.Context, name string) (ent.Value return m.OldRequestedModel(ctx) case usagelog.FieldUpstreamModel: return m.OldUpstreamModel(ctx) + case usagelog.FieldUpstreamResponseModel: + return m.OldUpstreamResponseModel(ctx) + case usagelog.FieldUpstreamModelMismatch: + return m.OldUpstreamModelMismatch(ctx) case usagelog.FieldChannelID: return m.OldChannelID(ctx) case usagelog.FieldModelMappingChain: @@ -46392,6 +47007,20 @@ func (m *UsageLogMutation) SetField(name string, value ent.Value) error { } m.SetUpstreamModel(v) return nil + case usagelog.FieldUpstreamResponseModel: + v, ok := value.(string) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetUpstreamResponseModel(v) + return nil + case usagelog.FieldUpstreamModelMismatch: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetUpstreamModelMismatch(v) + return nil case usagelog.FieldChannelID: v, ok := value.(int64) if !ok { @@ -46949,6 +47578,12 @@ func (m *UsageLogMutation) ClearedFields() []string { if m.FieldCleared(usagelog.FieldUpstreamModel) { fields = append(fields, usagelog.FieldUpstreamModel) } + if m.FieldCleared(usagelog.FieldUpstreamResponseModel) { + fields = append(fields, usagelog.FieldUpstreamResponseModel) + } + if m.FieldCleared(usagelog.FieldUpstreamModelMismatch) { + fields = append(fields, usagelog.FieldUpstreamModelMismatch) + } if m.FieldCleared(usagelog.FieldChannelID) { fields = append(fields, usagelog.FieldChannelID) } @@ -47023,6 +47658,12 @@ func (m *UsageLogMutation) ClearField(name string) error { case usagelog.FieldUpstreamModel: m.ClearUpstreamModel() return nil + case usagelog.FieldUpstreamResponseModel: + m.ClearUpstreamResponseModel() + return nil + case usagelog.FieldUpstreamModelMismatch: + m.ClearUpstreamModelMismatch() + return nil case usagelog.FieldChannelID: m.ClearChannelID() return nil @@ -47106,6 +47747,12 @@ func (m *UsageLogMutation) ResetField(name string) error { case usagelog.FieldUpstreamModel: m.ResetUpstreamModel() return nil + case usagelog.FieldUpstreamResponseModel: + m.ResetUpstreamResponseModel() + return nil + case usagelog.FieldUpstreamModelMismatch: + m.ResetUpstreamModelMismatch() + return nil case usagelog.FieldChannelID: m.ResetChannelID() return nil diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index e85660dcc1..8042a835a9 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -1117,80 +1117,96 @@ func init() { groupDescVideoRateMultiplier := groupFields[26].Descriptor() // group.DefaultVideoRateMultiplier holds the default value on creation for the video_rate_multiplier field. group.DefaultVideoRateMultiplier = groupDescVideoRateMultiplier.Default.(float64) + // groupDescSearchPricePer1k is the schema descriptor for search_price_per_1k field. + groupDescSearchPricePer1k := groupFields[32].Descriptor() + // group.SearchPricePer1kValidator is a validator for the "search_price_per_1k" field. It is called by the builders before save. + group.SearchPricePer1kValidator = groupDescSearchPricePer1k.Validators[0].(func(float64) error) + // groupDescAudioRealtimePricePerMin is the schema descriptor for audio_realtime_price_per_min field. + groupDescAudioRealtimePricePerMin := groupFields[33].Descriptor() + // group.AudioRealtimePricePerMinValidator is a validator for the "audio_realtime_price_per_min" field. It is called by the builders before save. + group.AudioRealtimePricePerMinValidator = groupDescAudioRealtimePricePerMin.Validators[0].(func(float64) error) + // groupDescAudioTtsPricePerMillionChars is the schema descriptor for audio_tts_price_per_million_chars field. + groupDescAudioTtsPricePerMillionChars := groupFields[34].Descriptor() + // group.AudioTtsPricePerMillionCharsValidator is a validator for the "audio_tts_price_per_million_chars" field. It is called by the builders before save. + group.AudioTtsPricePerMillionCharsValidator = groupDescAudioTtsPricePerMillionChars.Validators[0].(func(float64) error) + // groupDescAudioSttPricePerHour is the schema descriptor for audio_stt_price_per_hour field. + groupDescAudioSttPricePerHour := groupFields[35].Descriptor() + // group.AudioSttPricePerHourValidator is a validator for the "audio_stt_price_per_hour" field. It is called by the builders before save. + group.AudioSttPricePerHourValidator = groupDescAudioSttPricePerHour.Validators[0].(func(float64) error) // groupDescClaudeCodeOnly is the schema descriptor for claude_code_only field. - groupDescClaudeCodeOnly := groupFields[31].Descriptor() + groupDescClaudeCodeOnly := groupFields[36].Descriptor() // group.DefaultClaudeCodeOnly holds the default value on creation for the claude_code_only field. group.DefaultClaudeCodeOnly = groupDescClaudeCodeOnly.Default.(bool) // groupDescModelRoutingEnabled is the schema descriptor for model_routing_enabled field. - groupDescModelRoutingEnabled := groupFields[35].Descriptor() + groupDescModelRoutingEnabled := groupFields[40].Descriptor() // group.DefaultModelRoutingEnabled holds the default value on creation for the model_routing_enabled field. group.DefaultModelRoutingEnabled = groupDescModelRoutingEnabled.Default.(bool) // groupDescMcpXMLInject is the schema descriptor for mcp_xml_inject field. - groupDescMcpXMLInject := groupFields[36].Descriptor() + groupDescMcpXMLInject := groupFields[41].Descriptor() // group.DefaultMcpXMLInject holds the default value on creation for the mcp_xml_inject field. group.DefaultMcpXMLInject = groupDescMcpXMLInject.Default.(bool) // groupDescSupportedModelScopes is the schema descriptor for supported_model_scopes field. - groupDescSupportedModelScopes := groupFields[37].Descriptor() + groupDescSupportedModelScopes := groupFields[42].Descriptor() // group.DefaultSupportedModelScopes holds the default value on creation for the supported_model_scopes field. group.DefaultSupportedModelScopes = groupDescSupportedModelScopes.Default.([]string) // groupDescSortOrder is the schema descriptor for sort_order field. - groupDescSortOrder := groupFields[38].Descriptor() + groupDescSortOrder := groupFields[43].Descriptor() // group.DefaultSortOrder holds the default value on creation for the sort_order field. group.DefaultSortOrder = groupDescSortOrder.Default.(int) // groupDescAllowMessagesDispatch is the schema descriptor for allow_messages_dispatch field. - groupDescAllowMessagesDispatch := groupFields[39].Descriptor() + groupDescAllowMessagesDispatch := groupFields[44].Descriptor() // group.DefaultAllowMessagesDispatch holds the default value on creation for the allow_messages_dispatch field. group.DefaultAllowMessagesDispatch = groupDescAllowMessagesDispatch.Default.(bool) // groupDescAllowLive is the schema descriptor for allow_live field. - groupDescAllowLive := groupFields[40].Descriptor() + groupDescAllowLive := groupFields[45].Descriptor() // group.DefaultAllowLive holds the default value on creation for the allow_live field. group.DefaultAllowLive = groupDescAllowLive.Default.(bool) // groupDescRequireOauthOnly is the schema descriptor for require_oauth_only field. - groupDescRequireOauthOnly := groupFields[41].Descriptor() + groupDescRequireOauthOnly := groupFields[46].Descriptor() // group.DefaultRequireOauthOnly holds the default value on creation for the require_oauth_only field. group.DefaultRequireOauthOnly = groupDescRequireOauthOnly.Default.(bool) // groupDescRequirePrivacySet is the schema descriptor for require_privacy_set field. - groupDescRequirePrivacySet := groupFields[42].Descriptor() + groupDescRequirePrivacySet := groupFields[47].Descriptor() // group.DefaultRequirePrivacySet holds the default value on creation for the require_privacy_set field. group.DefaultRequirePrivacySet = groupDescRequirePrivacySet.Default.(bool) // groupDescDefaultMappedModel is the schema descriptor for default_mapped_model field. - groupDescDefaultMappedModel := groupFields[43].Descriptor() + groupDescDefaultMappedModel := groupFields[48].Descriptor() // group.DefaultDefaultMappedModel holds the default value on creation for the default_mapped_model field. group.DefaultDefaultMappedModel = groupDescDefaultMappedModel.Default.(string) // group.DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. group.DefaultMappedModelValidator = groupDescDefaultMappedModel.Validators[0].(func(string) error) // groupDescMessagesDispatchModelConfig is the schema descriptor for messages_dispatch_model_config field. - groupDescMessagesDispatchModelConfig := groupFields[44].Descriptor() + groupDescMessagesDispatchModelConfig := groupFields[49].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[45].Descriptor() + groupDescModelsListConfig := groupFields[50].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[46].Descriptor() + groupDescRpmLimit := groupFields[51].Descriptor() // group.DefaultRpmLimit holds the default value on creation for the rpm_limit field. group.DefaultRpmLimit = groupDescRpmLimit.Default.(int) // groupDescMaxReasoningEffort is the schema descriptor for max_reasoning_effort field. - groupDescMaxReasoningEffort := groupFields[47].Descriptor() + groupDescMaxReasoningEffort := groupFields[52].Descriptor() // group.DefaultMaxReasoningEffort holds the default value on creation for the max_reasoning_effort field. group.DefaultMaxReasoningEffort = groupDescMaxReasoningEffort.Default.(string) // group.MaxReasoningEffortValidator is a validator for the "max_reasoning_effort" field. It is called by the builders before save. group.MaxReasoningEffortValidator = groupDescMaxReasoningEffort.Validators[0].(func(string) error) // groupDescReasoningEffortMappings is the schema descriptor for reasoning_effort_mappings field. - groupDescReasoningEffortMappings := groupFields[48].Descriptor() + groupDescReasoningEffortMappings := groupFields[53].Descriptor() // group.DefaultReasoningEffortMappings holds the default value on creation for the reasoning_effort_mappings field. group.DefaultReasoningEffortMappings = groupDescReasoningEffortMappings.Default.([]domain.ReasoningEffortMapping) // groupDescProfitControlEnabled is the schema descriptor for profit_control_enabled field. - groupDescProfitControlEnabled := groupFields[49].Descriptor() + groupDescProfitControlEnabled := groupFields[54].Descriptor() // group.DefaultProfitControlEnabled holds the default value on creation for the profit_control_enabled field. group.DefaultProfitControlEnabled = groupDescProfitControlEnabled.Default.(bool) // groupDescProfitMinMargin is the schema descriptor for profit_min_margin field. - groupDescProfitMinMargin := groupFields[50].Descriptor() + groupDescProfitMinMargin := groupFields[55].Descriptor() // group.DefaultProfitMinMargin holds the default value on creation for the profit_min_margin field. group.DefaultProfitMinMargin = groupDescProfitMinMargin.Default.(float64) // groupDescProfitSafetyBuffer is the schema descriptor for profit_safety_buffer field. - groupDescProfitSafetyBuffer := groupFields[51].Descriptor() + groupDescProfitSafetyBuffer := groupFields[56].Descriptor() // group.DefaultProfitSafetyBuffer holds the default value on creation for the profit_safety_buffer field. group.DefaultProfitSafetyBuffer = groupDescProfitSafetyBuffer.Default.(float64) idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin() @@ -1982,124 +1998,128 @@ func init() { usagelogDescUpstreamModel := usagelogFields[6].Descriptor() // usagelog.UpstreamModelValidator is a validator for the "upstream_model" field. It is called by the builders before save. usagelog.UpstreamModelValidator = usagelogDescUpstreamModel.Validators[0].(func(string) error) + // usagelogDescUpstreamResponseModel is the schema descriptor for upstream_response_model field. + usagelogDescUpstreamResponseModel := usagelogFields[7].Descriptor() + // usagelog.UpstreamResponseModelValidator is a validator for the "upstream_response_model" field. It is called by the builders before save. + usagelog.UpstreamResponseModelValidator = usagelogDescUpstreamResponseModel.Validators[0].(func(string) error) // usagelogDescModelMappingChain is the schema descriptor for model_mapping_chain field. - usagelogDescModelMappingChain := usagelogFields[8].Descriptor() + usagelogDescModelMappingChain := usagelogFields[10].Descriptor() // usagelog.ModelMappingChainValidator is a validator for the "model_mapping_chain" field. It is called by the builders before save. usagelog.ModelMappingChainValidator = usagelogDescModelMappingChain.Validators[0].(func(string) error) // usagelogDescBillingTier is the schema descriptor for billing_tier field. - usagelogDescBillingTier := usagelogFields[9].Descriptor() + usagelogDescBillingTier := usagelogFields[11].Descriptor() // usagelog.BillingTierValidator is a validator for the "billing_tier" field. It is called by the builders before save. usagelog.BillingTierValidator = usagelogDescBillingTier.Validators[0].(func(string) error) // usagelogDescBillingMode is the schema descriptor for billing_mode field. - usagelogDescBillingMode := usagelogFields[10].Descriptor() + usagelogDescBillingMode := usagelogFields[12].Descriptor() // usagelog.BillingModeValidator is a validator for the "billing_mode" field. It is called by the builders before save. usagelog.BillingModeValidator = usagelogDescBillingMode.Validators[0].(func(string) error) // usagelogDescInputTokens is the schema descriptor for input_tokens field. - usagelogDescInputTokens := usagelogFields[13].Descriptor() + usagelogDescInputTokens := usagelogFields[15].Descriptor() // usagelog.DefaultInputTokens holds the default value on creation for the input_tokens field. usagelog.DefaultInputTokens = usagelogDescInputTokens.Default.(int) // usagelogDescOutputTokens is the schema descriptor for output_tokens field. - usagelogDescOutputTokens := usagelogFields[14].Descriptor() + usagelogDescOutputTokens := usagelogFields[16].Descriptor() // usagelog.DefaultOutputTokens holds the default value on creation for the output_tokens field. usagelog.DefaultOutputTokens = usagelogDescOutputTokens.Default.(int) // usagelogDescCacheCreationTokens is the schema descriptor for cache_creation_tokens field. - usagelogDescCacheCreationTokens := usagelogFields[15].Descriptor() + usagelogDescCacheCreationTokens := usagelogFields[17].Descriptor() // usagelog.DefaultCacheCreationTokens holds the default value on creation for the cache_creation_tokens field. usagelog.DefaultCacheCreationTokens = usagelogDescCacheCreationTokens.Default.(int) // usagelogDescCacheReadTokens is the schema descriptor for cache_read_tokens field. - usagelogDescCacheReadTokens := usagelogFields[16].Descriptor() + usagelogDescCacheReadTokens := usagelogFields[18].Descriptor() // usagelog.DefaultCacheReadTokens holds the default value on creation for the cache_read_tokens field. usagelog.DefaultCacheReadTokens = usagelogDescCacheReadTokens.Default.(int) // usagelogDescCacheCreation5mTokens is the schema descriptor for cache_creation_5m_tokens field. - usagelogDescCacheCreation5mTokens := usagelogFields[17].Descriptor() + usagelogDescCacheCreation5mTokens := usagelogFields[19].Descriptor() // usagelog.DefaultCacheCreation5mTokens holds the default value on creation for the cache_creation_5m_tokens field. usagelog.DefaultCacheCreation5mTokens = usagelogDescCacheCreation5mTokens.Default.(int) // usagelogDescCacheCreation1hTokens is the schema descriptor for cache_creation_1h_tokens field. - usagelogDescCacheCreation1hTokens := usagelogFields[18].Descriptor() + usagelogDescCacheCreation1hTokens := usagelogFields[20].Descriptor() // usagelog.DefaultCacheCreation1hTokens holds the default value on creation for the cache_creation_1h_tokens field. usagelog.DefaultCacheCreation1hTokens = usagelogDescCacheCreation1hTokens.Default.(int) // usagelogDescInputCost is the schema descriptor for input_cost field. - usagelogDescInputCost := usagelogFields[19].Descriptor() + usagelogDescInputCost := usagelogFields[21].Descriptor() // usagelog.DefaultInputCost holds the default value on creation for the input_cost field. usagelog.DefaultInputCost = usagelogDescInputCost.Default.(float64) // usagelogDescOutputCost is the schema descriptor for output_cost field. - usagelogDescOutputCost := usagelogFields[20].Descriptor() + usagelogDescOutputCost := usagelogFields[22].Descriptor() // usagelog.DefaultOutputCost holds the default value on creation for the output_cost field. usagelog.DefaultOutputCost = usagelogDescOutputCost.Default.(float64) // usagelogDescCacheCreationCost is the schema descriptor for cache_creation_cost field. - usagelogDescCacheCreationCost := usagelogFields[21].Descriptor() + usagelogDescCacheCreationCost := usagelogFields[23].Descriptor() // usagelog.DefaultCacheCreationCost holds the default value on creation for the cache_creation_cost field. usagelog.DefaultCacheCreationCost = usagelogDescCacheCreationCost.Default.(float64) // usagelogDescCacheReadCost is the schema descriptor for cache_read_cost field. - usagelogDescCacheReadCost := usagelogFields[22].Descriptor() + usagelogDescCacheReadCost := usagelogFields[24].Descriptor() // usagelog.DefaultCacheReadCost holds the default value on creation for the cache_read_cost field. usagelog.DefaultCacheReadCost = usagelogDescCacheReadCost.Default.(float64) // usagelogDescTotalCost is the schema descriptor for total_cost field. - usagelogDescTotalCost := usagelogFields[23].Descriptor() + usagelogDescTotalCost := usagelogFields[25].Descriptor() // usagelog.DefaultTotalCost holds the default value on creation for the total_cost field. usagelog.DefaultTotalCost = usagelogDescTotalCost.Default.(float64) // usagelogDescActualCost is the schema descriptor for actual_cost field. - usagelogDescActualCost := usagelogFields[24].Descriptor() + usagelogDescActualCost := usagelogFields[26].Descriptor() // usagelog.DefaultActualCost holds the default value on creation for the actual_cost field. usagelog.DefaultActualCost = usagelogDescActualCost.Default.(float64) // usagelogDescRateMultiplier is the schema descriptor for rate_multiplier field. - usagelogDescRateMultiplier := usagelogFields[25].Descriptor() + usagelogDescRateMultiplier := usagelogFields[27].Descriptor() // usagelog.DefaultRateMultiplier holds the default value on creation for the rate_multiplier field. usagelog.DefaultRateMultiplier = usagelogDescRateMultiplier.Default.(float64) // usagelogDescLongContextBillingApplied is the schema descriptor for long_context_billing_applied field. - usagelogDescLongContextBillingApplied := usagelogFields[26].Descriptor() + usagelogDescLongContextBillingApplied := usagelogFields[28].Descriptor() // usagelog.DefaultLongContextBillingApplied holds the default value on creation for the long_context_billing_applied field. usagelog.DefaultLongContextBillingApplied = usagelogDescLongContextBillingApplied.Default.(bool) // usagelogDescBillingType is the schema descriptor for billing_type field. - usagelogDescBillingType := usagelogFields[28].Descriptor() + usagelogDescBillingType := usagelogFields[30].Descriptor() // usagelog.DefaultBillingType holds the default value on creation for the billing_type field. usagelog.DefaultBillingType = usagelogDescBillingType.Default.(int8) // usagelogDescStream is the schema descriptor for stream field. - usagelogDescStream := usagelogFields[29].Descriptor() + usagelogDescStream := usagelogFields[31].Descriptor() // usagelog.DefaultStream holds the default value on creation for the stream field. usagelog.DefaultStream = usagelogDescStream.Default.(bool) // usagelogDescUserAgent is the schema descriptor for user_agent field. - usagelogDescUserAgent := usagelogFields[32].Descriptor() + usagelogDescUserAgent := usagelogFields[34].Descriptor() // usagelog.UserAgentValidator is a validator for the "user_agent" field. It is called by the builders before save. usagelog.UserAgentValidator = usagelogDescUserAgent.Validators[0].(func(string) error) // usagelogDescIPAddress is the schema descriptor for ip_address field. - usagelogDescIPAddress := usagelogFields[33].Descriptor() + usagelogDescIPAddress := usagelogFields[35].Descriptor() // usagelog.IPAddressValidator is a validator for the "ip_address" field. It is called by the builders before save. usagelog.IPAddressValidator = usagelogDescIPAddress.Validators[0].(func(string) error) // usagelogDescImageCount is the schema descriptor for image_count field. - usagelogDescImageCount := usagelogFields[34].Descriptor() + usagelogDescImageCount := usagelogFields[36].Descriptor() // usagelog.DefaultImageCount holds the default value on creation for the image_count field. usagelog.DefaultImageCount = usagelogDescImageCount.Default.(int) // usagelogDescImageSize is the schema descriptor for image_size field. - usagelogDescImageSize := usagelogFields[35].Descriptor() + usagelogDescImageSize := usagelogFields[37].Descriptor() // usagelog.ImageSizeValidator is a validator for the "image_size" field. It is called by the builders before save. usagelog.ImageSizeValidator = usagelogDescImageSize.Validators[0].(func(string) error) // usagelogDescImageInputSize is the schema descriptor for image_input_size field. - usagelogDescImageInputSize := usagelogFields[36].Descriptor() + usagelogDescImageInputSize := usagelogFields[38].Descriptor() // usagelog.ImageInputSizeValidator is a validator for the "image_input_size" field. It is called by the builders before save. usagelog.ImageInputSizeValidator = usagelogDescImageInputSize.Validators[0].(func(string) error) // usagelogDescImageOutputSize is the schema descriptor for image_output_size field. - usagelogDescImageOutputSize := usagelogFields[37].Descriptor() + usagelogDescImageOutputSize := usagelogFields[39].Descriptor() // usagelog.ImageOutputSizeValidator is a validator for the "image_output_size" field. It is called by the builders before save. usagelog.ImageOutputSizeValidator = usagelogDescImageOutputSize.Validators[0].(func(string) error) // usagelogDescImageSizeSource is the schema descriptor for image_size_source field. - usagelogDescImageSizeSource := usagelogFields[38].Descriptor() + usagelogDescImageSizeSource := usagelogFields[40].Descriptor() // usagelog.ImageSizeSourceValidator is a validator for the "image_size_source" field. It is called by the builders before save. usagelog.ImageSizeSourceValidator = usagelogDescImageSizeSource.Validators[0].(func(string) error) // usagelogDescVideoCount is the schema descriptor for video_count field. - usagelogDescVideoCount := usagelogFields[40].Descriptor() + usagelogDescVideoCount := usagelogFields[42].Descriptor() // usagelog.DefaultVideoCount holds the default value on creation for the video_count field. usagelog.DefaultVideoCount = usagelogDescVideoCount.Default.(int) // usagelogDescVideoResolution is the schema descriptor for video_resolution field. - usagelogDescVideoResolution := usagelogFields[41].Descriptor() + usagelogDescVideoResolution := usagelogFields[43].Descriptor() // usagelog.VideoResolutionValidator is a validator for the "video_resolution" field. It is called by the builders before save. usagelog.VideoResolutionValidator = usagelogDescVideoResolution.Validators[0].(func(string) error) // usagelogDescCacheTTLOverridden is the schema descriptor for cache_ttl_overridden field. - usagelogDescCacheTTLOverridden := usagelogFields[43].Descriptor() + usagelogDescCacheTTLOverridden := usagelogFields[45].Descriptor() // usagelog.DefaultCacheTTLOverridden holds the default value on creation for the cache_ttl_overridden field. usagelog.DefaultCacheTTLOverridden = usagelogDescCacheTTLOverridden.Default.(bool) // usagelogDescCreatedAt is the schema descriptor for created_at field. - usagelogDescCreatedAt := usagelogFields[44].Descriptor() + usagelogDescCreatedAt := usagelogFields[46].Descriptor() // usagelog.DefaultCreatedAt holds the default value on creation for the created_at field. usagelog.DefaultCreatedAt = usagelogDescCreatedAt.Default.(func() time.Time) userMixin := schema.User{}.Mixin() diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go index 8a9fc12ded..85f86891b8 100644 --- a/backend/ent/schema/group.go +++ b/backend/ent/schema/group.go @@ -148,12 +148,44 @@ func (Group) Fields() []ent.Field { Optional(). Nillable(). SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}), + field.JSON("video_model_prices", map[string]map[string]float64{}). + Optional(). + SchemaType(map[string]string{dialect.Postgres: "jsonb"}). + Comment("按模型族和分辨率覆盖视频每秒价格"), field.Float("web_search_price_per_call"). Optional(). Nillable(). SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}). Comment("Codex alpha/search 网页搜索单次价格(USD/次);nil 表示使用默认价 0.01(官方 $10/1000 次)"), + // 搜索/工具调用显式定价(per 1k calls),用于 Grok web_search 等。 + field.Float("search_price_per_1k"). + Optional(). + Nillable(). + Min(0). + SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}). + Comment("搜索工具价格 per 1000 calls(web_search 等)"), + + // Grok Voice 显式定价(realtime / TTS / STT),不按文本 RateMultiplier。 + field.Float("audio_realtime_price_per_min"). + Optional(). + Nillable(). + Min(0). + SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}). + Comment("Voice realtime 每分钟价格(USD)"), + field.Float("audio_tts_price_per_million_chars"). + Optional(). + Nillable(). + Min(0). + SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}). + Comment("TTS 每百万字符价格(USD)"), + field.Float("audio_stt_price_per_hour"). + Optional(). + Nillable(). + Min(0). + SchemaType(map[string]string{dialect.Postgres: "decimal(20,8)"}). + Comment("STT 每小时价格(USD)"), + // Claude Code 客户端限制 (added by migration 029) field.Bool("claude_code_only"). Default(false). diff --git a/backend/ent/schema/usage_log.go b/backend/ent/schema/usage_log.go index 6d8c2d4191..f13d07170c 100644 --- a/backend/ent/schema/usage_log.go +++ b/backend/ent/schema/usage_log.go @@ -53,6 +53,17 @@ func (UsageLog) Fields() []ent.Field { MaxLen(100). Optional(). Nillable(), + // UpstreamResponseModel stores the model name declared by the upstream + // response before any protocol conversion or client-facing rewrite. + field.String("upstream_response_model"). + MaxLen(200). + Optional(). + Nillable(), + // UpstreamModelMismatch is tri-state: NULL means the upstream response did + // not declare a model (or predates this field); false/true means observed. + field.Bool("upstream_model_mismatch"). + Optional(). + Nillable(), field.Int64("channel_id").Optional().Nillable().Comment("渠道 ID"), field.String("model_mapping_chain").MaxLen(500).Optional().Nillable().Comment("模型映射链"), field.String("billing_tier").MaxLen(50).Optional().Nillable().Comment("计费层级标签"), diff --git a/backend/ent/usagelog.go b/backend/ent/usagelog.go index b13e29b2f7..1bd424f8d5 100644 --- a/backend/ent/usagelog.go +++ b/backend/ent/usagelog.go @@ -37,6 +37,10 @@ type UsageLog struct { RequestedModel *string `json:"requested_model,omitempty"` // UpstreamModel holds the value of the "upstream_model" field. UpstreamModel *string `json:"upstream_model,omitempty"` + // UpstreamResponseModel holds the value of the "upstream_response_model" field. + UpstreamResponseModel *string `json:"upstream_response_model,omitempty"` + // UpstreamModelMismatch holds the value of the "upstream_model_mismatch" field. + UpstreamModelMismatch *bool `json:"upstream_model_mismatch,omitempty"` // 渠道 ID ChannelID *int64 `json:"channel_id,omitempty"` // 模型映射链 @@ -198,13 +202,13 @@ func (*UsageLog) scanValues(columns []string) ([]any, error) { switch columns[i] { case usagelog.FieldImageSizeBreakdown: values[i] = new([]byte) - case usagelog.FieldLongContextBillingApplied, usagelog.FieldStream, usagelog.FieldCacheTTLOverridden: + case usagelog.FieldUpstreamModelMismatch, usagelog.FieldLongContextBillingApplied, usagelog.FieldStream, usagelog.FieldCacheTTLOverridden: values[i] = new(sql.NullBool) case usagelog.FieldInputCost, usagelog.FieldOutputCost, usagelog.FieldCacheCreationCost, usagelog.FieldCacheReadCost, usagelog.FieldTotalCost, usagelog.FieldActualCost, usagelog.FieldRateMultiplier, usagelog.FieldAccountRateMultiplier: values[i] = new(sql.NullFloat64) case usagelog.FieldID, usagelog.FieldUserID, usagelog.FieldAPIKeyID, usagelog.FieldAccountID, usagelog.FieldChannelID, usagelog.FieldGroupID, usagelog.FieldSubscriptionID, usagelog.FieldInputTokens, usagelog.FieldOutputTokens, usagelog.FieldCacheCreationTokens, usagelog.FieldCacheReadTokens, usagelog.FieldCacheCreation5mTokens, usagelog.FieldCacheCreation1hTokens, usagelog.FieldBillingType, usagelog.FieldDurationMs, usagelog.FieldFirstTokenMs, usagelog.FieldImageCount, usagelog.FieldVideoCount, usagelog.FieldVideoDurationSeconds: values[i] = new(sql.NullInt64) - case usagelog.FieldRequestID, usagelog.FieldModel, usagelog.FieldRequestedModel, usagelog.FieldUpstreamModel, usagelog.FieldModelMappingChain, usagelog.FieldBillingTier, usagelog.FieldBillingMode, usagelog.FieldUserAgent, usagelog.FieldIPAddress, usagelog.FieldImageSize, usagelog.FieldImageInputSize, usagelog.FieldImageOutputSize, usagelog.FieldImageSizeSource, usagelog.FieldVideoResolution: + case usagelog.FieldRequestID, usagelog.FieldModel, usagelog.FieldRequestedModel, usagelog.FieldUpstreamModel, usagelog.FieldUpstreamResponseModel, usagelog.FieldModelMappingChain, usagelog.FieldBillingTier, usagelog.FieldBillingMode, usagelog.FieldUserAgent, usagelog.FieldIPAddress, usagelog.FieldImageSize, usagelog.FieldImageInputSize, usagelog.FieldImageOutputSize, usagelog.FieldImageSizeSource, usagelog.FieldVideoResolution: values[i] = new(sql.NullString) case usagelog.FieldCreatedAt: values[i] = new(sql.NullTime) @@ -273,6 +277,20 @@ func (_m *UsageLog) assignValues(columns []string, values []any) error { _m.UpstreamModel = new(string) *_m.UpstreamModel = value.String } + case usagelog.FieldUpstreamResponseModel: + if value, ok := values[i].(*sql.NullString); !ok { + return fmt.Errorf("unexpected type %T for field upstream_response_model", values[i]) + } else if value.Valid { + _m.UpstreamResponseModel = new(string) + *_m.UpstreamResponseModel = value.String + } + case usagelog.FieldUpstreamModelMismatch: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field upstream_model_mismatch", values[i]) + } else if value.Valid { + _m.UpstreamModelMismatch = new(bool) + *_m.UpstreamModelMismatch = value.Bool + } case usagelog.FieldChannelID: if value, ok := values[i].(*sql.NullInt64); !ok { return fmt.Errorf("unexpected type %T for field channel_id", values[i]) @@ -606,6 +624,16 @@ func (_m *UsageLog) String() string { builder.WriteString(*v) } builder.WriteString(", ") + if v := _m.UpstreamResponseModel; v != nil { + builder.WriteString("upstream_response_model=") + builder.WriteString(*v) + } + builder.WriteString(", ") + if v := _m.UpstreamModelMismatch; v != nil { + builder.WriteString("upstream_model_mismatch=") + builder.WriteString(fmt.Sprintf("%v", *v)) + } + builder.WriteString(", ") if v := _m.ChannelID; v != nil { builder.WriteString("channel_id=") builder.WriteString(fmt.Sprintf("%v", *v)) diff --git a/backend/ent/usagelog/usagelog.go b/backend/ent/usagelog/usagelog.go index a87d937195..d3f8dfa1b3 100644 --- a/backend/ent/usagelog/usagelog.go +++ b/backend/ent/usagelog/usagelog.go @@ -28,6 +28,10 @@ const ( FieldRequestedModel = "requested_model" // FieldUpstreamModel holds the string denoting the upstream_model field in the database. FieldUpstreamModel = "upstream_model" + // FieldUpstreamResponseModel holds the string denoting the upstream_response_model field in the database. + FieldUpstreamResponseModel = "upstream_response_model" + // FieldUpstreamModelMismatch holds the string denoting the upstream_model_mismatch field in the database. + FieldUpstreamModelMismatch = "upstream_model_mismatch" // FieldChannelID holds the string denoting the channel_id field in the database. FieldChannelID = "channel_id" // FieldModelMappingChain holds the string denoting the model_mapping_chain field in the database. @@ -163,6 +167,8 @@ var Columns = []string{ FieldModel, FieldRequestedModel, FieldUpstreamModel, + FieldUpstreamResponseModel, + FieldUpstreamModelMismatch, FieldChannelID, FieldModelMappingChain, FieldBillingTier, @@ -222,6 +228,8 @@ var ( RequestedModelValidator func(string) error // UpstreamModelValidator is a validator for the "upstream_model" field. It is called by the builders before save. UpstreamModelValidator func(string) error + // UpstreamResponseModelValidator is a validator for the "upstream_response_model" field. It is called by the builders before save. + UpstreamResponseModelValidator func(string) error // ModelMappingChainValidator is a validator for the "model_mapping_chain" field. It is called by the builders before save. ModelMappingChainValidator func(string) error // BillingTierValidator is a validator for the "billing_tier" field. It is called by the builders before save. @@ -327,6 +335,16 @@ func ByUpstreamModel(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldUpstreamModel, opts...).ToFunc() } +// ByUpstreamResponseModel orders the results by the upstream_response_model field. +func ByUpstreamResponseModel(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldUpstreamResponseModel, opts...).ToFunc() +} + +// ByUpstreamModelMismatch orders the results by the upstream_model_mismatch field. +func ByUpstreamModelMismatch(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldUpstreamModelMismatch, opts...).ToFunc() +} + // ByChannelID orders the results by the channel_id field. func ByChannelID(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldChannelID, opts...).ToFunc() diff --git a/backend/ent/usagelog/where.go b/backend/ent/usagelog/where.go index a9462e0d0e..7982d74c5f 100644 --- a/backend/ent/usagelog/where.go +++ b/backend/ent/usagelog/where.go @@ -90,6 +90,16 @@ func UpstreamModel(v string) predicate.UsageLog { return predicate.UsageLog(sql.FieldEQ(FieldUpstreamModel, v)) } +// UpstreamResponseModel applies equality check predicate on the "upstream_response_model" field. It's identical to UpstreamResponseModelEQ. +func UpstreamResponseModel(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldEQ(FieldUpstreamResponseModel, v)) +} + +// UpstreamModelMismatch applies equality check predicate on the "upstream_model_mismatch" field. It's identical to UpstreamModelMismatchEQ. +func UpstreamModelMismatch(v bool) predicate.UsageLog { + return predicate.UsageLog(sql.FieldEQ(FieldUpstreamModelMismatch, v)) +} + // ChannelID applies equality check predicate on the "channel_id" field. It's identical to ChannelIDEQ. func ChannelID(v int64) predicate.UsageLog { return predicate.UsageLog(sql.FieldEQ(FieldChannelID, v)) @@ -615,6 +625,101 @@ func UpstreamModelContainsFold(v string) predicate.UsageLog { return predicate.UsageLog(sql.FieldContainsFold(FieldUpstreamModel, v)) } +// UpstreamResponseModelEQ applies the EQ predicate on the "upstream_response_model" field. +func UpstreamResponseModelEQ(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldEQ(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelNEQ applies the NEQ predicate on the "upstream_response_model" field. +func UpstreamResponseModelNEQ(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldNEQ(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelIn applies the In predicate on the "upstream_response_model" field. +func UpstreamResponseModelIn(vs ...string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldIn(FieldUpstreamResponseModel, vs...)) +} + +// UpstreamResponseModelNotIn applies the NotIn predicate on the "upstream_response_model" field. +func UpstreamResponseModelNotIn(vs ...string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldNotIn(FieldUpstreamResponseModel, vs...)) +} + +// UpstreamResponseModelGT applies the GT predicate on the "upstream_response_model" field. +func UpstreamResponseModelGT(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldGT(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelGTE applies the GTE predicate on the "upstream_response_model" field. +func UpstreamResponseModelGTE(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldGTE(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelLT applies the LT predicate on the "upstream_response_model" field. +func UpstreamResponseModelLT(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldLT(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelLTE applies the LTE predicate on the "upstream_response_model" field. +func UpstreamResponseModelLTE(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldLTE(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelContains applies the Contains predicate on the "upstream_response_model" field. +func UpstreamResponseModelContains(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldContains(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelHasPrefix applies the HasPrefix predicate on the "upstream_response_model" field. +func UpstreamResponseModelHasPrefix(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldHasPrefix(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelHasSuffix applies the HasSuffix predicate on the "upstream_response_model" field. +func UpstreamResponseModelHasSuffix(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldHasSuffix(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelIsNil applies the IsNil predicate on the "upstream_response_model" field. +func UpstreamResponseModelIsNil() predicate.UsageLog { + return predicate.UsageLog(sql.FieldIsNull(FieldUpstreamResponseModel)) +} + +// UpstreamResponseModelNotNil applies the NotNil predicate on the "upstream_response_model" field. +func UpstreamResponseModelNotNil() predicate.UsageLog { + return predicate.UsageLog(sql.FieldNotNull(FieldUpstreamResponseModel)) +} + +// UpstreamResponseModelEqualFold applies the EqualFold predicate on the "upstream_response_model" field. +func UpstreamResponseModelEqualFold(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldEqualFold(FieldUpstreamResponseModel, v)) +} + +// UpstreamResponseModelContainsFold applies the ContainsFold predicate on the "upstream_response_model" field. +func UpstreamResponseModelContainsFold(v string) predicate.UsageLog { + return predicate.UsageLog(sql.FieldContainsFold(FieldUpstreamResponseModel, v)) +} + +// UpstreamModelMismatchEQ applies the EQ predicate on the "upstream_model_mismatch" field. +func UpstreamModelMismatchEQ(v bool) predicate.UsageLog { + return predicate.UsageLog(sql.FieldEQ(FieldUpstreamModelMismatch, v)) +} + +// UpstreamModelMismatchNEQ applies the NEQ predicate on the "upstream_model_mismatch" field. +func UpstreamModelMismatchNEQ(v bool) predicate.UsageLog { + return predicate.UsageLog(sql.FieldNEQ(FieldUpstreamModelMismatch, v)) +} + +// UpstreamModelMismatchIsNil applies the IsNil predicate on the "upstream_model_mismatch" field. +func UpstreamModelMismatchIsNil() predicate.UsageLog { + return predicate.UsageLog(sql.FieldIsNull(FieldUpstreamModelMismatch)) +} + +// UpstreamModelMismatchNotNil applies the NotNil predicate on the "upstream_model_mismatch" field. +func UpstreamModelMismatchNotNil() predicate.UsageLog { + return predicate.UsageLog(sql.FieldNotNull(FieldUpstreamModelMismatch)) +} + // ChannelIDEQ applies the EQ predicate on the "channel_id" field. func ChannelIDEQ(v int64) predicate.UsageLog { return predicate.UsageLog(sql.FieldEQ(FieldChannelID, v)) diff --git a/backend/ent/usagelog_create.go b/backend/ent/usagelog_create.go index 31cf45328e..8178fa0710 100644 --- a/backend/ent/usagelog_create.go +++ b/backend/ent/usagelog_create.go @@ -85,6 +85,34 @@ func (_c *UsageLogCreate) SetNillableUpstreamModel(v *string) *UsageLogCreate { return _c } +// SetUpstreamResponseModel sets the "upstream_response_model" field. +func (_c *UsageLogCreate) SetUpstreamResponseModel(v string) *UsageLogCreate { + _c.mutation.SetUpstreamResponseModel(v) + return _c +} + +// SetNillableUpstreamResponseModel sets the "upstream_response_model" field if the given value is not nil. +func (_c *UsageLogCreate) SetNillableUpstreamResponseModel(v *string) *UsageLogCreate { + if v != nil { + _c.SetUpstreamResponseModel(*v) + } + return _c +} + +// SetUpstreamModelMismatch sets the "upstream_model_mismatch" field. +func (_c *UsageLogCreate) SetUpstreamModelMismatch(v bool) *UsageLogCreate { + _c.mutation.SetUpstreamModelMismatch(v) + return _c +} + +// SetNillableUpstreamModelMismatch sets the "upstream_model_mismatch" field if the given value is not nil. +func (_c *UsageLogCreate) SetNillableUpstreamModelMismatch(v *bool) *UsageLogCreate { + if v != nil { + _c.SetUpstreamModelMismatch(*v) + } + return _c +} + // SetChannelID sets the "channel_id" field. func (_c *UsageLogCreate) SetChannelID(v int64) *UsageLogCreate { _c.mutation.SetChannelID(v) @@ -788,6 +816,11 @@ func (_c *UsageLogCreate) check() error { return &ValidationError{Name: "upstream_model", err: fmt.Errorf(`ent: validator failed for field "UsageLog.upstream_model": %w`, err)} } } + if v, ok := _c.mutation.UpstreamResponseModel(); ok { + if err := usagelog.UpstreamResponseModelValidator(v); err != nil { + return &ValidationError{Name: "upstream_response_model", err: fmt.Errorf(`ent: validator failed for field "UsageLog.upstream_response_model": %w`, err)} + } + } if v, ok := _c.mutation.ModelMappingChain(); ok { if err := usagelog.ModelMappingChainValidator(v); err != nil { return &ValidationError{Name: "model_mapping_chain", err: fmt.Errorf(`ent: validator failed for field "UsageLog.model_mapping_chain": %w`, err)} @@ -950,6 +983,14 @@ func (_c *UsageLogCreate) createSpec() (*UsageLog, *sqlgraph.CreateSpec) { _spec.SetField(usagelog.FieldUpstreamModel, field.TypeString, value) _node.UpstreamModel = &value } + if value, ok := _c.mutation.UpstreamResponseModel(); ok { + _spec.SetField(usagelog.FieldUpstreamResponseModel, field.TypeString, value) + _node.UpstreamResponseModel = &value + } + if value, ok := _c.mutation.UpstreamModelMismatch(); ok { + _spec.SetField(usagelog.FieldUpstreamModelMismatch, field.TypeBool, value) + _node.UpstreamModelMismatch = &value + } if value, ok := _c.mutation.ChannelID(); ok { _spec.SetField(usagelog.FieldChannelID, field.TypeInt64, value) _node.ChannelID = &value @@ -1327,6 +1368,42 @@ func (u *UsageLogUpsert) ClearUpstreamModel() *UsageLogUpsert { return u } +// SetUpstreamResponseModel sets the "upstream_response_model" field. +func (u *UsageLogUpsert) SetUpstreamResponseModel(v string) *UsageLogUpsert { + u.Set(usagelog.FieldUpstreamResponseModel, v) + return u +} + +// UpdateUpstreamResponseModel sets the "upstream_response_model" field to the value that was provided on create. +func (u *UsageLogUpsert) UpdateUpstreamResponseModel() *UsageLogUpsert { + u.SetExcluded(usagelog.FieldUpstreamResponseModel) + return u +} + +// ClearUpstreamResponseModel clears the value of the "upstream_response_model" field. +func (u *UsageLogUpsert) ClearUpstreamResponseModel() *UsageLogUpsert { + u.SetNull(usagelog.FieldUpstreamResponseModel) + return u +} + +// SetUpstreamModelMismatch sets the "upstream_model_mismatch" field. +func (u *UsageLogUpsert) SetUpstreamModelMismatch(v bool) *UsageLogUpsert { + u.Set(usagelog.FieldUpstreamModelMismatch, v) + return u +} + +// UpdateUpstreamModelMismatch sets the "upstream_model_mismatch" field to the value that was provided on create. +func (u *UsageLogUpsert) UpdateUpstreamModelMismatch() *UsageLogUpsert { + u.SetExcluded(usagelog.FieldUpstreamModelMismatch) + return u +} + +// ClearUpstreamModelMismatch clears the value of the "upstream_model_mismatch" field. +func (u *UsageLogUpsert) ClearUpstreamModelMismatch() *UsageLogUpsert { + u.SetNull(usagelog.FieldUpstreamModelMismatch) + return u +} + // SetChannelID sets the "channel_id" field. func (u *UsageLogUpsert) SetChannelID(v int64) *UsageLogUpsert { u.Set(usagelog.FieldChannelID, v) @@ -2162,6 +2239,48 @@ func (u *UsageLogUpsertOne) ClearUpstreamModel() *UsageLogUpsertOne { }) } +// SetUpstreamResponseModel sets the "upstream_response_model" field. +func (u *UsageLogUpsertOne) SetUpstreamResponseModel(v string) *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.SetUpstreamResponseModel(v) + }) +} + +// UpdateUpstreamResponseModel sets the "upstream_response_model" field to the value that was provided on create. +func (u *UsageLogUpsertOne) UpdateUpstreamResponseModel() *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.UpdateUpstreamResponseModel() + }) +} + +// ClearUpstreamResponseModel clears the value of the "upstream_response_model" field. +func (u *UsageLogUpsertOne) ClearUpstreamResponseModel() *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.ClearUpstreamResponseModel() + }) +} + +// SetUpstreamModelMismatch sets the "upstream_model_mismatch" field. +func (u *UsageLogUpsertOne) SetUpstreamModelMismatch(v bool) *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.SetUpstreamModelMismatch(v) + }) +} + +// UpdateUpstreamModelMismatch sets the "upstream_model_mismatch" field to the value that was provided on create. +func (u *UsageLogUpsertOne) UpdateUpstreamModelMismatch() *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.UpdateUpstreamModelMismatch() + }) +} + +// ClearUpstreamModelMismatch clears the value of the "upstream_model_mismatch" field. +func (u *UsageLogUpsertOne) ClearUpstreamModelMismatch() *UsageLogUpsertOne { + return u.Update(func(s *UsageLogUpsert) { + s.ClearUpstreamModelMismatch() + }) +} + // SetChannelID sets the "channel_id" field. func (u *UsageLogUpsertOne) SetChannelID(v int64) *UsageLogUpsertOne { return u.Update(func(s *UsageLogUpsert) { @@ -3276,6 +3395,48 @@ func (u *UsageLogUpsertBulk) ClearUpstreamModel() *UsageLogUpsertBulk { }) } +// SetUpstreamResponseModel sets the "upstream_response_model" field. +func (u *UsageLogUpsertBulk) SetUpstreamResponseModel(v string) *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.SetUpstreamResponseModel(v) + }) +} + +// UpdateUpstreamResponseModel sets the "upstream_response_model" field to the value that was provided on create. +func (u *UsageLogUpsertBulk) UpdateUpstreamResponseModel() *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.UpdateUpstreamResponseModel() + }) +} + +// ClearUpstreamResponseModel clears the value of the "upstream_response_model" field. +func (u *UsageLogUpsertBulk) ClearUpstreamResponseModel() *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.ClearUpstreamResponseModel() + }) +} + +// SetUpstreamModelMismatch sets the "upstream_model_mismatch" field. +func (u *UsageLogUpsertBulk) SetUpstreamModelMismatch(v bool) *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.SetUpstreamModelMismatch(v) + }) +} + +// UpdateUpstreamModelMismatch sets the "upstream_model_mismatch" field to the value that was provided on create. +func (u *UsageLogUpsertBulk) UpdateUpstreamModelMismatch() *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.UpdateUpstreamModelMismatch() + }) +} + +// ClearUpstreamModelMismatch clears the value of the "upstream_model_mismatch" field. +func (u *UsageLogUpsertBulk) ClearUpstreamModelMismatch() *UsageLogUpsertBulk { + return u.Update(func(s *UsageLogUpsert) { + s.ClearUpstreamModelMismatch() + }) +} + // SetChannelID sets the "channel_id" field. func (u *UsageLogUpsertBulk) SetChannelID(v int64) *UsageLogUpsertBulk { return u.Update(func(s *UsageLogUpsert) { diff --git a/backend/ent/usagelog_update.go b/backend/ent/usagelog_update.go index 2a60d6f44d..6c23ebd355 100644 --- a/backend/ent/usagelog_update.go +++ b/backend/ent/usagelog_update.go @@ -142,6 +142,46 @@ func (_u *UsageLogUpdate) ClearUpstreamModel() *UsageLogUpdate { return _u } +// SetUpstreamResponseModel sets the "upstream_response_model" field. +func (_u *UsageLogUpdate) SetUpstreamResponseModel(v string) *UsageLogUpdate { + _u.mutation.SetUpstreamResponseModel(v) + return _u +} + +// SetNillableUpstreamResponseModel sets the "upstream_response_model" field if the given value is not nil. +func (_u *UsageLogUpdate) SetNillableUpstreamResponseModel(v *string) *UsageLogUpdate { + if v != nil { + _u.SetUpstreamResponseModel(*v) + } + return _u +} + +// ClearUpstreamResponseModel clears the value of the "upstream_response_model" field. +func (_u *UsageLogUpdate) ClearUpstreamResponseModel() *UsageLogUpdate { + _u.mutation.ClearUpstreamResponseModel() + return _u +} + +// SetUpstreamModelMismatch sets the "upstream_model_mismatch" field. +func (_u *UsageLogUpdate) SetUpstreamModelMismatch(v bool) *UsageLogUpdate { + _u.mutation.SetUpstreamModelMismatch(v) + return _u +} + +// SetNillableUpstreamModelMismatch sets the "upstream_model_mismatch" field if the given value is not nil. +func (_u *UsageLogUpdate) SetNillableUpstreamModelMismatch(v *bool) *UsageLogUpdate { + if v != nil { + _u.SetUpstreamModelMismatch(*v) + } + return _u +} + +// ClearUpstreamModelMismatch clears the value of the "upstream_model_mismatch" field. +func (_u *UsageLogUpdate) ClearUpstreamModelMismatch() *UsageLogUpdate { + _u.mutation.ClearUpstreamModelMismatch() + return _u +} + // SetChannelID sets the "channel_id" field. func (_u *UsageLogUpdate) SetChannelID(v int64) *UsageLogUpdate { _u.mutation.ResetChannelID() @@ -1016,6 +1056,11 @@ func (_u *UsageLogUpdate) check() error { return &ValidationError{Name: "upstream_model", err: fmt.Errorf(`ent: validator failed for field "UsageLog.upstream_model": %w`, err)} } } + if v, ok := _u.mutation.UpstreamResponseModel(); ok { + if err := usagelog.UpstreamResponseModelValidator(v); err != nil { + return &ValidationError{Name: "upstream_response_model", err: fmt.Errorf(`ent: validator failed for field "UsageLog.upstream_response_model": %w`, err)} + } + } if v, ok := _u.mutation.ModelMappingChain(); ok { if err := usagelog.ModelMappingChainValidator(v); err != nil { return &ValidationError{Name: "model_mapping_chain", err: fmt.Errorf(`ent: validator failed for field "UsageLog.model_mapping_chain": %w`, err)} @@ -1108,6 +1153,18 @@ func (_u *UsageLogUpdate) sqlSave(ctx context.Context) (_node int, err error) { if _u.mutation.UpstreamModelCleared() { _spec.ClearField(usagelog.FieldUpstreamModel, field.TypeString) } + if value, ok := _u.mutation.UpstreamResponseModel(); ok { + _spec.SetField(usagelog.FieldUpstreamResponseModel, field.TypeString, value) + } + if _u.mutation.UpstreamResponseModelCleared() { + _spec.ClearField(usagelog.FieldUpstreamResponseModel, field.TypeString) + } + if value, ok := _u.mutation.UpstreamModelMismatch(); ok { + _spec.SetField(usagelog.FieldUpstreamModelMismatch, field.TypeBool, value) + } + if _u.mutation.UpstreamModelMismatchCleared() { + _spec.ClearField(usagelog.FieldUpstreamModelMismatch, field.TypeBool) + } if value, ok := _u.mutation.ChannelID(); ok { _spec.SetField(usagelog.FieldChannelID, field.TypeInt64, value) } @@ -1599,6 +1656,46 @@ func (_u *UsageLogUpdateOne) ClearUpstreamModel() *UsageLogUpdateOne { return _u } +// SetUpstreamResponseModel sets the "upstream_response_model" field. +func (_u *UsageLogUpdateOne) SetUpstreamResponseModel(v string) *UsageLogUpdateOne { + _u.mutation.SetUpstreamResponseModel(v) + return _u +} + +// SetNillableUpstreamResponseModel sets the "upstream_response_model" field if the given value is not nil. +func (_u *UsageLogUpdateOne) SetNillableUpstreamResponseModel(v *string) *UsageLogUpdateOne { + if v != nil { + _u.SetUpstreamResponseModel(*v) + } + return _u +} + +// ClearUpstreamResponseModel clears the value of the "upstream_response_model" field. +func (_u *UsageLogUpdateOne) ClearUpstreamResponseModel() *UsageLogUpdateOne { + _u.mutation.ClearUpstreamResponseModel() + return _u +} + +// SetUpstreamModelMismatch sets the "upstream_model_mismatch" field. +func (_u *UsageLogUpdateOne) SetUpstreamModelMismatch(v bool) *UsageLogUpdateOne { + _u.mutation.SetUpstreamModelMismatch(v) + return _u +} + +// SetNillableUpstreamModelMismatch sets the "upstream_model_mismatch" field if the given value is not nil. +func (_u *UsageLogUpdateOne) SetNillableUpstreamModelMismatch(v *bool) *UsageLogUpdateOne { + if v != nil { + _u.SetUpstreamModelMismatch(*v) + } + return _u +} + +// ClearUpstreamModelMismatch clears the value of the "upstream_model_mismatch" field. +func (_u *UsageLogUpdateOne) ClearUpstreamModelMismatch() *UsageLogUpdateOne { + _u.mutation.ClearUpstreamModelMismatch() + return _u +} + // SetChannelID sets the "channel_id" field. func (_u *UsageLogUpdateOne) SetChannelID(v int64) *UsageLogUpdateOne { _u.mutation.ResetChannelID() @@ -2486,6 +2583,11 @@ func (_u *UsageLogUpdateOne) check() error { return &ValidationError{Name: "upstream_model", err: fmt.Errorf(`ent: validator failed for field "UsageLog.upstream_model": %w`, err)} } } + if v, ok := _u.mutation.UpstreamResponseModel(); ok { + if err := usagelog.UpstreamResponseModelValidator(v); err != nil { + return &ValidationError{Name: "upstream_response_model", err: fmt.Errorf(`ent: validator failed for field "UsageLog.upstream_response_model": %w`, err)} + } + } if v, ok := _u.mutation.ModelMappingChain(); ok { if err := usagelog.ModelMappingChainValidator(v); err != nil { return &ValidationError{Name: "model_mapping_chain", err: fmt.Errorf(`ent: validator failed for field "UsageLog.model_mapping_chain": %w`, err)} @@ -2595,6 +2697,18 @@ func (_u *UsageLogUpdateOne) sqlSave(ctx context.Context) (_node *UsageLog, err if _u.mutation.UpstreamModelCleared() { _spec.ClearField(usagelog.FieldUpstreamModel, field.TypeString) } + if value, ok := _u.mutation.UpstreamResponseModel(); ok { + _spec.SetField(usagelog.FieldUpstreamResponseModel, field.TypeString, value) + } + if _u.mutation.UpstreamResponseModelCleared() { + _spec.ClearField(usagelog.FieldUpstreamResponseModel, field.TypeString) + } + if value, ok := _u.mutation.UpstreamModelMismatch(); ok { + _spec.SetField(usagelog.FieldUpstreamModelMismatch, field.TypeBool, value) + } + if _u.mutation.UpstreamModelMismatchCleared() { + _spec.ClearField(usagelog.FieldUpstreamModelMismatch, field.TypeBool) + } if value, ok := _u.mutation.ChannelID(); ok { _spec.SetField(usagelog.FieldChannelID, field.TypeInt64, value) } diff --git a/backend/go.sum b/backend/go.sum index e9a4c1829e..cdea1a5b91 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -250,6 +250,8 @@ github.com/google/go-tpm-tools v0.3.13-0.20230620182252-4639ecce2aba/go.mod h1:E github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= +github.com/google/subcommands v1.2.0 h1:vWQspBTo2nEqTUFita5/KeEWlUL8kQObDFbub/EN9oE= +github.com/google/subcommands v1.2.0/go.mod h1:ZjhPrFU+Olkh9WazFPsl27BQ4UPiG37m3yTrtFlrHVk= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/wire v0.7.0 h1:JxUKI6+CVBgCO2WToKy/nQk0sS+amI9z9EjVmdaocj4= @@ -310,6 +312,8 @@ github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovk github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/mattn/go-runewidth v0.0.15 h1:UNAjwbU9l54TA3KzvqLGxwWjHmMgBUVhBiTjelZgg3U= +github.com/mattn/go-runewidth v0.0.15/go.mod h1:Jdepj2loyihRzMpdS35Xk/zdY8IAYHsh153qUoGf23w= github.com/mattn/go-sqlite3 v1.14.17 h1:mCRHCLDUBXgpKAqIKsaAaAsrAlbkeomtRFKXh2L6YIM= github.com/mattn/go-sqlite3 v1.14.17/go.mod h1:2eHXhiwb8IkHr+BDWZGa96P6+rkvnG63S2DGjv9HUNg= github.com/mdelapenya/tlscert v0.2.0 h1:7H81W6Z/4weDvZBNOfQte5GpIMo0lGYEeWbkGp5LJHI= @@ -346,6 +350,8 @@ github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7P github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= github.com/niemeyer/pretty v0.0.0-20200227124842-a10e7caefd8e/go.mod h1:zD1mROLANZcx1PVRCS0qkT7pwLkGfwJo4zjcN/Tysno= +github.com/olekukonko/tablewriter v0.0.5 h1:P2Ga83D34wi1o9J6Wh1mRuqd4mF/x/lgBS7N7AbDhec= +github.com/olekukonko/tablewriter v0.0.5/go.mod h1:hPp6KlRPjbx+hW8ykQs1w3UBbZlj6HuIJcUGPhkA7kY= github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U= github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM= github.com/opencontainers/image-spec v1.1.1 h1:y0fUlFfIZhPF1W537XOLg0/fcx6zcHCJwooC2xJA040= @@ -380,6 +386,8 @@ github.com/refraction-networking/utls v1.8.2 h1:j4Q1gJj0xngdeH+Ox/qND11aEfhpgoEv github.com/refraction-networking/utls v1.8.2/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +github.com/rivo/uniseg v0.2.0 h1:S1pD9weZBuJdFmowNwbpi7BJ8TNftyUImj/0WQi72jY= +github.com/rivo/uniseg v0.2.0/go.mod h1:J6wj4VEh+S6ZtnVlnTBMWIodfgj8LQOQFoIToxlJtxc= github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= @@ -415,6 +423,8 @@ github.com/spf13/afero v1.11.0 h1:WJQKhtpdm3v2IzqG8VMqrr6Rf3UYpEF239Jy9wNepM8= github.com/spf13/afero v1.11.0/go.mod h1:GH9Y3pIexgf1MTIWtNGyogA5MwRIDXGUr+hbWNoBjkY= github.com/spf13/cast v1.6.0 h1:GEiTHELF+vaR5dhz3VqZfFSzZjYbgeKDpBxQVS4GYJ0= github.com/spf13/cast v1.6.0/go.mod h1:ancEpBxwJDODSW/UG4rDrAqiKolqNNh2DX3mk86cAdo= +github.com/spf13/cobra v1.7.0 h1:hyqWnYt1ZQShIddO5kBpj3vu05/++x6tJ6dg8EC572I= +github.com/spf13/cobra v1.7.0/go.mod h1:uLxZILRyS/50WlhOIKD7W6V5bgeIt+4sICxh6uRMrb0= github.com/spf13/pflag v1.0.5 h1:iy+VFUOCP1a+8yFto/drg2CJ5u0yRoB7fZw3DKv/JXA= github.com/spf13/pflag v1.0.5/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= github.com/spf13/viper v1.18.2 h1:LUXCnvUvSM6FXAsj6nnfc8Q2tp1dIgUfY9Kc8GsSOiQ= diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index ab83333eb4..0ed2b17213 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -1022,6 +1022,39 @@ type GatewayConfig struct { // UserMessageQueue: 用户消息串行队列配置 // 对 role:"user" 的真实用户消息实施账号级串行化 + RPM 自适应延迟 UserMessageQueue UserMessageQueueConfig `mapstructure:"user_message_queue"` + + // Grok: Grok/xAI gateway scheduling and free-tier soft-gate settings. + Grok GatewayGrokConfig `mapstructure:"grok"` +} + +// GatewayGrokConfig holds Grok-specific gateway scheduling knobs. +// +// Free-quota soft gate keys (gateway.grok.*): +// - free_quota_soft_gate_enabled: enable local rolling-window scheduling guard for +// OAuth accounts whose subscription_tier/plan_type is explicitly "free". +// Default true is safe only because free-tier detection is strict (unknown/paid fail open). +// - free_quota_token_limit: nominal rolling-window token allowance. +// - free_quota_soft_gate_percent: stop new scheduling before the nominal limit (1-100). +// - free_quota_window_hours: local usage rolling window length in hours. +// - free_quota_stats_cache_seconds: cache TTL for free-tier usage stats +// (hot path never blocks on DB; misses fail open and refresh in background). +type GatewayGrokConfig struct { + // PasswordAuthEnabled controls the optional password-to-SSO OAuth flow. + // It defaults to false and must be explicitly enabled by the operator. + // When true, POST /admin/grok/oauth/password is functional (not ignored). + PasswordAuthEnabled bool `mapstructure:"password_auth_enabled"` + // FreeQuotaSoftGateEnabled enables a local rolling-window scheduling guard + // for explicitly free Grok OAuth accounts only. + FreeQuotaSoftGateEnabled bool `mapstructure:"free_quota_soft_gate_enabled"` + // FreeQuotaTokenLimit is the nominal rolling-window allowance. + FreeQuotaTokenLimit int64 `mapstructure:"free_quota_token_limit"` + // FreeQuotaSoftGatePercent stops new scheduling before the nominal limit. + FreeQuotaSoftGatePercent int `mapstructure:"free_quota_soft_gate_percent"` + // FreeQuotaWindowHours controls the local rolling usage window. + FreeQuotaWindowHours int `mapstructure:"free_quota_window_hours"` + // FreeQuotaStatsCacheSeconds is the soft-gate stats cache TTL. Hot path never + // waits on usage_logs; misses fail open and refresh asynchronously. + FreeQuotaStatsCacheSeconds int `mapstructure:"free_quota_stats_cache_seconds"` } type GatewayLiveConfig struct { @@ -2309,6 +2342,15 @@ func setDefaults() { viper.SetDefault("gateway.openai_proxy_stream_circuit.failure_threshold", 2) viper.SetDefault("gateway.openai_proxy_stream_circuit.window_seconds", 60) viper.SetDefault("gateway.openai_proxy_stream_circuit.ttl_seconds", 600) + // Grok free-tier local soft gate (scheduler-only; admin QueryQuota does not use this). + // Enabled by default because free detection requires an explicit free tier marker. + viper.SetDefault("gateway.grok.free_quota_soft_gate_enabled", true) + viper.SetDefault("gateway.grok.password_auth_enabled", false) + // Free soft-gate nominal limit: 500k tokens / rolling 24h (operator policy). + viper.SetDefault("gateway.grok.free_quota_token_limit", int64(500_000)) + viper.SetDefault("gateway.grok.free_quota_soft_gate_percent", 95) + viper.SetDefault("gateway.grok.free_quota_window_hours", 24) + viper.SetDefault("gateway.grok.free_quota_stats_cache_seconds", 60) viper.SetDefault("gateway.image_concurrency.enabled", false) viper.SetDefault("gateway.image_concurrency.max_concurrent_requests", 0) viper.SetDefault("gateway.image_concurrency.overflow_mode", ImageConcurrencyOverflowModeReject) @@ -3518,6 +3560,20 @@ func (c *Config) Validate() error { if c.Concurrency.PingInterval < 5 || c.Concurrency.PingInterval > 30 { return fmt.Errorf("concurrency.ping_interval must be between 5-30 seconds") } + if c.Gateway.Grok.FreeQuotaSoftGateEnabled { + if c.Gateway.Grok.FreeQuotaTokenLimit <= 0 { + return fmt.Errorf("gateway.grok.free_quota_token_limit must be positive") + } + if c.Gateway.Grok.FreeQuotaSoftGatePercent < 1 || c.Gateway.Grok.FreeQuotaSoftGatePercent > 100 { + return fmt.Errorf("gateway.grok.free_quota_soft_gate_percent must be between 1 and 100") + } + if c.Gateway.Grok.FreeQuotaWindowHours <= 0 { + return fmt.Errorf("gateway.grok.free_quota_window_hours must be positive") + } + } + if c.Gateway.Grok.FreeQuotaStatsCacheSeconds < 0 { + return fmt.Errorf("gateway.grok.free_quota_stats_cache_seconds must be non-negative") + } if err := ValidateDingTalkConfig(c.DingTalk); err != nil { return fmt.Errorf("dingtalk_connect: %w", err) } diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index b6ade9c4b2..718913a044 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -538,6 +538,19 @@ func TestLoadOpenAICompactModelFromEnv(t *testing.T) { require.Equal(t, "gpt-5.3-codex", cfg.Gateway.OpenAICompactModel) } +func TestLoadDefaultGrokFreeQuotaSoftGate(t *testing.T) { + resetViperWithJWTSecret(t) + + cfg, err := Load() + require.NoError(t, err) + require.False(t, cfg.Gateway.Grok.PasswordAuthEnabled) + require.True(t, cfg.Gateway.Grok.FreeQuotaSoftGateEnabled) + require.Equal(t, int64(500_000), cfg.Gateway.Grok.FreeQuotaTokenLimit) + require.Equal(t, 95, cfg.Gateway.Grok.FreeQuotaSoftGatePercent) + require.Equal(t, 24, cfg.Gateway.Grok.FreeQuotaWindowHours) + require.Equal(t, 60, cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds) +} + func TestLoadDefaultOpenAIHTTP2Enabled(t *testing.T) { resetViperWithJWTSecret(t) diff --git a/backend/internal/domain/constants.go b/backend/internal/domain/constants.go index d79b49a4a7..148c1cd5d3 100644 --- a/backend/internal/domain/constants.go +++ b/backend/internal/domain/constants.go @@ -117,6 +117,12 @@ var DefaultAntigravityModelMapping = map[string]string{ "gemini-3.1-flash-image": "gemini-3.1-flash-image", // Gemini 3.1 image preview 映射 "gemini-3.1-flash-image-preview": "gemini-3.1-flash-image", + // Gemini 3.6 Flash tiered models + "gemini-3.6-flash": "gemini-3.6-flash", + "gemini-3.6-flash-high": "gemini-3.6-flash-high", + "gemini-3.6-flash-low": "gemini-3.6-flash-low", + "gemini-3.6-flash-medium": "gemini-3.6-flash-medium", + "gemini-3.6-flash-tiered": "gemini-3.6-flash-tiered", // Gemini 3 image 兼容映射(向 3.1 image 迁移) "gemini-3-pro-image": "gemini-3.1-flash-image", "gemini-3-pro-image-preview": "gemini-3.1-flash-image", diff --git a/backend/internal/domain/constants_test.go b/backend/internal/domain/constants_test.go index 0fb9054f7e..e847e22f30 100644 --- a/backend/internal/domain/constants_test.go +++ b/backend/internal/domain/constants_test.go @@ -65,6 +65,14 @@ func TestDefaultAntigravityModelMapping_Gemini31ProAliases(t *testing.T) { } } +func TestDefaultAntigravityModelMapping_Gemini36FlashModels(t *testing.T) { + for _, model := range []string{"gemini-3.6-flash", "gemini-3.6-flash-high", "gemini-3.6-flash-low", "gemini-3.6-flash-medium", "gemini-3.6-flash-tiered"} { + if got := DefaultAntigravityModelMapping[model]; got != model { + t.Fatalf("expected %s to map to itself, got %q", model, got) + } + } +} + func TestDefaultBedrockModelMapping_ContainsNewClaudeModels(t *testing.T) { t.Parallel() diff --git a/backend/internal/handler/admin/account_handler.go b/backend/internal/handler/admin/account_handler.go index 9b80dd05d1..30fc6b9828 100644 --- a/backend/internal/handler/admin/account_handler.go +++ b/backend/internal/handler/admin/account_handler.go @@ -1065,6 +1065,10 @@ type TestAccountRequest struct { ModelID string `json:"model_id"` Prompt string `json:"prompt"` Mode string `json:"mode"` + // Optional media for Grok (and future) real generation tests. + // ImageDataURL / AudioDataURL are data:;base64,... payloads. + ImageDataURL string `json:"image_data_url"` + AudioDataURL string `json:"audio_data_url"` } type SyncFromCRSRequest struct { @@ -1094,8 +1098,13 @@ func (h *AccountHandler) Test(c *gin.Context) { // Allow empty body, model_id is optional _ = c.ShouldBindJSON(&req) + opts := service.AccountTestOptions{ + ImageDataURL: req.ImageDataURL, + AudioDataURL: req.AudioDataURL, + } + // Use AccountTestService to test the account with SSE streaming - if err := h.accountTestService.TestAccountConnection(c, accountID, req.ModelID, req.Prompt, req.Mode); err != nil { + if err := h.accountTestService.TestAccountConnection(c, accountID, req.ModelID, req.Prompt, req.Mode, opts); err != nil { // Error already sent via SSE, just log return } @@ -1415,6 +1424,9 @@ func (h *AccountHandler) ApplyOAuthCredentials(c *gin.Context) { return } + // Drop SSO/password residue; re-auth must leave only OAuth tokens on disk. + req.Credentials = service.SanitizeStoredCredentials(existing.Platform, req.Credentials) + updatedAccount, err := h.adminService.UpdateAccount(ctx, accountID, &service.UpdateAccountInput{ Type: req.Type, Credentials: req.Credentials, @@ -1442,6 +1454,20 @@ func (h *AccountHandler) ApplyOAuthCredentials(c *gin.Context) { } } + // Successful re-auth clears the soft spending-limit reauth flag for Grok. + if existing.Platform == service.PlatformGrok { + if clearErr := h.adminService.UpdateAccountExtra(ctx, accountID, map[string]any{ + "grok_needs_reauth": false, + "grok_needs_reauth_reason": "", + "grok_needs_reauth_at": "", + }); clearErr != nil { + slog.Warn("apply_oauth_credentials.clear_grok_reauth_failed", + "account_id", accountID, + "err", clearErr, + ) + } + } + if cleared, clearErr := h.adminService.ClearAccountError(ctx, accountID); clearErr != nil { slog.Warn("apply_oauth_credentials.clear_error_failed", "account_id", accountID, @@ -2422,6 +2448,11 @@ type BatchTodayStatsRequest struct { AccountIDs []int64 `json:"account_ids" binding:"required"` } +type BatchUsageRequest struct { + AccountIDs []int64 `json:"account_ids" binding:"required"` + Force bool `json:"force"` +} + // GetBatchTodayStats 批量获取多个账号的今日统计。 // POST /api/v1/admin/accounts/today-stats/batch func (h *AccountHandler) GetBatchTodayStats(c *gin.Context) { @@ -2468,6 +2499,36 @@ func (h *AccountHandler) GetBatchTodayStats(c *gin.Context) { response.Success(c, payload) } +// GetBatchUsage 批量获取多个账号的 current usage。 +// POST /api/v1/admin/accounts/usage/batch +func (h *AccountHandler) GetBatchUsage(c *gin.Context) { + var req BatchUsageRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + + accountIDs := normalizeInt64IDList(req.AccountIDs) + if len(accountIDs) == 0 { + response.Success(c, gin.H{ + "usage": map[string]any{}, + "errors": map[string]string{}, + }) + return + } + + usageByAccount, errorsByAccount, err := h.accountUsageService.GetUsageBatch(c.Request.Context(), accountIDs, req.Force) + if err != nil { + response.ErrorFrom(c, err) + return + } + + response.Success(c, gin.H{ + "usage": usageByAccount, + "errors": errorsByAccount, + }) +} + // SetSchedulableRequest represents the request body for setting schedulable status type SetSchedulableRequest struct { Schedulable bool `json:"schedulable"` diff --git a/backend/internal/handler/admin/dashboard_handler.go b/backend/internal/handler/admin/dashboard_handler.go index 8f55fb4165..c94ed033fa 100644 --- a/backend/internal/handler/admin/dashboard_handler.go +++ b/backend/internal/handler/admin/dashboard_handler.go @@ -64,6 +64,18 @@ func parseTimeRange(c *gin.Context) (time.Time, time.Time) { return startTime, endTime } +func parseOptionalBoolDashboardFilter(c *gin.Context, name string) (*bool, error) { + raw := strings.TrimSpace(c.Query(name)) + if raw == "" { + return nil, nil + } + value, err := strconv.ParseBool(raw) + if err != nil { + return nil, err + } + return &value, nil +} + // GetStats handles getting dashboard statistics // GET /api/v1/admin/dashboard/stats func (h *DashboardHandler) GetStats(c *gin.Context) { @@ -200,6 +212,7 @@ func (h *DashboardHandler) GetUsageTrend(c *gin.Context) { var requestType *int16 var stream *bool var billingType *int8 + var upstreamModelMismatch *bool if userIDStr := c.Query("user_id"); userIDStr != "" { if id, err := strconv.ParseInt(userIDStr, 10, 64); err == nil { @@ -249,8 +262,13 @@ func (h *DashboardHandler) GetUsageTrend(c *gin.Context) { return } } + upstreamModelMismatch, err := parseOptionalBoolDashboardFilter(c, "upstream_model_mismatch") + if err != nil { + response.BadRequest(c, "Invalid upstream_model_mismatch value, use true or false") + return + } - trend, hit, err := h.getUsageTrendCached(c.Request.Context(), startTime, endTime, granularity, userID, apiKeyID, accountID, groupID, model, requestType, stream, billingType) + trend, hit, err := h.getUsageTrendCached(c.Request.Context(), startTime, endTime, granularity, userID, apiKeyID, accountID, groupID, model, requestType, stream, billingType, upstreamModelMismatch) if err != nil { response.Error(c, 500, "Failed to get usage trend") return @@ -277,6 +295,7 @@ func (h *DashboardHandler) GetModelStats(c *gin.Context) { var requestType *int16 var stream *bool var billingType *int8 + var upstreamModelMismatch *bool if userIDStr := c.Query("user_id"); userIDStr != "" { if id, err := strconv.ParseInt(userIDStr, 10, 64); err == nil { @@ -330,8 +349,13 @@ func (h *DashboardHandler) GetModelStats(c *gin.Context) { return } } + upstreamModelMismatch, err := parseOptionalBoolDashboardFilter(c, "upstream_model_mismatch") + if err != nil { + response.BadRequest(c, "Invalid upstream_model_mismatch value, use true or false") + return + } - stats, hit, err := h.getModelStatsCached(c.Request.Context(), startTime, endTime, userID, apiKeyID, accountID, groupID, modelSource, requestType, stream, billingType) + stats, hit, err := h.getModelStatsCached(c.Request.Context(), startTime, endTime, userID, apiKeyID, accountID, groupID, modelSource, requestType, stream, billingType, upstreamModelMismatch) if err != nil { response.Error(c, 500, "Failed to get model statistics") return @@ -355,6 +379,7 @@ func (h *DashboardHandler) GetGroupStats(c *gin.Context) { var requestType *int16 var stream *bool var billingType *int8 + var upstreamModelMismatch *bool if userIDStr := c.Query("user_id"); userIDStr != "" { if id, err := strconv.ParseInt(userIDStr, 10, 64); err == nil { @@ -401,8 +426,13 @@ func (h *DashboardHandler) GetGroupStats(c *gin.Context) { return } } + upstreamModelMismatch, err := parseOptionalBoolDashboardFilter(c, "upstream_model_mismatch") + if err != nil { + response.BadRequest(c, "Invalid upstream_model_mismatch value, use true or false") + return + } - stats, hit, err := h.getGroupStatsCached(c.Request.Context(), startTime, endTime, userID, apiKeyID, accountID, groupID, requestType, stream, billingType) + stats, hit, err := h.getGroupStatsCached(c.Request.Context(), startTime, endTime, userID, apiKeyID, accountID, groupID, requestType, stream, billingType, upstreamModelMismatch) if err != nil { response.Error(c, 500, "Failed to get group statistics") return diff --git a/backend/internal/handler/admin/dashboard_handler_request_type_test.go b/backend/internal/handler/admin/dashboard_handler_request_type_test.go index 6056f725b5..3316557dba 100644 --- a/backend/internal/handler/admin/dashboard_handler_request_type_test.go +++ b/backend/internal/handler/admin/dashboard_handler_request_type_test.go @@ -19,11 +19,26 @@ type dashboardUsageRepoCapture struct { trendStream *bool modelRequestType *int16 modelStream *bool + trendMismatch *bool + modelMismatch *bool + groupMismatch *bool rankingLimit int ranking []usagestats.UserSpendingRankingItem rankingTotal float64 } +func (s *dashboardUsageRepoCapture) GetUsageTrendWithUsageFilters( + ctx context.Context, + startTime, endTime time.Time, + granularity string, + filters usagestats.UsageLogFilters, +) ([]usagestats.TrendDataPoint, error) { + s.trendRequestType = filters.RequestType + s.trendStream = filters.Stream + s.trendMismatch = filters.UpstreamModelMismatch + return []usagestats.TrendDataPoint{}, nil +} + func (s *dashboardUsageRepoCapture) GetUsageTrendWithFilters( ctx context.Context, startTime, endTime time.Time, @@ -39,6 +54,27 @@ func (s *dashboardUsageRepoCapture) GetUsageTrendWithFilters( return []usagestats.TrendDataPoint{}, nil } +func (s *dashboardUsageRepoCapture) GetModelStatsWithUsageFiltersBySource( + ctx context.Context, + startTime, endTime time.Time, + filters usagestats.UsageLogFilters, + source string, +) ([]usagestats.ModelStat, error) { + s.modelRequestType = filters.RequestType + s.modelStream = filters.Stream + s.modelMismatch = filters.UpstreamModelMismatch + return []usagestats.ModelStat{}, nil +} + +func (s *dashboardUsageRepoCapture) GetGroupStatsWithUsageFilters( + ctx context.Context, + startTime, endTime time.Time, + filters usagestats.UsageLogFilters, +) ([]usagestats.GroupStat, error) { + s.groupMismatch = filters.UpstreamModelMismatch + return []usagestats.GroupStat{}, nil +} + func (s *dashboardUsageRepoCapture) GetModelStatsWithFilters( ctx context.Context, startTime, endTime time.Time, @@ -73,6 +109,7 @@ func newDashboardRequestTypeTestRouter(repo *dashboardUsageRepoCapture) *gin.Eng router := gin.New() router.GET("/admin/dashboard/trend", handler.GetUsageTrend) router.GET("/admin/dashboard/models", handler.GetModelStats) + router.GET("/admin/dashboard/groups", handler.GetGroupStats) router.GET("/admin/dashboard/users-ranking", handler.GetUserSpendingRanking) return router } @@ -171,6 +208,46 @@ func TestDashboardModelStatsValidModelSource(t *testing.T) { require.Equal(t, http.StatusOK, rec.Code) } +func TestDashboardModelAuditFilterPropagatesToTrendModelAndGroupQueries(t *testing.T) { + resetDashboardReadCachesForTest() + repo := &dashboardUsageRepoCapture{} + router := newDashboardRequestTypeTestRouter(repo) + + for _, path := range []string{ + "/admin/dashboard/trend?upstream_model_mismatch=true", + "/admin/dashboard/models?upstream_model_mismatch=true", + "/admin/dashboard/groups?upstream_model_mismatch=true", + } { + req := httptest.NewRequest(http.MethodGet, path, nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusOK, rec.Code, path) + } + + require.NotNil(t, repo.trendMismatch) + require.True(t, *repo.trendMismatch) + require.NotNil(t, repo.modelMismatch) + require.True(t, *repo.modelMismatch) + require.NotNil(t, repo.groupMismatch) + require.True(t, *repo.groupMismatch) +} + +func TestDashboardModelAuditFilterRejectsInvalidBoolean(t *testing.T) { + repo := &dashboardUsageRepoCapture{} + router := newDashboardRequestTypeTestRouter(repo) + + for _, path := range []string{ + "/admin/dashboard/trend?upstream_model_mismatch=invalid", + "/admin/dashboard/models?upstream_model_mismatch=invalid", + "/admin/dashboard/groups?upstream_model_mismatch=invalid", + } { + req := httptest.NewRequest(http.MethodGet, path, nil) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, req) + require.Equal(t, http.StatusBadRequest, rec.Code, path) + } +} + func TestDashboardUsersRankingLimitAndCache(t *testing.T) { dashboardUsersRankingCache = newSnapshotCache(5 * time.Minute) repo := &dashboardUsageRepoCapture{ diff --git a/backend/internal/handler/admin/dashboard_query_cache.go b/backend/internal/handler/admin/dashboard_query_cache.go index 815c516154..4f79fb4500 100644 --- a/backend/internal/handler/admin/dashboard_query_cache.go +++ b/backend/internal/handler/admin/dashboard_query_cache.go @@ -18,30 +18,32 @@ var ( ) type dashboardTrendCacheKey struct { - StartTime string `json:"start_time"` - EndTime string `json:"end_time"` - Granularity string `json:"granularity"` - 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"` - RequestType *int16 `json:"request_type"` - Stream *bool `json:"stream"` - BillingType *int8 `json:"billing_type"` + StartTime string `json:"start_time"` + EndTime string `json:"end_time"` + Granularity string `json:"granularity"` + 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"` + RequestType *int16 `json:"request_type"` + Stream *bool `json:"stream"` + BillingType *int8 `json:"billing_type"` + UpstreamModelMismatch *bool `json:"upstream_model_mismatch"` } type dashboardModelGroupCacheKey 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"` - ModelSource string `json:"model_source,omitempty"` - RequestType *int16 `json:"request_type"` - Stream *bool `json:"stream"` - BillingType *int8 `json:"billing_type"` + 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"` + ModelSource string `json:"model_source,omitempty"` + RequestType *int16 `json:"request_type"` + Stream *bool `json:"stream"` + BillingType *int8 `json:"billing_type"` + UpstreamModelMismatch *bool `json:"upstream_model_mismatch"` } type dashboardEntityTrendCacheKey struct { @@ -84,22 +86,28 @@ func (h *DashboardHandler) getUsageTrendCached( requestType *int16, stream *bool, billingType *int8, + upstreamModelMismatch *bool, ) ([]usagestats.TrendDataPoint, bool, error) { key := mustMarshalDashboardCacheKey(dashboardTrendCacheKey{ - StartTime: startTime.UTC().Format(time.RFC3339), - EndTime: endTime.UTC().Format(time.RFC3339), - Granularity: granularity, - UserID: userID, - APIKeyID: apiKeyID, - AccountID: accountID, - GroupID: groupID, - Model: model, - RequestType: requestType, - Stream: stream, - BillingType: billingType, + StartTime: startTime.UTC().Format(time.RFC3339), + EndTime: endTime.UTC().Format(time.RFC3339), + Granularity: granularity, + UserID: userID, + APIKeyID: apiKeyID, + AccountID: accountID, + GroupID: groupID, + Model: model, + RequestType: requestType, + Stream: stream, + BillingType: billingType, + UpstreamModelMismatch: upstreamModelMismatch, }) entry, hit, err := dashboardTrendCache.GetOrLoad(key, func() (any, error) { - return h.dashboardService.GetUsageTrendWithFilters(ctx, startTime, endTime, granularity, userID, apiKeyID, accountID, groupID, model, requestType, stream, billingType) + return h.dashboardService.GetUsageTrendWithUsageFilters(ctx, startTime, endTime, granularity, usagestats.UsageLogFilters{ + UserID: userID, APIKeyID: apiKeyID, AccountID: accountID, GroupID: groupID, + Model: model, RequestType: requestType, Stream: stream, BillingType: billingType, + UpstreamModelMismatch: upstreamModelMismatch, + }) }) if err != nil { return nil, hit, err @@ -116,21 +124,27 @@ func (h *DashboardHandler) getModelStatsCached( requestType *int16, stream *bool, billingType *int8, + upstreamModelMismatch *bool, ) ([]usagestats.ModelStat, bool, error) { key := mustMarshalDashboardCacheKey(dashboardModelGroupCacheKey{ - StartTime: startTime.UTC().Format(time.RFC3339), - EndTime: endTime.UTC().Format(time.RFC3339), - UserID: userID, - APIKeyID: apiKeyID, - AccountID: accountID, - GroupID: groupID, - ModelSource: usagestats.NormalizeModelSource(modelSource), - RequestType: requestType, - Stream: stream, - BillingType: billingType, + StartTime: startTime.UTC().Format(time.RFC3339), + EndTime: endTime.UTC().Format(time.RFC3339), + UserID: userID, + APIKeyID: apiKeyID, + AccountID: accountID, + GroupID: groupID, + ModelSource: usagestats.NormalizeModelSource(modelSource), + RequestType: requestType, + Stream: stream, + BillingType: billingType, + UpstreamModelMismatch: upstreamModelMismatch, }) entry, hit, err := dashboardModelStatsCache.GetOrLoad(key, func() (any, error) { - return h.dashboardService.GetModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, requestType, stream, billingType, modelSource) + return h.dashboardService.GetModelStatsWithUsageFiltersBySource(ctx, startTime, endTime, usagestats.UsageLogFilters{ + UserID: userID, APIKeyID: apiKeyID, AccountID: accountID, GroupID: groupID, + RequestType: requestType, Stream: stream, BillingType: billingType, + UpstreamModelMismatch: upstreamModelMismatch, + }, modelSource) }) if err != nil { return nil, hit, err @@ -146,20 +160,26 @@ func (h *DashboardHandler) getGroupStatsCached( requestType *int16, stream *bool, billingType *int8, + upstreamModelMismatch *bool, ) ([]usagestats.GroupStat, bool, error) { key := mustMarshalDashboardCacheKey(dashboardModelGroupCacheKey{ - StartTime: startTime.UTC().Format(time.RFC3339), - EndTime: endTime.UTC().Format(time.RFC3339), - UserID: userID, - APIKeyID: apiKeyID, - AccountID: accountID, - GroupID: groupID, - RequestType: requestType, - Stream: stream, - BillingType: billingType, + StartTime: startTime.UTC().Format(time.RFC3339), + EndTime: endTime.UTC().Format(time.RFC3339), + UserID: userID, + APIKeyID: apiKeyID, + AccountID: accountID, + GroupID: groupID, + RequestType: requestType, + Stream: stream, + BillingType: billingType, + UpstreamModelMismatch: upstreamModelMismatch, }) entry, hit, err := dashboardGroupStatsCache.GetOrLoad(key, func() (any, error) { - return h.dashboardService.GetGroupStatsWithFilters(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, requestType, stream, billingType) + return h.dashboardService.GetGroupStatsWithUsageFilters(ctx, startTime, endTime, usagestats.UsageLogFilters{ + UserID: userID, APIKeyID: apiKeyID, AccountID: accountID, GroupID: groupID, + RequestType: requestType, Stream: stream, BillingType: billingType, + UpstreamModelMismatch: upstreamModelMismatch, + }) }) if err != nil { return nil, hit, err diff --git a/backend/internal/handler/admin/dashboard_snapshot_v2_handler.go b/backend/internal/handler/admin/dashboard_snapshot_v2_handler.go index 517ae7bd14..b19c9360aa 100644 --- a/backend/internal/handler/admin/dashboard_snapshot_v2_handler.go +++ b/backend/internal/handler/admin/dashboard_snapshot_v2_handler.go @@ -37,34 +37,36 @@ type dashboardSnapshotV2Response struct { } type dashboardSnapshotV2Filters struct { - UserID int64 - APIKeyID int64 - AccountID int64 - GroupID int64 - Model string - RequestType *int16 - Stream *bool - BillingType *int8 + UserID int64 + APIKeyID int64 + AccountID int64 + GroupID int64 + Model string + RequestType *int16 + Stream *bool + BillingType *int8 + UpstreamModelMismatch *bool } type dashboardSnapshotV2CacheKey struct { - StartTime string `json:"start_time"` - EndTime string `json:"end_time"` - Granularity string `json:"granularity"` - 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"` - RequestType *int16 `json:"request_type"` - Stream *bool `json:"stream"` - BillingType *int8 `json:"billing_type"` - IncludeStats bool `json:"include_stats"` - IncludeTrend bool `json:"include_trend"` - IncludeModels bool `json:"include_models"` - IncludeGroups bool `json:"include_groups"` - IncludeUsersTrend bool `json:"include_users_trend"` - UsersTrendLimit int `json:"users_trend_limit"` + StartTime string `json:"start_time"` + EndTime string `json:"end_time"` + Granularity string `json:"granularity"` + 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"` + RequestType *int16 `json:"request_type"` + Stream *bool `json:"stream"` + BillingType *int8 `json:"billing_type"` + UpstreamModelMismatch *bool `json:"upstream_model_mismatch"` + IncludeStats bool `json:"include_stats"` + IncludeTrend bool `json:"include_trend"` + IncludeModels bool `json:"include_models"` + IncludeGroups bool `json:"include_groups"` + IncludeUsersTrend bool `json:"include_users_trend"` + UsersTrendLimit int `json:"users_trend_limit"` } func (h *DashboardHandler) GetSnapshotV2(c *gin.Context) { @@ -93,23 +95,24 @@ func (h *DashboardHandler) GetSnapshotV2(c *gin.Context) { } keyRaw, _ := json.Marshal(dashboardSnapshotV2CacheKey{ - StartTime: startTime.UTC().Format(time.RFC3339), - EndTime: endTime.UTC().Format(time.RFC3339), - Granularity: granularity, - UserID: filters.UserID, - APIKeyID: filters.APIKeyID, - AccountID: filters.AccountID, - GroupID: filters.GroupID, - Model: filters.Model, - RequestType: filters.RequestType, - Stream: filters.Stream, - BillingType: filters.BillingType, - IncludeStats: includeStats, - IncludeTrend: includeTrend, - IncludeModels: includeModels, - IncludeGroups: includeGroups, - IncludeUsersTrend: includeUsersTrend, - UsersTrendLimit: usersTrendLimit, + StartTime: startTime.UTC().Format(time.RFC3339), + EndTime: endTime.UTC().Format(time.RFC3339), + Granularity: granularity, + UserID: filters.UserID, + APIKeyID: filters.APIKeyID, + AccountID: filters.AccountID, + GroupID: filters.GroupID, + Model: filters.Model, + RequestType: filters.RequestType, + Stream: filters.Stream, + BillingType: filters.BillingType, + UpstreamModelMismatch: filters.UpstreamModelMismatch, + IncludeStats: includeStats, + IncludeTrend: includeTrend, + IncludeModels: includeModels, + IncludeGroups: includeGroups, + IncludeUsersTrend: includeUsersTrend, + UsersTrendLimit: usersTrendLimit, }) cacheKey := string(keyRaw) @@ -184,6 +187,7 @@ func (h *DashboardHandler) buildSnapshotV2Response( filters.RequestType, filters.Stream, filters.BillingType, + filters.UpstreamModelMismatch, ) if err != nil { return nil, errors.New("failed to get usage trend") @@ -204,6 +208,7 @@ func (h *DashboardHandler) buildSnapshotV2Response( filters.RequestType, filters.Stream, filters.BillingType, + filters.UpstreamModelMismatch, ) if err != nil { return nil, errors.New("failed to get model statistics") @@ -223,6 +228,7 @@ func (h *DashboardHandler) buildSnapshotV2Response( filters.RequestType, filters.Stream, filters.BillingType, + filters.UpstreamModelMismatch, ) if err != nil { return nil, errors.New("failed to get group statistics") @@ -299,5 +305,13 @@ func parseDashboardSnapshotV2Filters(c *gin.Context) (*dashboardSnapshotV2Filter filters.BillingType = &bt } + if mismatchStr := strings.TrimSpace(c.Query("upstream_model_mismatch")); mismatchStr != "" { + value, err := strconv.ParseBool(mismatchStr) + if err != nil { + return nil, err + } + filters.UpstreamModelMismatch = &value + } + return filters, nil } diff --git a/backend/internal/handler/admin/grok_import_probe.go b/backend/internal/handler/admin/grok_import_probe.go index 9ef4ce8cd0..da547b82f8 100644 --- a/backend/internal/handler/admin/grok_import_probe.go +++ b/backend/internal/handler/admin/grok_import_probe.go @@ -13,6 +13,7 @@ import ( const ( grokImportProbeConcurrency = 3 grokImportProbeTimeout = 25 * time.Second + grokImportProbeQueueLimit = 64 ) type grokImportProber interface { @@ -27,6 +28,8 @@ type grokImportProbeTask struct { type grokImportProbeScheduler struct { mu sync.Mutex queue []grokImportProbeTask + pending map[int64]struct{} + inFlight map[int64]struct{} concurrency int workers int maxWorkers int @@ -48,6 +51,8 @@ func newGrokImportProbeScheduler(concurrency int, timeout time.Duration) *grokIm return &grokImportProbeScheduler{ concurrency: concurrency, timeout: timeout, + pending: make(map[int64]struct{}), + inFlight: make(map[int64]struct{}), } } @@ -60,7 +65,21 @@ func (s *grokImportProbeScheduler) schedule(prober grokImportProber, account *se } s.mu.Lock() + if _, exists := s.pending[account.ID]; exists { + s.mu.Unlock() + return + } + if _, exists := s.inFlight[account.ID]; exists { + s.mu.Unlock() + return + } + if len(s.queue) >= grokImportProbeQueueLimit { + s.mu.Unlock() + slog.Debug("grok_import_active_probe_dropped", "account_id", account.ID, "reason", "queue_full") + return + } s.queue = append(s.queue, grokImportProbeTask{prober: prober, accountID: account.ID}) + s.pending[account.ID] = struct{}{} if s.workers < s.concurrency { s.workers++ if s.workers > s.maxWorkers { @@ -78,6 +97,7 @@ func (s *grokImportProbeScheduler) worker() { return } s.run(task.prober, task.accountID) + s.finish(task.accountID) } } @@ -94,9 +114,17 @@ func (s *grokImportProbeScheduler) nextTask() (grokImportProbeTask, bool) { if len(s.queue) == 0 { s.queue = nil } + delete(s.pending, task.accountID) + s.inFlight[task.accountID] = struct{}{} return task, true } +func (s *grokImportProbeScheduler) finish(accountID int64) { + s.mu.Lock() + delete(s.inFlight, accountID) + s.mu.Unlock() +} + func (s *grokImportProbeScheduler) run(prober grokImportProber, accountID int64) { defer func() { if recovered := recover(); recovered != nil { @@ -108,8 +136,6 @@ func (s *grokImportProbeScheduler) run(prober grokImportProber, accountID int64) } }() - // Queue time is intentionally excluded: every imported account is probed, - // while this timeout only bounds the actual upstream probe execution. ctx, cancel := context.WithTimeout(context.Background(), s.timeout) defer cancel() result, err := prober.QueryQuota(ctx, accountID) diff --git a/backend/internal/handler/admin/grok_import_probe_handler_test.go b/backend/internal/handler/admin/grok_import_probe_handler_test.go index c671d2ff1c..a8b7d91250 100644 --- a/backend/internal/handler/admin/grok_import_probe_handler_test.go +++ b/backend/internal/handler/admin/grok_import_probe_handler_test.go @@ -4,6 +4,7 @@ package admin import ( "context" + "errors" "net/http" "net/http/httptest" "strings" @@ -62,6 +63,10 @@ func (grokImportOAuthClientStub) RefreshToken(context.Context, string, string, s return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil } +func (grokImportOAuthClientStub) LoginWithPassword(context.Context, string, string, string) (*service.GrokPasswordLoginResult, error) { + return nil, errors.New("unexpected password login") +} + func (grokImportOAuthClientStub) ConvertSSOToBuild(context.Context, string, string) (*xai.TokenResponse, error) { return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil } diff --git a/backend/internal/handler/admin/grok_import_probe_test.go b/backend/internal/handler/admin/grok_import_probe_test.go index 255f04d0dc..f4d8567463 100644 --- a/backend/internal/handler/admin/grok_import_probe_test.go +++ b/backend/internal/handler/admin/grok_import_probe_test.go @@ -142,7 +142,7 @@ func TestGrokImportProbeSchedulerProbesSingleAccountOnce(t *testing.T) { } func TestGrokImportProbeSchedulerQueuesBatchWithoutPerTaskGoroutines(t *testing.T) { - const taskCount = 100 + const taskCount = 50 release := make(chan struct{}) scheduler := newGrokImportProbeScheduler(3, time.Second) prober := newGrokImportProbeStub(taskCount) @@ -156,7 +156,7 @@ func TestGrokImportProbeSchedulerQueuesBatchWithoutPerTaskGoroutines(t *testing. awaitGrokProbeSignal(t, prober.started) } snapshot := snapshotGrokImportProbeScheduler(scheduler) - require.Equal(t, 97, snapshot.queued) + require.Equal(t, taskCount-3, snapshot.queued) require.Equal(t, 3, snapshot.workers) require.Equal(t, 3, snapshot.maxWorkers) select { @@ -182,6 +182,55 @@ func TestGrokImportProbeSchedulerQueuesBatchWithoutPerTaskGoroutines(t *testing. require.Equal(t, 3, snapshot.maxWorkers) } +func TestGrokImportProbeSchedulerDeduplicatesPendingAndInFlightAccounts(t *testing.T) { + scheduler := newGrokImportProbeScheduler(1, time.Second) + prober := newGrokImportProbeStub(2) + release := make(chan struct{}) + prober.block = release + account := newGrokOAuthImportAccount(501) + queued := newGrokOAuthImportAccount(502) + + scheduler.schedule(prober, account) + require.Equal(t, int64(501), awaitGrokProbeSignal(t, prober.started)) + scheduler.schedule(prober, account) + scheduler.schedule(prober, queued) + scheduler.schedule(prober, queued) + + scheduler.mu.Lock() + require.Len(t, scheduler.queue, 1) + require.Contains(t, scheduler.inFlight, int64(501)) + require.Contains(t, scheduler.pending, int64(502)) + scheduler.mu.Unlock() + + close(release) + require.Equal(t, int64(501), awaitGrokProbeSignal(t, prober.done)) + require.Equal(t, int64(502), awaitGrokProbeSignal(t, prober.done)) + calls, _, _ := prober.snapshot() + require.Equal(t, 1, calls[501]) + require.Equal(t, 1, calls[502]) +} + +func TestGrokImportProbeSchedulerBoundsPendingQueue(t *testing.T) { + scheduler := newGrokImportProbeScheduler(1, time.Second) + prober := newGrokImportProbeStub(grokImportProbeQueueLimit + 1) + release := make(chan struct{}) + prober.block = release + scheduler.schedule(prober, newGrokOAuthImportAccount(600)) + require.Equal(t, int64(600), awaitGrokProbeSignal(t, prober.started)) + for id := int64(601); id < 601+grokImportProbeQueueLimit+10; id++ { + scheduler.schedule(prober, newGrokOAuthImportAccount(id)) + } + + scheduler.mu.Lock() + require.Len(t, scheduler.queue, grokImportProbeQueueLimit) + scheduler.mu.Unlock() + + close(release) + for i := 0; i < grokImportProbeQueueLimit+1; i++ { + awaitGrokProbeSignal(t, prober.done) + } +} + func TestGrokImportProbeSchedulerTimeoutCancelsProbe(t *testing.T) { neverRelease := make(chan struct{}) scheduler := newGrokImportProbeScheduler(1, 20*time.Millisecond) diff --git a/backend/internal/handler/admin/grok_oauth_handler.go b/backend/internal/handler/admin/grok_oauth_handler.go index 1f679b9566..696a24b419 100644 --- a/backend/internal/handler/admin/grok_oauth_handler.go +++ b/backend/internal/handler/admin/grok_oauth_handler.go @@ -47,6 +47,10 @@ type GrokGenerateAuthURLRequest struct { RedirectURI string `json:"redirect_uri"` } +func (h *GrokOAuthHandler) GetCapabilities(c *gin.Context) { + response.Success(c, h.grokOAuthService.GetCapabilities()) +} + func (h *GrokOAuthHandler) GenerateAuthURL(c *gin.Context) { var req GrokGenerateAuthURLRequest if err := c.ShouldBindJSON(&req); err != nil { @@ -95,6 +99,17 @@ type GrokRefreshTokenRequest struct { ProxyID *int64 `json:"proxy_id"` } +type GrokSSOTokenRequest struct { + SSOToken string `json:"sso_token"` + ProxyID *int64 `json:"proxy_id"` +} + +type GrokPasswordAuthorizeRequest struct { + Email string `json:"email"` + Password string `json:"password"` + ProxyID *int64 `json:"proxy_id"` +} + func (h *GrokOAuthHandler) RefreshToken(c *gin.Context) { var req GrokRefreshTokenRequest if err := c.ShouldBindJSON(&req); err != nil { @@ -113,9 +128,15 @@ func (h *GrokOAuthHandler) RefreshToken(c *gin.Context) { var proxyURL string if req.ProxyID != nil { proxy, err := h.adminService.GetProxy(c.Request.Context(), *req.ProxyID) - if err == nil && proxy != nil { - proxyURL = proxy.URL() + if err != nil { + response.ErrorFrom(c, err) + return } + if proxy == nil { + response.BadRequest(c, "GROK_OAUTH_PROXY_NOT_FOUND: proxy not found") + return + } + proxyURL = proxy.URL() } tokenInfo, err := h.grokOAuthService.RefreshToken(c.Request.Context(), refreshToken, proxyURL, req.ClientID) if err != nil { @@ -125,6 +146,38 @@ func (h *GrokOAuthHandler) RefreshToken(c *gin.Context) { response.Success(c, tokenInfo) } +// ValidateSSOToken converts a Web SSO cookie into Build OAuth tokens. +// Response contains OAuth token info only — never echoes sso_token. +func (h *GrokOAuthHandler) ValidateSSOToken(c *gin.Context) { + var req GrokSSOTokenRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + tokenInfo, err := h.grokOAuthService.ValidateSSOToken(c.Request.Context(), req.SSOToken, req.ProxyID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, tokenInfo) +} + +// AuthorizePassword exchanges email/password for Build OAuth tokens via SSO conversion. +// Response never includes password or raw sso_token. +func (h *GrokOAuthHandler) AuthorizePassword(c *gin.Context) { + var req GrokPasswordAuthorizeRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.BadRequest(c, "Invalid request: "+err.Error()) + return + } + tokenInfo, err := h.grokOAuthService.AuthorizePassword(c.Request.Context(), req.Email, req.Password, req.ProxyID) + if err != nil { + response.ErrorFrom(c, err) + return + } + response.Success(c, tokenInfo) +} + func (h *GrokOAuthHandler) RefreshAccountToken(c *gin.Context) { accountID, err := strconv.ParseInt(c.Param("id"), 10, 64) if err != nil { @@ -412,11 +465,38 @@ func (h *GrokOAuthHandler) createAccountFromSSOToken(ctx context.Context, req Gr // 配置且 Build 恒写官方地址,会吞掉导入时指定的自定义转发地址——与 // RefreshAccountToken 的保留逻辑对齐,请求显式提供时以请求为准。 func grokSSOImportCredentials(built map[string]any, reqCredentials map[string]any) map[string]any { - credentials := service.MergeCredentials(cloneGrokSSOMap(reqCredentials), built) + // Only merge operator config from the request — never free-form secrets + // (password / sso_token / cookie / etc.) into stored credentials. + allowedReqKeys := map[string]struct{}{ + "base_url": {}, "model_mapping": {}, + "header_override": {}, "header_overrides": {}, "header_override_enabled": {}, + "custom_headers": {}, + } + ops := map[string]any{} + for k, v := range reqCredentials { + if _, ok := allowedReqKeys[k]; !ok { + continue + } + if service.IsSensitiveCredentialKey(k) { + continue + } + ops[k] = v + } + credentials := service.MergeCredentials(ops, built) + // Strip any sensitive keys that might have slipped in via older callers. + for k := range credentials { + if service.IsSensitiveCredentialKey(k) { + // Keep only keys produced by BuildAccountCredentials (tokens). + if k == "access_token" || k == "refresh_token" || k == "id_token" { + continue + } + delete(credentials, k) + } + } if reqBaseURL, ok := reqCredentials["base_url"].(string); ok && strings.TrimSpace(reqBaseURL) != "" { credentials["base_url"] = strings.TrimSpace(reqBaseURL) } - return credentials + return service.SanitizeStoredCredentials(service.PlatformGrok, credentials) } func grokSSOImportExpiry(requestExpiresAt *int64, requestAutoPause *bool, tokenInfo *service.GrokTokenInfo) (*int64, *bool) { diff --git a/backend/internal/handler/admin/grok_oauth_handler_test.go b/backend/internal/handler/admin/grok_oauth_handler_test.go index 6214394882..873773220f 100644 --- a/backend/internal/handler/admin/grok_oauth_handler_test.go +++ b/backend/internal/handler/admin/grok_oauth_handler_test.go @@ -4,6 +4,7 @@ package admin import ( "context" + "errors" "io" "net/http" "net/http/httptest" @@ -12,6 +13,7 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/config" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" @@ -128,14 +130,22 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) { require.Contains(t, rec.Body.String(), `"snapshot":`) require.Contains(t, rec.Body.String(), `"headers_observed":true`) require.NotContains(t, rec.Body.String(), "access-token") + require.Eventually(t, func() bool { + upstream.mu.Lock() + defer upstream.mu.Unlock() + return len(upstream.requests) == 4 + }, time.Second, 10*time.Millisecond) upstream.mu.Lock() requests := append([]*http.Request(nil), upstream.requests...) bodies := append([][]byte(nil), upstream.bodies...) upstream.mu.Unlock() - require.Len(t, requests, 3) + require.Len(t, requests, 4) + responsesProbeSeen := false + modelsSyncSeen := false for i, upstreamReq := range requests { require.Equal(t, "Bearer access-token", upstreamReq.Header.Get("Authorization")) if upstreamReq.URL.String() == xai.DefaultCLIBaseURL+"/responses" { + responsesProbeSeen = true require.Equal(t, "application/json, text/event-stream", upstreamReq.Header.Get("Accept")) require.Contains(t, string(bodies[i]), `"model":"grok-4.5"`) require.Contains(t, string(bodies[i]), `"input":"hi"`) @@ -143,7 +153,12 @@ func TestGrokOAuthHandlerQueryQuotaProbesUpstream(t *testing.T) { require.NotContains(t, string(bodies[i]), `"max_output_tokens"`) require.NotContains(t, string(bodies[i]), `"store"`) } + if upstreamReq.URL.String() == xai.DefaultCLIBaseURL+"/models" { + modelsSyncSeen = true + } } + require.True(t, responsesProbeSeen) + require.True(t, modelsSyncSeen) require.NotNil(t, repo.updates[42]) } @@ -189,6 +204,84 @@ func TestGrokOAuthHandlerRuntimeSanityDoesNotExposeSecrets(t *testing.T) { require.NotContains(t, rec.Body.String(), "client-secret-like-value") } +type grokOAuthHandlerClient struct{} + +func (c *grokOAuthHandlerClient) ExchangeCode(context.Context, string, string, string, string, string) (*xai.TokenResponse, error) { + return nil, errors.New("unexpected exchange") +} + +func (c *grokOAuthHandlerClient) RefreshToken(context.Context, string, string, string) (*xai.TokenResponse, error) { + return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil +} + +func (c *grokOAuthHandlerClient) LoginWithPassword(_ context.Context, email, _ string, _ string) (*service.GrokPasswordLoginResult, error) { + return &service.GrokPasswordLoginResult{ + Email: email, + SSOToken: "sso-from-password", + }, nil +} + +func (c *grokOAuthHandlerClient) ConvertSSOToBuild(context.Context, string, string) (*xai.TokenResponse, error) { + return &xai.TokenResponse{AccessToken: "access-token", RefreshToken: "refresh-token", ExpiresIn: 3600}, nil +} + +func TestGrokOAuthHandlerValidateSSOTokenReturnsTokenInfo(t *testing.T) { + gin.SetMode(gin.TestMode) + + oauthClient := &grokOAuthHandlerClient{} + oauthService := service.NewGrokOAuthService(nil, oauthClient) + defer oauthService.Stop() + handler := NewGrokOAuthHandler(oauthService, nil, nil, nil) + + router := gin.New() + router.POST("/api/v1/admin/grok/oauth/sso-token", handler.ValidateSSOToken) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/grok/oauth/sso-token", strings.NewReader(`{"sso_token":"sso-token"}`)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Contains(t, rec.Body.String(), `"access_token":"access-token"`) + require.NotContains(t, rec.Body.String(), `"sso_token"`) +} + +func TestGrokOAuthHandlerAuthorizePasswordReturnsTokenInfoWithoutPassword(t *testing.T) { + gin.SetMode(gin.TestMode) + + oauthClient := &grokOAuthHandlerClient{} + cfg := &config.Config{} + cfg.Gateway.Grok.PasswordAuthEnabled = true + oauthService := service.NewGrokOAuthService(nil, oauthClient, cfg) + defer oauthService.Stop() + handler := NewGrokOAuthHandler(oauthService, nil, nil, nil) + + router := gin.New() + router.POST("/api/v1/admin/grok/oauth/password", handler.AuthorizePassword) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/grok/oauth/password", strings.NewReader(`{"email":"user@example.com","password":"super-secret"}`)) + req.Header.Set("Content-Type", "application/json") + router.ServeHTTP(rec, req) + + require.Equal(t, http.StatusOK, rec.Code) + require.Contains(t, rec.Body.String(), `"access_token":"access-token"`) + require.NotContains(t, rec.Body.String(), "super-secret") +} + +func TestGrokOAuthHandlerPasswordCapabilityDefaultsToDisabled(t *testing.T) { + gin.SetMode(gin.TestMode) + oauthService := service.NewGrokOAuthService(nil, &grokOAuthHandlerClient{}) + defer oauthService.Stop() + handler := NewGrokOAuthHandler(oauthService, nil, nil, nil) + + router := gin.New() + router.GET("/api/v1/admin/grok/oauth/capabilities", handler.GetCapabilities) + rec := httptest.NewRecorder() + router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/v1/admin/grok/oauth/capabilities", nil)) + + require.Equal(t, http.StatusOK, rec.Code) + require.Contains(t, rec.Body.String(), `"password_auth_enabled":false`) +} + func TestGrokSSOImportExpiryUsesTokenExpiryWithoutRefreshToken(t *testing.T) { tokenExpiry := time.Now().Add(6 * time.Hour).Unix() expiresAt, autoPause := grokSSOImportExpiry(nil, nil, &service.GrokTokenInfo{ @@ -266,14 +359,12 @@ func TestGrokSSOImportCredentialsDefaultsToOfficialBaseURL(t *testing.T) { require.Equal(t, "at-2", credentials["access_token"]) } -func TestGrokSSOImportWorkerRecoversPanic(t *testing.T) { +func TestGrokSSOImportWorkerHandlesMissingOAuthService(t *testing.T) { h := &GrokOAuthHandler{} result := h.safeCreateAccountFromSSOToken(context.Background(), GrokSSOToOAuthRequest{}, "token", 2, 3) - // Without a service, createAccountFromSSOToken would panic on nil service access. - // Recovery must convert that into a failed item and keep the worker alive. require.False(t, result.created) require.Equal(t, 2, result.item.Index) - require.Contains(t, result.item.Error, "internal worker panic") + require.Contains(t, result.item.Error, "GROK_OAUTH_CLIENT_NOT_CONFIGURED") } func TestGrokOAuthHandlerReconcileDefaultsToDryRun(t *testing.T) { diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index e99ba7b755..8ddbfef183 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -106,31 +106,36 @@ type CreateGroupRequest struct { WeeklyLimitUSD optionalLimitField `json:"weekly_limit_usd"` MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` // 图片生成计费配置(antigravity 和 gemini 平台使用,负数表示清除配置) - AllowImageGeneration bool `json:"allow_image_generation"` - AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` - ImageRateIndependent bool `json:"image_rate_independent"` - ImageRateMultiplier *float64 `json:"image_rate_multiplier"` - BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` - BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` - VideoRateIndependent bool `json:"video_rate_independent"` - VideoRateMultiplier *float64 `json:"video_rate_multiplier"` - PeakRateEnabled bool `json:"peak_rate_enabled"` - PeakStart string `json:"peak_start"` - PeakEnd string `json:"peak_end"` - PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` - ProfitControlEnabled bool `json:"profit_control_enabled"` - ProfitMinMargin *float64 `json:"profit_min_margin"` - ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` - ImagePrice1K *float64 `json:"image_price_1k"` - ImagePrice2K *float64 `json:"image_price_2k"` - ImagePrice4K *float64 `json:"image_price_4k"` - VideoPrice480P *float64 `json:"video_price_480p"` - VideoPrice720P *float64 `json:"video_price_720p"` - VideoPrice1080P *float64 `json:"video_price_1080p"` - WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` - ClaudeCodeOnly bool `json:"claude_code_only"` - FallbackGroupID *int64 `json:"fallback_group_id"` - FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` + AllowImageGeneration bool `json:"allow_image_generation"` + AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` + ImageRateIndependent bool `json:"image_rate_independent"` + ImageRateMultiplier *float64 `json:"image_rate_multiplier"` + BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` + BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` + VideoRateIndependent bool `json:"video_rate_independent"` + VideoRateMultiplier *float64 `json:"video_rate_multiplier"` + PeakRateEnabled bool `json:"peak_rate_enabled"` + PeakStart string `json:"peak_start"` + PeakEnd string `json:"peak_end"` + PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` + ProfitControlEnabled bool `json:"profit_control_enabled"` + ProfitMinMargin *float64 `json:"profit_min_margin"` + ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` + ImagePrice1K *float64 `json:"image_price_1k"` + ImagePrice2K *float64 `json:"image_price_2k"` + ImagePrice4K *float64 `json:"image_price_4k"` + VideoPrice480P *float64 `json:"video_price_480p"` + VideoPrice720P *float64 `json:"video_price_720p"` + VideoPrice1080P *float64 `json:"video_price_1080p"` + VideoModelPrices map[string]map[string]float64 `json:"video_model_prices,omitempty"` + WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` + SearchPricePer1k *float64 `json:"search_price_per_1k"` + AudioRealtimePricePerMin *float64 `json:"audio_realtime_price_per_min"` + AudioTtsPricePerMillionChars *float64 `json:"audio_tts_price_per_million_chars"` + AudioSttPricePerHour *float64 `json:"audio_stt_price_per_hour"` + ClaudeCodeOnly bool `json:"claude_code_only"` + FallbackGroupID *int64 `json:"fallback_group_id"` + FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` // 模型路由配置(仅 anthropic 平台使用) ModelRouting map[string][]int64 `json:"model_routing"` ModelRoutingEnabled bool `json:"model_routing_enabled"` @@ -168,31 +173,36 @@ type UpdateGroupRequest struct { WeeklyLimitUSD optionalLimitField `json:"weekly_limit_usd"` MonthlyLimitUSD optionalLimitField `json:"monthly_limit_usd"` // 图片生成计费配置(antigravity 和 gemini 平台使用,负数表示清除配置) - AllowImageGeneration *bool `json:"allow_image_generation"` - AllowBatchImageGeneration *bool `json:"allow_batch_image_generation"` - ImageRateIndependent *bool `json:"image_rate_independent"` - ImageRateMultiplier *float64 `json:"image_rate_multiplier"` - BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` - BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` - VideoRateIndependent *bool `json:"video_rate_independent"` - VideoRateMultiplier *float64 `json:"video_rate_multiplier"` - PeakRateEnabled *bool `json:"peak_rate_enabled"` - PeakStart *string `json:"peak_start"` - PeakEnd *string `json:"peak_end"` - PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` - ProfitControlEnabled *bool `json:"profit_control_enabled"` - ProfitMinMargin *float64 `json:"profit_min_margin"` - ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` - ImagePrice1K *float64 `json:"image_price_1k"` - ImagePrice2K *float64 `json:"image_price_2k"` - ImagePrice4K *float64 `json:"image_price_4k"` - VideoPrice480P *float64 `json:"video_price_480p"` - VideoPrice720P *float64 `json:"video_price_720p"` - VideoPrice1080P *float64 `json:"video_price_1080p"` - WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` - ClaudeCodeOnly *bool `json:"claude_code_only"` - FallbackGroupID *int64 `json:"fallback_group_id"` - FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` + AllowImageGeneration *bool `json:"allow_image_generation"` + AllowBatchImageGeneration *bool `json:"allow_batch_image_generation"` + ImageRateIndependent *bool `json:"image_rate_independent"` + ImageRateMultiplier *float64 `json:"image_rate_multiplier"` + BatchImageDiscountMultiplier *float64 `json:"batch_image_discount_multiplier"` + BatchImageHoldMultiplier *float64 `json:"batch_image_hold_multiplier"` + VideoRateIndependent *bool `json:"video_rate_independent"` + VideoRateMultiplier *float64 `json:"video_rate_multiplier"` + PeakRateEnabled *bool `json:"peak_rate_enabled"` + PeakStart *string `json:"peak_start"` + PeakEnd *string `json:"peak_end"` + PeakRateMultiplier *float64 `json:"peak_rate_multiplier"` + ProfitControlEnabled *bool `json:"profit_control_enabled"` + ProfitMinMargin *float64 `json:"profit_min_margin"` + ProfitSafetyBuffer *float64 `json:"profit_safety_buffer"` + ImagePrice1K *float64 `json:"image_price_1k"` + ImagePrice2K *float64 `json:"image_price_2k"` + ImagePrice4K *float64 `json:"image_price_4k"` + VideoPrice480P *float64 `json:"video_price_480p"` + VideoPrice720P *float64 `json:"video_price_720p"` + VideoPrice1080P *float64 `json:"video_price_1080p"` + VideoModelPrices map[string]map[string]float64 `json:"video_model_prices,omitempty"` + WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` + SearchPricePer1k *float64 `json:"search_price_per_1k"` + AudioRealtimePricePerMin *float64 `json:"audio_realtime_price_per_min"` + AudioTtsPricePerMillionChars *float64 `json:"audio_tts_price_per_million_chars"` + AudioSttPricePerHour *float64 `json:"audio_stt_price_per_hour"` + ClaudeCodeOnly *bool `json:"claude_code_only"` + FallbackGroupID *int64 `json:"fallback_group_id"` + FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request"` // 模型路由配置(仅 anthropic 平台使用) ModelRouting map[string][]int64 `json:"model_routing"` ModelRoutingEnabled *bool `json:"model_routing_enabled"` @@ -519,7 +529,12 @@ func (h *GroupHandler) Create(c *gin.Context) { VideoPrice480P: req.VideoPrice480P, VideoPrice720P: req.VideoPrice720P, VideoPrice1080P: req.VideoPrice1080P, + VideoModelPrices: req.VideoModelPrices, WebSearchPricePerCall: req.WebSearchPricePerCall, + SearchPricePer1k: req.SearchPricePer1k, + AudioRealtimePricePerMin: req.AudioRealtimePricePerMin, + AudioTTSPricePerMillionChars: req.AudioTtsPricePerMillionChars, + AudioSTTPricePerHour: req.AudioSttPricePerHour, ClaudeCodeOnly: req.ClaudeCodeOnly, FallbackGroupID: req.FallbackGroupID, FallbackGroupIDOnInvalidRequest: req.FallbackGroupIDOnInvalidRequest, @@ -641,7 +656,12 @@ func (h *GroupHandler) Update(c *gin.Context) { VideoPrice480P: req.VideoPrice480P, VideoPrice720P: req.VideoPrice720P, VideoPrice1080P: req.VideoPrice1080P, + VideoModelPrices: req.VideoModelPrices, WebSearchPricePerCall: req.WebSearchPricePerCall, + SearchPricePer1k: req.SearchPricePer1k, + AudioRealtimePricePerMin: req.AudioRealtimePricePerMin, + AudioTTSPricePerMillionChars: req.AudioTtsPricePerMillionChars, + AudioSTTPricePerHour: req.AudioSttPricePerHour, ClaudeCodeOnly: req.ClaudeCodeOnly, FallbackGroupID: req.FallbackGroupID, FallbackGroupIDOnInvalidRequest: req.FallbackGroupIDOnInvalidRequest, diff --git a/backend/internal/handler/admin/setting_handler.go b/backend/internal/handler/admin/setting_handler.go index da21a0f734..2c12398baa 100644 --- a/backend/internal/handler/admin/setting_handler.go +++ b/backend/internal/handler/admin/setting_handler.go @@ -374,6 +374,10 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { ChannelMonitorDefaultIntervalSeconds: settings.ChannelMonitorDefaultIntervalSeconds, ChannelMonitorHideThroughput: settings.ChannelMonitorHideThroughput, + GrokDefaultTextModel: settings.GrokDefaultTextModel, + GrokCrossClientModelMapEnabled: settings.GrokCrossClientModelMapEnabled, + GrokDefaultBaseURLMode: settings.GrokDefaultBaseURLMode, + AvailableChannelsEnabled: settings.AvailableChannelsEnabled, ModelPlazaEnabled: settings.ModelPlazaEnabled, @@ -382,7 +386,8 @@ func (h *SettingHandler) GetSettings(c *gin.Context) { AffiliateEnabled: settings.AffiliateEnabled, - AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests, + AccountSchedulingThresholds: settings.AccountSchedulingThresholds, + AllowUserViewErrorRequests: settings.AllowUserViewErrorRequests, } // OpenAI fast policy (stored under a dedicated setting key) diff --git a/backend/internal/handler/admin/setting_handler_audit.go b/backend/internal/handler/admin/setting_handler_audit.go index 5594b6c9c1..e5370bf299 100644 --- a/backend/internal/handler/admin/setting_handler_audit.go +++ b/backend/internal/handler/admin/setting_handler_audit.go @@ -598,6 +598,9 @@ func diffSettings(before *service.SystemSettings, after *service.SystemSettings, if !equalPlatformQuotaSettings(before.DefaultPlatformQuotas, after.DefaultPlatformQuotas) { changed = append(changed, service.SettingKeyDefaultPlatformQuotas) } + if !equalAccountSchedulingThresholds(before.AccountSchedulingThresholds, after.AccountSchedulingThresholds) { + changed = append(changed, service.SettingKeyAccountSchedulingThresholds) + } changed = appendAuthSourceDefaultChanges(changed, beforeAuthSourceDefaults, afterAuthSourceDefaults) return changed } @@ -811,6 +814,27 @@ func slotOf(s *service.DefaultPlatformQuotaSetting, win string) *float64 { } // equalPlatformQuotaSettings reports whether two platform-quota maps are identical across all allowed slots. +func equalAccountSchedulingThresholds(before, after map[string]int) bool { + for _, platform := range service.AllowedSchedulingThresholdPlatforms { + beforeValue := 100 + if before != nil { + if value, ok := before[platform]; ok { + beforeValue = value + } + } + afterValue := 100 + if after != nil { + if value, ok := after[platform]; ok { + afterValue = value + } + } + if beforeValue != afterValue { + return false + } + } + return true +} + func equalPlatformQuotaSettings(before, after map[string]*service.DefaultPlatformQuotaSetting) bool { for _, platform := range service.AllowedQuotaPlatforms { b := before[platform] diff --git a/backend/internal/handler/admin/setting_handler_partial_payload_test.go b/backend/internal/handler/admin/setting_handler_partial_payload_test.go index 4d009cb0b2..b4c94809bb 100644 --- a/backend/internal/handler/admin/setting_handler_partial_payload_test.go +++ b/backend/internal/handler/admin/setting_handler_partial_payload_test.go @@ -66,6 +66,18 @@ func TestUpdateSettingsSMTPFromAliasIsWritable(t *testing.T) { require.Equal(t, "new@example.com", repo.values[service.SettingKeySMTPFrom]) } +func TestUpdateSettingsGrokDefaultBaseURLModeIsWritable(t *testing.T) { + h, repo := newStepUpSwitchTestHandler(t, map[string]string{ + service.SettingKeyGrokDefaultBaseURLMode: service.GrokDefaultBaseURLModeCLI, + }) + + rec := doUpdateSettings(t, h, map[string]any{ + "grok_default_base_url_mode": service.GrokDefaultBaseURLModeEUWest1, + }, nil) + require.Equal(t, http.StatusOK, rec.Code) + require.Equal(t, service.GrokDefaultBaseURLModeEUWest1, repo.values[service.SettingKeyGrokDefaultBaseURLMode]) +} + func TestUpdateSettingsRejectsTwoCaptchaProviders(t *testing.T) { h, _ := newStepUpSwitchTestHandler(t, map[string]string{ service.SettingKeyTurnstileEnabled: "true", diff --git a/backend/internal/handler/admin/setting_handler_update.go b/backend/internal/handler/admin/setting_handler_update.go index 99eb1f21b6..37971eb826 100644 --- a/backend/internal/handler/admin/setting_handler_update.go +++ b/backend/internal/handler/admin/setting_handler_update.go @@ -332,6 +332,11 @@ type UpdateSettingsRequest struct { ChannelMonitorDefaultIntervalSeconds *int `json:"channel_monitor_default_interval_seconds"` ChannelMonitorHideThroughput *bool `json:"channel_monitor_hide_throughput"` + // Grok model mapping policy + GrokDefaultTextModel *string `json:"grok_default_text_model"` + GrokCrossClientModelMapEnabled *bool `json:"grok_cross_client_model_map_enabled"` + GrokDefaultBaseURLMode *string `json:"grok_default_base_url_mode"` + // Available Channels feature switch (user-facing) AvailableChannelsEnabled *bool `json:"available_channels_enabled"` @@ -356,6 +361,9 @@ type UpdateSettingsRequest struct { // 系统全局 platform quota 默认值(整体替换语义:nil = 不修改,non-nil = 整体覆盖)。 DefaultPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"default_platform_quotas"` + // 各平台账号自动停调阈值(整体替换语义:nil = 不修改,non-nil = 整体覆盖)。 + AccountSchedulingThresholds map[string]int `json:"account_scheduling_thresholds"` + // auth-source 层 platform quota 覆盖(override 语义:nil = 不修改,non-nil = 整体覆盖该 source 的 quota 配置)。 AuthSourceEmailPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_email_platform_quotas"` AuthSourceLinuxDoPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"auth_source_default_linuxdo_platform_quotas"` @@ -1478,7 +1486,8 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { settings := &service.SystemSettings{ // 系统全局 platform quota 默认值(整体替换语义) - DefaultPlatformQuotas: req.DefaultPlatformQuotas, + DefaultPlatformQuotas: req.DefaultPlatformQuotas, + AccountSchedulingThresholds: req.AccountSchedulingThresholds, RegistrationEnabled: req.RegistrationEnabled, EmailVerifyEnabled: req.EmailVerifyEnabled, @@ -1874,6 +1883,24 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { } return previousSettings.ChannelMonitorHideThroughput }(), + GrokDefaultTextModel: func() string { + if req.GrokDefaultTextModel != nil { + return *req.GrokDefaultTextModel + } + return previousSettings.GrokDefaultTextModel + }(), + GrokCrossClientModelMapEnabled: func() bool { + if req.GrokCrossClientModelMapEnabled != nil { + return *req.GrokCrossClientModelMapEnabled + } + return previousSettings.GrokCrossClientModelMapEnabled + }(), + GrokDefaultBaseURLMode: func() string { + if req.GrokDefaultBaseURLMode != nil { + return strings.TrimSpace(*req.GrokDefaultBaseURLMode) + } + return previousSettings.GrokDefaultBaseURLMode + }(), AvailableChannelsEnabled: func() bool { if req.AvailableChannelsEnabled != nil { return *req.AvailableChannelsEnabled @@ -2309,6 +2336,10 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { ChannelMonitorDefaultIntervalSeconds: updatedSettings.ChannelMonitorDefaultIntervalSeconds, ChannelMonitorHideThroughput: updatedSettings.ChannelMonitorHideThroughput, + GrokDefaultTextModel: updatedSettings.GrokDefaultTextModel, + GrokCrossClientModelMapEnabled: updatedSettings.GrokCrossClientModelMapEnabled, + GrokDefaultBaseURLMode: updatedSettings.GrokDefaultBaseURLMode, + AvailableChannelsEnabled: updatedSettings.AvailableChannelsEnabled, ModelPlazaEnabled: updatedSettings.ModelPlazaEnabled, @@ -2320,6 +2351,7 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) { RiskControlEnabled: updatedSettings.RiskControlEnabled, CyberSessionBlockEnabled: updatedSettings.CyberSessionBlockEnabled, CyberSessionBlockTTLSeconds: updatedSettings.CyberSessionBlockTTLSeconds, + AccountSchedulingThresholds: updatedSettings.AccountSchedulingThresholds, AllowUserViewErrorRequests: updatedSettings.AllowUserViewErrorRequests, } if fastPolicy, err := h.settingService.GetOpenAIFastPolicySettings(c.Request.Context()); err != nil { diff --git a/backend/internal/handler/admin/usage_handler.go b/backend/internal/handler/admin/usage_handler.go index ad2bbecd22..829730c647 100644 --- a/backend/internal/handler/admin/usage_handler.go +++ b/backend/internal/handler/admin/usage_handler.go @@ -143,6 +143,16 @@ func (h *UsageHandler) List(c *gin.Context) { billingType = &bt } + var upstreamModelMismatch *bool + if raw := strings.TrimSpace(c.Query("upstream_model_mismatch")); raw != "" { + value, err := strconv.ParseBool(raw) + if err != nil { + response.BadRequest(c, "Invalid upstream_model_mismatch value, use true or false") + return + } + upstreamModelMismatch = &value + } + // Parse date range var startTime, endTime *time.Time userTZ := c.Query("timezone") // Get user's timezone from request @@ -173,20 +183,21 @@ func (h *UsageHandler) List(c *gin.Context) { SortOrder: c.DefaultQuery("sort_order", "desc"), } filters := usagestats.UsageLogFilters{ - UserID: userID, - APIKeyID: apiKeyID, - AccountID: accountID, - GroupID: groupID, - RequestID: requestID, - Model: model, - ModelFilterSource: usagestats.ModelSourceRequested, - RequestType: requestType, - Stream: stream, - BillingType: billingType, - BillingMode: billingMode, - StartTime: startTime, - EndTime: endTime, - ExactTotal: exactTotal, + UserID: userID, + APIKeyID: apiKeyID, + AccountID: accountID, + GroupID: groupID, + RequestID: requestID, + Model: model, + ModelFilterSource: usagestats.ModelSourceRequested, + RequestType: requestType, + Stream: stream, + BillingType: billingType, + BillingMode: billingMode, + UpstreamModelMismatch: upstreamModelMismatch, + StartTime: startTime, + EndTime: endTime, + ExactTotal: exactTotal, } records, result, err := h.usageService.ListWithFilters(c.Request.Context(), params, filters) @@ -276,6 +287,16 @@ func (h *UsageHandler) Stats(c *gin.Context) { billingType = &bt } + var upstreamModelMismatch *bool + if raw := strings.TrimSpace(c.Query("upstream_model_mismatch")); raw != "" { + value, err := strconv.ParseBool(raw) + if err != nil { + response.BadRequest(c, "Invalid upstream_model_mismatch value, use true or false") + return + } + upstreamModelMismatch = &value + } + // Parse date range userTZ := c.Query("timezone") now := timezone.NowInUserLocation(userTZ) @@ -315,18 +336,19 @@ func (h *UsageHandler) Stats(c *gin.Context) { // Build filters and call GetStatsWithFilters filters := usagestats.UsageLogFilters{ - UserID: userID, - APIKeyID: apiKeyID, - AccountID: accountID, - GroupID: groupID, - Model: model, - ModelFilterSource: usagestats.ModelSourceRequested, - RequestType: requestType, - Stream: stream, - BillingType: billingType, - BillingMode: billingMode, - StartTime: &startTime, - EndTime: &endTime, + UserID: userID, + APIKeyID: apiKeyID, + AccountID: accountID, + GroupID: groupID, + Model: model, + ModelFilterSource: usagestats.ModelSourceRequested, + RequestType: requestType, + Stream: stream, + BillingType: billingType, + BillingMode: billingMode, + UpstreamModelMismatch: upstreamModelMismatch, + StartTime: &startTime, + EndTime: &endTime, } var stats *usagestats.UsageStats diff --git a/backend/internal/handler/admin/usage_query_cache.go b/backend/internal/handler/admin/usage_query_cache.go index b288a95ba4..f6b7746e1b 100644 --- a/backend/internal/handler/admin/usage_query_cache.go +++ b/backend/internal/handler/admin/usage_query_cache.go @@ -11,17 +11,18 @@ import ( 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"` + 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"` + UpstreamModelMismatch *bool `json:"upstream_model_mismatch"` } func usageStatsCacheKey(filters usagestats.UsageLogFilters) string { @@ -34,17 +35,18 @@ func usageStatsCacheKey(filters usagestats.UsageLogFilters) string { 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, + 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, + UpstreamModelMismatch: filters.UpstreamModelMismatch, }) } diff --git a/backend/internal/handler/auth_oauth_pending_flow.go b/backend/internal/handler/auth_oauth_pending_flow.go index 32e669b226..7b1d71dca4 100644 --- a/backend/internal/handler/auth_oauth_pending_flow.go +++ b/backend/internal/handler/auth_oauth_pending_flow.go @@ -1998,6 +1998,21 @@ func (h *AuthHandler) ExchangePendingOAuthCompletion(c *gin.Context) { response.Success(c, payload) return } + // ─── 安全修复(账号接管 0day)──────────────────────────────────────────── + // 非终态 session(如 choose_account_action_required)的 TargetUserID 可能来自 + // 攻击者提交的他人邮箱:createPendingOAuthAccount / SendPendingOAuthVerifyCode + // 发现邮箱已存在时会把本 session 指向该邮箱用户,全程无密码、无邮箱验证码、 + // 无账号所有权证明。若此时带着 adoption decision 继续执行,下方的 + // applyPendingOAuthAdoption 会把本 OAuth identity 直接绑定到 TargetUserID, + // 攻击者随后再次 OAuth 登录即被系统识别为受害者本人(完整账号接管)。 + // 只有两类 session 允许在此处执行 adoption/binding: + // 1. canIssueTokenPair == true —— 登录终态,identity 已安全绑定该用户; + // 2. intent == bind_current_user —— 已登录用户主动发起绑定(绑定目标来自登录态 cookie)。 + // 其余状态一律只返回 payload,不绑定、不消费 session。 + if !canIssueTokenPair && !strings.EqualFold(strings.TrimSpace(session.Intent), oauthIntentBindCurrentUser) { + response.Success(c, payload) + return + } if !adoptionDecision.hasDecision() { adoptionRequired, _ := payload["adoption_required"].(bool) if adoptionRequired { diff --git a/backend/internal/handler/auth_oauth_pending_flow_test.go b/backend/internal/handler/auth_oauth_pending_flow_test.go index 4da8db2ab6..76c29602e8 100644 --- a/backend/internal/handler/auth_oauth_pending_flow_test.go +++ b/backend/internal/handler/auth_oauth_pending_flow_test.go @@ -910,6 +910,92 @@ func TestExchangePendingOAuthCompletionRejectsDisabledTargetUser(t *testing.T) { require.Nil(t, storedSession.ConsumedAt) } +func TestExchangePendingOAuthCompletionChoiceStateDoesNotBindIdentity(t *testing.T) { + // 回归测试:复刻"补邮箱/创建账户"路径的账号接管 0day。 + // 攻击者用自己的 OAuth 账号登录后,在 create-account 步骤提交受害者邮箱, + // 后端发现邮箱已存在会把 pending session 转入 choice 状态并指向受害者 + // (TargetUserID=受害者、无密码/验证码证明)。此时带 adoption decision 调 + // exchange 绝不能把 OAuth identity 绑定到受害者账号。 + handler, client := newOAuthPendingFlowTestHandler(t, false) + ctx := context.Background() + + victim, err := client.User.Create(). + SetEmail("victim@example.com"). + SetUsername("victim-user"). + SetPasswordHash("hash"). + SetRole(service.RoleUser). + SetStatus(service.StatusActive). + Save(ctx) + require.NoError(t, err) + + session, err := client.PendingAuthSession.Create(). + SetSessionToken("choice-state-attack-session-token"). + SetIntent("login"). + SetProviderType("linuxdo"). + SetProviderKey("linuxdo"). + SetProviderSubject("attacker-subject-123"). + SetTargetUserID(victim.ID). + SetResolvedEmail(victim.Email). + SetBrowserSessionKey("choice-state-attack-browser-session-key"). + SetUpstreamIdentityClaims(map[string]any{ + "username": "attacker_linuxdo_user", + "suggested_display_name": "Attacker Display Name", + "suggested_avatar_url": "https://cdn.example/attacker.png", + }). + SetLocalFlowState(map[string]any{ + oauthCompletionResponseKey: map[string]any{ + "step": oauthPendingChoiceStep, + "adoption_required": true, + "force_email_on_signup": true, + "email_binding_required": true, + "existing_account_bindable": true, + "email": victim.Email, + "resolved_email": victim.Email, + "redirect": "/dashboard", + }, + }). + SetExpiresAt(time.Now().UTC().Add(10 * time.Minute)). + Save(ctx) + require.NoError(t, err) + + body := bytes.NewBufferString(`{"adopt_display_name":true,"adopt_avatar":true}`) + recorder := httptest.NewRecorder() + ginCtx, _ := gin.CreateTestContext(recorder) + req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/oauth/pending/exchange", body) + req.Header.Set("Content-Type", "application/json") + req.AddCookie(&http.Cookie{Name: oauthPendingSessionCookieName, Value: encodeCookieValue(session.SessionToken)}) + req.AddCookie(&http.Cookie{Name: oauthPendingBrowserCookieName, Value: encodeCookieValue("choice-state-attack-browser-session-key")}) + ginCtx.Request = req + + handler.ExchangePendingOAuthCompletion(ginCtx) + + require.Equal(t, http.StatusOK, recorder.Code) + data := decodeJSONResponseData(t, recorder) + require.NotContains(t, data, "access_token") + require.Equal(t, oauthPendingChoiceStep, data["step"]) + + // 攻击者的 OAuth identity 绝不能绑定到受害者账号 + identityCount, err := client.AuthIdentity.Query(). + Where( + authidentity.ProviderTypeEQ("linuxdo"), + authidentity.ProviderKeyEQ("linuxdo"), + authidentity.ProviderSubjectEQ("attacker-subject-123"), + ). + Count(ctx) + require.NoError(t, err) + require.Zero(t, identityCount) + + // 受害者资料不得被 adoption 篡改 + storedVictim, err := client.User.Get(ctx, victim.ID) + require.NoError(t, err) + require.Equal(t, "victim-user", storedVictim.Username) + + // session 不得被消费(攻击者无法进入下一环) + storedSession, err := client.PendingAuthSession.Get(ctx, session.ID) + require.NoError(t, err) + require.Nil(t, storedSession.ConsumedAt) +} + func TestNormalizePendingOAuthCompletionResponseScrubsLegacyTokenPayload(t *testing.T) { payload := normalizePendingOAuthCompletionResponse(map[string]any{ "access_token": "legacy-access-token", diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index c6f6cf25e2..0c7d892297 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -202,7 +202,12 @@ func groupFromServiceBase(g *service.Group) Group { VideoPrice480P: g.VideoPrice480P, VideoPrice720P: g.VideoPrice720P, VideoPrice1080P: g.VideoPrice1080P, + VideoModelPrices: g.VideoModelPrices, WebSearchPricePerCall: g.WebSearchPricePerCall, + SearchPricePer1k: g.SearchPricePer1k, + AudioRealtimePricePerMin: g.AudioRealtimePricePerMin, + AudioTtsPricePerMillionChars: g.AudioTTSPricePerMillionChars, + AudioSttPricePerHour: g.AudioSTTPricePerHour, ClaudeCodeOnly: g.ClaudeCodeOnly, FallbackGroupID: g.FallbackGroupID, FallbackGroupIDOnInvalidRequest: g.FallbackGroupIDOnInvalidRequest, @@ -705,6 +710,8 @@ func UsageLogFromServiceAdmin(l *service.UsageLog) *AdminUsageLog { return &AdminUsageLog{ UsageLog: usageLog, UpstreamModel: l.UpstreamModel, + UpstreamResponseModel: l.UpstreamResponseModel, + UpstreamModelMismatch: l.UpstreamModelMismatch, ChannelID: l.ChannelID, ModelMappingChain: l.ModelMappingChain, BillingTier: l.BillingTier, diff --git a/backend/internal/handler/dto/mappers_usage_test.go b/backend/internal/handler/dto/mappers_usage_test.go index 6cbcbada56..d396bbfa2f 100644 --- a/backend/internal/handler/dto/mappers_usage_test.go +++ b/backend/internal/handler/dto/mappers_usage_test.go @@ -110,11 +110,15 @@ func TestUsageLogFromService_UsesRequestedModelAndKeepsUpstreamAdminOnly(t *test t.Parallel() upstreamModel := "claude-sonnet-4-20250514" + upstreamResponseModel := "claude-sonnet-4-20250513" + upstreamModelMismatch := true log := &service.UsageLog{ - RequestID: "req_4", - Model: upstreamModel, - RequestedModel: "claude-sonnet-4", - UpstreamModel: &upstreamModel, + RequestID: "req_4", + Model: upstreamModel, + RequestedModel: "claude-sonnet-4", + UpstreamModel: &upstreamModel, + UpstreamResponseModel: &upstreamResponseModel, + UpstreamModelMismatch: &upstreamModelMismatch, } userDTO := UsageLogFromService(log) @@ -126,10 +130,14 @@ func TestUsageLogFromService_UsesRequestedModelAndKeepsUpstreamAdminOnly(t *test userJSON, err := json.Marshal(userDTO) require.NoError(t, err) require.NotContains(t, string(userJSON), "upstream_model") + require.NotContains(t, string(userJSON), "upstream_response_model") + require.NotContains(t, string(userJSON), "upstream_model_mismatch") adminJSON, err := json.Marshal(adminDTO) require.NoError(t, err) require.Contains(t, string(adminJSON), `"upstream_model":"claude-sonnet-4-20250514"`) + require.Contains(t, string(adminJSON), `"upstream_response_model":"claude-sonnet-4-20250513"`) + require.Contains(t, string(adminJSON), `"upstream_model_mismatch":true`) } func TestUsageLogFromService_KeepsUserBillingAndIPWithoutAdminCostFields(t *testing.T) { diff --git a/backend/internal/handler/dto/settings.go b/backend/internal/handler/dto/settings.go index e849f74940..6e67cd2ac5 100644 --- a/backend/internal/handler/dto/settings.go +++ b/backend/internal/handler/dto/settings.go @@ -305,6 +305,11 @@ type SystemSettings struct { ChannelMonitorDefaultIntervalSeconds int `json:"channel_monitor_default_interval_seconds"` ChannelMonitorHideThroughput bool `json:"channel_monitor_hide_throughput"` + // Grok model mapping policy (admin settings; empty account mapping falls back to these). + GrokDefaultTextModel string `json:"grok_default_text_model"` + GrokCrossClientModelMapEnabled bool `json:"grok_cross_client_model_map_enabled"` + GrokDefaultBaseURLMode string `json:"grok_default_base_url_mode"` + // Available Channels feature switch (user-facing aggregate view) AvailableChannelsEnabled bool `json:"available_channels_enabled"` @@ -329,6 +334,9 @@ type SystemSettings struct { // 系统全局默认平台配额(key = platform,nil/缺省 = 不限制) DefaultPlatformQuotas map[string]*service.DefaultPlatformQuotaSetting `json:"default_platform_quotas,omitempty"` + // 系统全局账号自动停调阈值(key = platform,100 = disabled) + AccountSchedulingThresholds map[string]int `json:"account_scheduling_thresholds,omitempty"` + // 允许终端用户在用量页查看自己的失败请求 AllowUserViewErrorRequests bool `json:"allow_user_view_error_requests"` } diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 37927c2b48..bc2c502c1f 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -121,8 +121,14 @@ type Group struct { VideoPrice480P *float64 `json:"video_price_480p"` VideoPrice720P *float64 `json:"video_price_720p"` VideoPrice1080P *float64 `json:"video_price_1080p"` + // VideoModelPrices 可选按模型族×分辨率覆盖视频每秒单价 (USD/s)。 + VideoModelPrices map[string]map[string]float64 `json:"video_model_prices,omitempty"` // Codex alpha/search 网页搜索单次价格(USD/次);null 表示使用默认价 0.01 - WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` + WebSearchPricePerCall *float64 `json:"web_search_price_per_call"` + SearchPricePer1k *float64 `json:"search_price_per_1k"` + AudioRealtimePricePerMin *float64 `json:"audio_realtime_price_per_min"` + AudioTtsPricePerMillionChars *float64 `json:"audio_tts_price_per_million_chars"` + AudioSttPricePerHour *float64 `json:"audio_stt_price_per_hour"` // Claude Code 客户端限制 ClaudeCodeOnly bool `json:"claude_code_only"` @@ -555,6 +561,10 @@ type AdminUsageLog struct { // UpstreamModel is the actual model sent to the upstream provider after mapping. // Omitted when no mapping was applied (requested model was used as-is). UpstreamModel *string `json:"upstream_model,omitempty"` + // UpstreamResponseModel is the raw model declared by the upstream response. + UpstreamResponseModel *string `json:"upstream_response_model,omitempty"` + // UpstreamModelMismatch is nil when the upstream did not declare a model. + UpstreamModelMismatch *bool `json:"upstream_model_mismatch,omitempty"` // ChannelID 渠道 ID ChannelID *int64 `json:"channel_id,omitempty"` diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 287d59fd9f..1e2d0fbe3d 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -2416,6 +2416,33 @@ func (h *GatewayHandler) submitUsageRecordTask(parent context.Context, task serv task(ctx) } +// submitMandatoryUsageRecordTask never silently drops billing work on pool overflow. +func (h *GatewayHandler) 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.Dropped() { + return + } + logger.L().With( + zap.String("component", "handler.gateway.usage"), + ).Warn("gateway.usage_record_task_mandatory_sync_fallback") + } + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + defer func() { + if recovered := recover(); recovered != nil { + logger.L().With( + zap.String("component", "handler.gateway.usage"), + zap.Any("panic", recovered), + ).Error("gateway.usage_record_task_panic_recovered") + } + }() + task(ctx) +} + // getUserMsgQueueMode 获取当前请求的 UMQ 模式 // 返回 "serialize" | "throttle" | "" func (h *GatewayHandler) getUserMsgQueueMode(account *service.Account, parsed *service.ParsedRequest) string { diff --git a/backend/internal/handler/gateway_web_search.go b/backend/internal/handler/gateway_web_search.go new file mode 100644 index 0000000000..a0a249f252 --- /dev/null +++ b/backend/internal/handler/gateway_web_search.go @@ -0,0 +1,509 @@ +package handler + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "net/url" + "strconv" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ip" + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/websearch" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/google/uuid" + "github.com/tidwall/gjson" + "go.uber.org/zap" +) + +const ( + defaultGrokWebSearchResults = 5 + maxGrokWebSearchResults = 20 +) + +func (h *GatewayHandler) WebSearch(c *gin.Context) { + type webSearchReq struct { + Query string `json:"query" binding:"required"` + MaxResults int `json:"max_results"` + } + + var req webSearchReq + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ + "type": "invalid_request_error", + "message": err.Error(), + }}) + return + } + req.MaxResults = normalizeGrokWebSearchMaxResults(req.MaxResults) + + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey == nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": gin.H{ + "type": "authentication_error", + "message": "API key required", + }}) + return + } + + if apiKey.Group == nil || apiKey.Group.Platform != "grok" { + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ + "type": "invalid_request_error", + "message": "web search is only supported for grok groups", + }}) + return + } + + // Billing eligibility (same as other requests) + subscription, _ := middleware2.GetSubscriptionFromContext(c) + if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil { + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + c.JSON(status, gin.H{"error": gin.H{"type": code, "message": message}}) + return + } + + subject, _ := middleware2.GetAuthSubjectFromContext(c) + reqLog := requestLogger(c, "handler.gateway.web_search") + // Audit user search query before upstream Grok web_search traffic. + auditBody, _ := json.Marshal(map[string]any{ + "messages": []map[string]any{{ + "role": "user", "content": req.Query, + }}, + }) + if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, xai.DefaultTextModel, auditBody); decision != nil && !decision.AllowNextStage { + status := decision.HTTPStatus + if status == 0 { + status = http.StatusForbidden + } + code := decision.ErrorCode + if code == "" { + code = "content_policy_violation" + } + msg := decision.ClientMessage + if msg == "" { + msg = "Request blocked by content policy" + } + c.JSON(status, gin.H{"error": gin.H{"type": code, "message": msg}}) + return + } + + // Use exactly the same scheduling as other requests (SelectAccountWithLoadAwareness handles load, rate limit, sticky, etc.) + groupID := apiKey.GroupID + if groupID == nil { + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{ + "type": "invalid_request_error", + "message": "group required", + }}) + return + } + + failedAccounts := make(map[int64]struct{}) + var account *service.Account + var accountReleaseFunc func() + var nativeResp *websearch.SearchResponse + var providerName string + var err error + + // Acquire + release holder for the whole handler (including failover retries). + defer func() { + if accountReleaseFunc != nil { + accountReleaseFunc() + } + }() + + // First attempt + up to 3 failover accounts (max 4 total). + for attempt := 0; attempt < 4; attempt++ { + selected, selectErr := h.gatewayService.SelectAccountWithLoadAwareness( + c.Request.Context(), groupID, "", xai.DefaultTextModel, failedAccounts, "", 0, + ) + if selectErr != nil { + if attempt == 0 { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": gin.H{ + "type": "scheduling_error", + "message": selectErr.Error(), + }}) + return + } + break + } + if selected == nil || selected.Account == nil { + if attempt == 0 { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": gin.H{ + "type": "scheduling_error", + "message": "No available accounts", + }}) + return + } + break + } + + release, acquireOK, acquireErr := h.acquireWebSearchAccountSlot(c, selected) + if !acquireOK { + // First hop: surface concurrency errors; later hops try another account. + if attempt == 0 && acquireErr != nil { + h.handleConcurrencyError(c, acquireErr, "account", false) + return + } + failedAccounts[selected.Account.ID] = struct{}{} + continue + } + account = selected.Account + accountReleaseFunc = release + + nativeResp, providerName, err = h.doGrokNativeWebSearch(c.Request.Context(), c, account, req.Query, req.MaxResults) + if err == nil { + break + } + var failoverErr *service.UpstreamFailoverError + if !errors.As(err, &failoverErr) || !failoverErr.ShouldRetryNextAccount() { + break + } + failedAccounts[account.ID] = struct{}{} + if accountReleaseFunc != nil { + accountReleaseFunc() + accountReleaseFunc = nil + } + account = nil + } + if err != nil || nativeResp == nil { + msg := "web search failed" + if err != nil { + msg = err.Error() + } + c.JSON(http.StatusBadGateway, gin.H{"error": gin.H{"type": "web_search_error", "message": msg}}) + return + } + if account == nil { + c.JSON(http.StatusServiceUnavailable, gin.H{"error": gin.H{ + "type": "scheduling_error", + "message": "No available accounts", + }}) + return + } + + userAgent := c.GetHeader("User-Agent") + clientIP := ip.GetClientIP(c) + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) + requestPayloadHash := service.HashUsageRequestPayload([]byte(req.Query)) + quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) + // Request IDs are billing idempotency keys, so they must be unique per invocation. + // Query/IP/UA hashes would collapse repeated identical searches into one charge. + searchRequestID := "web_search:" + uuid.NewString() + if apiKey.Group != nil { + if p := apiKey.Group.GetSearchPricePer1k(); p != nil && *p == 0 { + logger.L().With( + zap.String("component", "handler.gateway.web_search"), + zap.Int64("group_id", apiKey.Group.ID), + ).Info("gateway.web_search.search_price_per_1k_explicit_free") + } + } + h.submitMandatoryUsageRecordTask(c.Request.Context(), func(ctx context.Context) { + if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{ + Result: &service.ForwardResult{ + RequestID: searchRequestID, + Model: "grok-web-search", + SearchCount: 1, + Duration: 0, + }, + APIKey: apiKey, + User: apiKey.User, + Account: account, + Subscription: subscription, + InboundEndpoint: inboundEndpoint, + UpstreamEndpoint: upstreamEndpoint, + UserAgent: userAgent, + IPAddress: clientIP, + RequestPayloadHash: requestPayloadHash, + APIKeyService: h.apiKeyService, + QuotaPlatform: quotaPlatform, + }); err != nil { + logger.L().With( + zap.String("component", "handler.gateway.web_search"), + zap.Int64("user_id", apiKey.User.ID), + zap.Int64("api_key_id", apiKey.ID), + zap.Int64("account_id", account.ID), + ).Error("gateway.web_search.record_usage_failed", zap.Error(err)) + } + }) + + c.JSON(http.StatusOK, gin.H{ + "query": req.Query, + "results": nativeResp.Results, + "provider": providerName, + "max_results": req.MaxResults, + }) +} + +// acquireWebSearchAccountSlot resolves an immediate slot or WaitPlan wait. +// On failure returns (nil, false, err); err is non-nil for concurrency acquire +// failures so the first hop can map them to HTTP. Wait-queue full returns +// (nil, false, nil) so failover can try another account. +func (h *GatewayHandler) acquireWebSearchAccountSlot( + c *gin.Context, + selected *service.AccountSelectionResult, +) (release func(), ok bool, acquireErr error) { + if selected == nil || selected.Account == nil { + return nil, false, nil + } + if selected.Acquired { + return selected.ReleaseFunc, true, nil + } + if selected.WaitPlan == nil || h.concurrencyHelper == nil { + return nil, false, nil + } + account := selected.Account + accountWaitCounted := false + canWait, waitErr := h.concurrencyHelper.IncrementAccountWaitCount(c.Request.Context(), account.ID, selected.WaitPlan.MaxWaiting) + if waitErr != nil { + logger.L().Warn("gateway.web_search.account_wait_counter_increment_failed", + zap.Int64("account_id", account.ID), + zap.Error(waitErr), + ) + // Best-effort wait without counter (same as first-hop legacy path). + } else if !canWait { + return nil, false, nil + } else { + accountWaitCounted = true + } + releaseWait := func() { + if accountWaitCounted { + h.concurrencyHelper.DecrementAccountWaitCount(c.Request.Context(), account.ID) + accountWaitCounted = false + } + } + streamStarted := false + slotRelease, err := h.concurrencyHelper.AcquireAccountSlotWithWaitTimeout( + c, + account.ID, + selected.WaitPlan.MaxConcurrency, + selected.WaitPlan.Timeout, + false, + &streamStarted, + ) + releaseWait() + if err != nil { + return nil, false, err + } + return slotRelease, true, nil +} + +// doGrokNativeWebSearch executes web search using the Grok account's native capability +// by calling the responses endpoint with web_search tool, then normalizes sources to unified format. +func (h *GatewayHandler) doGrokNativeWebSearch(ctx context.Context, c *gin.Context, account *service.Account, query string, maxResults int) (*websearch.SearchResponse, string, error) { + maxResults = normalizeGrokWebSearchMaxResults(maxResults) + + // Build a minimal responses request that triggers Grok web search tool. + // Ask for structured metadata because xAI action.sources commonly contains URLs only. + searchBody := map[string]any{ + "model": xai.DefaultTextModel, + "input": buildGrokWebSearchPrompt(query, maxResults), + "tools": []map[string]any{{"type": "web_search"}}, + "include": []string{"web_search_call.action.sources"}, + "store": false, + "stream": false, + } + bodyBytes, _ := json.Marshal(searchBody) + + respBytes, err := h.gatewayService.DoGrokNativeResponsesJSON(ctx, account, bodyBytes) + if err != nil { + return nil, "", err + } + + // Extract sources from Grok responses output. + // Prefer web_search_call.action.sources (standardized), fallback to annotations or text links. + results := extractGrokWebSearchSources(respBytes, maxResults) + + return &websearch.SearchResponse{ + Results: results, + Query: query, + }, "grok-native", nil +} + +func normalizeGrokWebSearchMaxResults(maxResults int) int { + if maxResults <= 0 { + return defaultGrokWebSearchResults + } + if maxResults > maxGrokWebSearchResults { + return maxGrokWebSearchResults + } + return maxResults +} + +func buildGrokWebSearchPrompt(query string, maxResults int) string { + return fmt.Sprintf(`Search the web for the user query below. Return ONLY valid JSON with this exact shape: {"results":[{"url":"https://...","title":"page title","snippet":"concise factual summary"}]}. Return at most %d unique results. Every URL must be an actual web_search source. Populate a non-empty title and snippet for every result. Do not wrap the JSON in markdown. + +User query: +%s`, normalizeGrokWebSearchMaxResults(maxResults), query) +} + +// extractGrokWebSearchSources returns model-enriched results only when their URLs +// are present in the actual web_search sources, then falls back to raw sources. +func extractGrokWebSearchSources(body []byte, maxResults int) []websearch.SearchResult { + if len(body) == 0 || !gjson.ValidBytes(body) { + return nil + } + maxResults = normalizeGrokWebSearchMaxResults(maxResults) + + sources := make(map[string]websearch.SearchResult) + var sourceOrder []string + addSource := func(rawURL, title, snippet string) { + key, ok := normalizeGrokWebSearchURL(rawURL) + if !ok { + return + } + result, exists := sources[key] + if !exists { + result.URL = strings.TrimSpace(rawURL) + sourceOrder = append(sourceOrder, key) + } + if result.Title == "" { + result.Title = usableGrokWebSearchTitle(title, result.URL) + } + if result.Snippet == "" { + result.Snippet = strings.TrimSpace(snippet) + } + sources[key] = result + } + + output := gjson.GetBytes(body, "output") + output.ForEach(func(_, item gjson.Result) bool { + if item.Get("type").String() == "web_search_call" { + sources := item.Get("action.sources") + if sources.IsArray() { + sources.ForEach(func(_, src gjson.Result) bool { + addSource(src.Get("url").String(), src.Get("title").String(), src.Get("snippet").String()) + return true + }) + } + } + if item.Get("type").String() == "message" { + item.Get("content").ForEach(func(_, part gjson.Result) bool { + if part.Get("type").String() != "output_text" { + return true + } + part.Get("annotations").ForEach(func(_, ann gjson.Result) bool { + if ann.Get("type").String() == "url_citation" || ann.Get("type").String() == "web" { + addSource(ann.Get("url").String(), ann.Get("title").String(), "") + } + return true + }) + return true + }) + } + return true + }) + + var out []websearch.SearchResult + seen := make(map[string]bool) + output.ForEach(func(_, item gjson.Result) bool { + if item.Get("type").String() != "message" { + return true + } + item.Get("content").ForEach(func(_, part gjson.Result) bool { + if part.Get("type").String() != "output_text" || len(out) >= maxResults { + return true + } + for _, result := range parseGrokWebSearchStructuredResults(part.Get("text").String()) { + key, ok := normalizeGrokWebSearchURL(result.URL) + if !ok || seen[key] { + continue + } + source, allowed := sources[key] + if !allowed { + continue + } + seen[key] = true + result.URL = source.URL + result.Title = usableGrokWebSearchTitle(result.Title, result.URL) + if result.Title == "" { + result.Title = source.Title + } + result.Snippet = strings.TrimSpace(result.Snippet) + if result.Snippet == "" { + result.Snippet = source.Snippet + } + out = append(out, result) + if len(out) >= maxResults { + break + } + } + return true + }) + return len(out) < maxResults + }) + + for _, key := range sourceOrder { + if len(out) >= maxResults { + break + } + if seen[key] { + continue + } + result := sources[key] + if result.Title == "" { + result.Title = grokWebSearchTitleFromURL(result.URL) + } + seen[key] = true + out = append(out, result) + } + return out +} + +func parseGrokWebSearchStructuredResults(text string) []websearch.SearchResult { + text = strings.TrimSpace(text) + start := strings.IndexByte(text, '{') + end := strings.LastIndexByte(text, '}') + if start < 0 || end < start { + return nil + } + var payload struct { + Results []websearch.SearchResult `json:"results"` + } + if err := json.Unmarshal([]byte(text[start:end+1]), &payload); err != nil { + return nil + } + return payload.Results +} + +func normalizeGrokWebSearchURL(rawURL string) (string, bool) { + u, err := url.Parse(strings.TrimSpace(rawURL)) + if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") { + return "", false + } + u.Scheme = strings.ToLower(u.Scheme) + u.Host = strings.ToLower(u.Host) + u.Fragment = "" + if u.Path == "" { + u.Path = "/" + } + return u.String(), true +} + +func usableGrokWebSearchTitle(title, rawURL string) string { + title = strings.TrimSpace(title) + if title == "" || title == rawURL { + return "" + } + if _, err := strconv.Atoi(title); err == nil { + return "" + } + return title +} + +func grokWebSearchTitleFromURL(rawURL string) string { + u, err := url.Parse(rawURL) + if err != nil || u.Host == "" { + return rawURL + } + return strings.TrimPrefix(strings.ToLower(u.Host), "www.") +} diff --git a/backend/internal/handler/grok_audio.go b/backend/internal/handler/grok_audio.go new file mode 100644 index 0000000000..542d70d423 --- /dev/null +++ b/backend/internal/handler/grok_audio.go @@ -0,0 +1,338 @@ +package handler + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "strconv" + "strings" + "time" + + "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" + coderws "github.com/coder/websocket" + "github.com/gin-gonic/gin" + "go.uber.org/zap" +) + +// GrokRealtime exposes xAI's native Voice Realtime WebSocket. +// Only Grok-platform API keys may use this endpoint. +func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) { + if c == nil || c.Request == nil || !isOpenAIWSUpgradeRequest(c.Request) { + h.errorResponse(c, http.StatusUpgradeRequired, "invalid_request_error", "WebSocket upgrade required (Upgrade: websocket)") + return + } + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey.Group == nil || apiKey.Group.Platform != service.PlatformGrok { + h.errorResponse(c, http.StatusNotFound, "not_found_error", "Realtime API is not supported for this platform") + return + } + if !h.ensureResponsesDependencies(c, nil) { + return + } + subscription, _ := middleware2.GetSubscriptionFromContext(c) + if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil { + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + h.errorResponse(c, status, code, message) + return + } + + selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability( + c.Request.Context(), + apiKey.GroupID, + "", + "", + "grok-4.5", + nil, + service.OpenAIUpstreamTransportHTTPSSE, + // Grok only advertises chat_completions + media capabilities on HEAD. + service.OpenAIEndpointCapabilityChatCompletions, + false, + false, + false, + service.PlatformGrok, + ) + if err != nil || selection == nil || selection.Account == nil { + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available Grok accounts") + return + } + + var streamStarted bool + reqLog := requestLogger(c, "handler.openai_gateway.grok_realtime") + release, slotStatus := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, true, &streamStarted, reqLog) + if slotStatus != openAISlotAcquireOK { + return + } + defer release() + + token, _, err := h.gatewayService.GetRequestCredential(c.Request.Context(), c, selection.Account) + if err != nil { + h.errorResponse(c, http.StatusBadGateway, "upstream_error", "Grok credential unavailable") + return + } + + conn, err := coderws.Accept(c.Writer, c.Request, &coderws.AcceptOptions{CompressionMode: coderws.CompressionContextTakeover}) + if err != nil { + return + } + defer func() { _ = conn.CloseNow() }() + + model := c.Query("model") + if strings.TrimSpace(model) == "" { + model = "grok-voice-latest" + } + started := time.Now() + proxyErr := h.gatewayService.ProxyGrokRealtime(c.Request.Context(), c, conn, selection.Account, token, model) + elapsed := time.Since(started) + if proxyErr != nil { + reqLog.Info("grok_realtime.proxy_failed", zap.Error(proxyErr)) + if !isExpectedGrokRealtimeClose(proxyErr) { + _ = conn.Close(coderws.StatusInternalError, "upstream realtime websocket failed") + return + } + } + // A relay normally returns a close error when either side closes normally. + // Those sessions still consumed upstream audio time and must be billed. + if elapsed > 0 { + result := &service.OpenAIForwardResult{ + // One durable id per WS session so retries cannot collapse or double under client ids. + RequestID: service.StableGrokRealtimeBillingRequestID(""), + Model: model, + Duration: elapsed, + AudioUsage: &service.AudioUsage{Mode: "realtime", DurationOrUnits: elapsed.Minutes()}, + } + h.recordGrokVoiceUsage(c, apiKey, selection.Account, subscription, "realtime", nil, result) + } +} + +func isExpectedGrokRealtimeClose(err error) bool { + if err == nil { + return true + } + switch coderws.CloseStatus(err) { + case coderws.StatusNormalClosure, coderws.StatusGoingAway, + coderws.StatusNoStatusRcvd, coderws.StatusAbnormalClosure: + return true + default: + return false + } +} + +// GrokVoice handles xAI Voice HTTP endpoints. endpoint is "tts", "stt", or "custom-voices". +func (h *OpenAIGatewayHandler) GrokVoice(c *gin.Context, endpoint string) { + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok || apiKey.Group == nil || apiKey.Group.Platform != service.PlatformGrok { + h.errorResponse(c, http.StatusNotFound, "not_found_error", "Voice API is not supported for this platform") + return + } + if !h.ensureResponsesDependencies(c, nil) { + return + } + subscription, _ := middleware2.GetSubscriptionFromContext(c) + if err := h.billingCacheService.CheckBillingEligibility(c.Request.Context(), apiKey.User, apiKey, apiKey.Group, subscription, service.QuotaPlatform(c.Request.Context(), apiKey)); err != nil { + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + h.errorResponse(c, status, code, message) + return + } + + body, err := readGrokVoiceGatewayBody(c) + if err != nil { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error()) + return + } + if endpoint == "tts" { + subject, _ := middleware2.GetAuthSubjectFromContext(c) + reqLog := requestLogger(c, "handler.openai_gateway.grok_voice", zap.String("endpoint", endpoint)) + // TTS bodies use {"input":"..."} (and variants). Normalize to chat messages so + // content moderation extractors see the spoken text. + auditBody := body + if input := extractGrokTTSInputText(body); input != "" { + if b, err := json.Marshal(map[string]any{ + "messages": []map[string]any{{"role": "user", "content": input}}, + }); err == nil { + auditBody = b + } + } + if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, "grok-4.5", auditBody); decision != nil && !decision.AllowNextStage { + h.openAISecurityAuditError(c, decision) + return + } + } + contentType := c.GetHeader("Content-Type") + if strings.TrimSpace(contentType) == "" { + contentType = "application/json" + } + + failed := map[int64]struct{}{} + var last *service.UpstreamFailoverError + reqLog := requestLogger(c, "handler.openai_gateway.grok_voice", zap.String("endpoint", endpoint)) + selectionModel := "grok-4.5" + + for attempts := 0; attempts < 4; attempts++ { + selection, _, selectErr := h.gatewayService.SelectAccountWithSchedulerForCapability( + c.Request.Context(), + apiKey.GroupID, + "", + "", + selectionModel, + failed, + service.OpenAIUpstreamTransportHTTPSSE, + service.OpenAIEndpointCapabilityChatCompletions, + false, + false, + false, + service.PlatformGrok, + ) + if selectErr != nil || selection == nil || selection.Account == nil { + if last != nil { + h.handleFailoverExhausted(c, last, false) + } else { + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "No available Grok accounts") + } + return + } + account := selection.Account + var started bool + release, status := h.acquireResponsesAccountSlot(c, apiKey.GroupID, "", selection, false, &started, reqLog) + if status == openAISlotAcquireProfitVetoed { + failed[account.ID] = struct{}{} + continue + } + if status != openAISlotAcquireOK { + // Failed already wrote error response (or transient reject). + if status == openAISlotAcquireFailed && len(failed) == 0 { + // Slot path wrote the response; stop. + return + } + failed[account.ID] = struct{}{} + continue + } + result, forwardErr := func() (*service.OpenAIForwardResult, error) { + defer release() + return h.gatewayService.ForwardGrokVoice(c.Request.Context(), c, account, endpoint, body, contentType) + }() + if forwardErr == nil { + h.recordGrokVoiceUsage(c, apiKey, account, subscription, endpoint, body, result) + return + } + var failoverErr *service.UpstreamFailoverError + if errors.As(forwardErr, &failoverErr) && failoverErr.ShouldRetryNextAccount() { + failed[account.ID] = struct{}{} + last = failoverErr + continue + } + // Non-failover errors: handleGrokMediaErrorResponse / transport already wrote response. + return + } + if last != nil { + h.handleFailoverExhausted(c, last, false) + } +} + +// recordGrokVoiceUsage bills TTS/STT/realtime via group audio prices when AudioUsage is set. +func (h *OpenAIGatewayHandler) recordGrokVoiceUsage( + c *gin.Context, + apiKey *service.APIKey, + account *service.Account, + subscription *service.UserSubscription, + endpoint string, + body []byte, + result *service.OpenAIForwardResult, +) { + if h == nil || c == nil || apiKey == nil || account == nil || result == nil { + return + } + if result.AudioUsage == nil { + return + } + // Ensure forced durable request ids even if callers forget (realtime/tts/stt money path). + if mode := strings.TrimSpace(result.AudioUsage.Mode); mode == "realtime" { + result.RequestID = service.StableGrokRealtimeBillingRequestID(result.RequestID) + } else { + result.RequestID = service.StableGrokAudioBillingRequestID(result.RequestID) + } + userAgent := c.GetHeader("User-Agent") + clientIP := ip.GetClientIP(c) + sessionID := service.ExtractClientSessionID(c) + requestPayloadHash := service.HashUsageRequestPayload(body) + if requestPayloadHash == "" { + requestPayloadHash = service.HashUsageRequestPayload([]byte(endpoint)) + } + inboundEndpoint := GetInboundEndpoint(c) + upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform) + quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) + model := strings.TrimSpace(result.Model) + if model == "" { + model = endpoint + } + + h.submitMandatoryUsageRecordTask(c.Request.Context(), 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, + RequestPayloadHash: requestPayloadHash, + APIKeyService: h.apiKeyService, + QuotaPlatform: quotaPlatform, + SessionID: sessionID, + ChannelUsageFields: clientRequestedUsageFields(c, service.ChannelMappingResult{}, model, result.UpstreamModel), + }); err != nil { + logger.L().With( + zap.String("component", "handler.openai_gateway.grok_voice"), + zap.Int64("user_id", apiKey.User.ID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + zap.String("endpoint", endpoint), + zap.Int64("account_id", account.ID), + ).Error("grok_voice.record_usage_failed", zap.Error(err)) + } + }) +} + +func readGrokVoiceGatewayBody(c *gin.Context) ([]byte, error) { + if c == nil || c.Request == nil { + return nil, errors.New("request body is required") + } + if c.Request.Body == nil { + if c.Request.Method == http.MethodGet || c.Request.Method == http.MethodDelete { + return nil, nil + } + return nil, errors.New("request body is required") + } + return io.ReadAll(c.Request.Body) +} + +// extractGrokTTSInputText pulls the primary spoken text from a TTS JSON body. +func extractGrokTTSInputText(body []byte) string { + if len(body) == 0 { + return "" + } + var payload map[string]any + if err := json.Unmarshal(body, &payload); err != nil { + return "" + } + for _, key := range []string{"input", "text", "prompt"} { + if v, ok := payload[key]; ok { + if s, ok := v.(string); ok { + return strings.TrimSpace(s) + } + } + } + return "" +} diff --git a/backend/internal/handler/grok_audio_billing_test.go b/backend/internal/handler/grok_audio_billing_test.go new file mode 100644 index 0000000000..6b8a6c81be --- /dev/null +++ b/backend/internal/handler/grok_audio_billing_test.go @@ -0,0 +1,25 @@ +//go:build unit + +package handler + +import ( + "testing" + + coderws "github.com/coder/websocket" +) + +func TestIsExpectedGrokRealtimeClose(t *testing.T) { + for _, status := range []coderws.StatusCode{ + coderws.StatusNormalClosure, + coderws.StatusGoingAway, + coderws.StatusNoStatusRcvd, + coderws.StatusAbnormalClosure, + } { + if !isExpectedGrokRealtimeClose(coderws.CloseError{Code: status}) { + t.Fatalf("status %v should be treated as an expected session close", status) + } + } + if isExpectedGrokRealtimeClose(coderws.CloseError{Code: coderws.StatusPolicyViolation}) { + t.Fatal("policy violations must not be treated as billable normal closes") + } +} diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index e11bf7f08c..3cef97ed3f 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -188,6 +188,10 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. var oauth429FailoverState service.OpenAIOAuth429FailoverState mediaEligibilityRejected := false switchCount := 0 + videoCreateStartedAt := "" + if isGrokVideoCreateEndpoint(endpoint) { + videoCreateStartedAt = service.GrokVideoPendingCreatedAtNow() + } maxAccountSwitches := h.maxAccountSwitches if maxAccountSwitches <= 0 { maxAccountSwitches = 3 @@ -406,7 +410,7 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. } h.gatewayService.ReportOpenAIAccountScheduleResult(account.ID, grokMediaScheduleModel(account, routingModel, result), true, nil) - if endpoint.IsGenerationRequest() && strings.TrimSpace(result.ResponseID) != "" { + if isGrokVideoCreateEndpoint(endpoint) && strings.TrimSpace(result.ResponseID) != "" { if err := h.gatewayService.BindGrokMediaVideoRequestAccount( requestCtx, apiKey.GroupID, result.ResponseID, subject.UserID, apiKey.ID, account.ID, ); err != nil { @@ -416,8 +420,44 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service. zap.Error(err), ) } + // Defer billing until status polling observes video.url. Persist create-time + // model/duration/resolution so status can still price if upstream omits them. + // Retry once: missing pending causes silent underpricing (status omits resolution). + pending := service.GrokVideoPendingBilling{ + Model: requestModel, + BillingModel: firstNonEmptyString(result.BillingModel, requestModel), + UpstreamModel: result.UpstreamModel, + VideoResolution: result.VideoResolution, + VideoDurationSeconds: result.VideoDurationSeconds, + OriginalModel: clientRequestedModel(c, requestModel), + // Wall-clock start for usage duration_ms: create accepted → first done discovery. + CreatedAt: videoCreateStartedAt, + } + if err := h.gatewayService.StoreGrokVideoPendingBilling(requestCtx, result.ResponseID, subject.UserID, apiKey.ID, pending); err != nil { + reqLog.Warn("grok_media.store_video_pending_billing_failed_retrying", + zap.Int64("account_id", account.ID), + zap.String("request_id", result.ResponseID), + zap.Error(err), + ) + if err2 := h.gatewayService.StoreGrokVideoPendingBilling(requestCtx, result.ResponseID, subject.UserID, apiKey.ID, pending); err2 != nil { + // Response body may already be committed; completion path will fail-closed + // when pending is still missing and status cannot price duration. + reqLog.Error("grok_media.store_video_pending_billing_failed", + zap.Int64("account_id", account.ID), + zap.String("request_id", result.ResponseID), + zap.Error(err2), + ) + } + } } - if shouldRecordGrokMediaUsage(endpoint, requestModel) { + // Status poll OR content download can observe official done+video.url. + // Both paths share the same claim key so the customer is charged once. + if endpoint == service.GrokMediaEndpointVideoStatus || endpoint == service.GrokMediaEndpointVideoContent { + taskID := strings.TrimSpace(requestID) + if billResult := prepareGrokVideoCompletionBilling(requestCtx, h, reqLog, apiKey, subject, taskID, result); billResult != nil { + recordGrokMediaUsage(c, h, reqLog, apiKey, subject, subscription, account, billResult, billResult.Model, body, taskID) + } + } else if shouldRecordGrokMediaUsage(endpoint, requestModel, result) { recordGrokMediaUsage(c, h, reqLog, apiKey, subject, subscription, account, result, requestModel, body, requestID) } reqLog.Debug("grok_media.request_completed", @@ -459,8 +499,147 @@ func grokMediaScheduleModel(account *service.Account, routingModel string, resul return account.GetMappedModel(routingModel) } -func shouldRecordGrokMediaUsage(endpoint service.GrokMediaEndpoint, requestModel string) bool { - return endpoint.IsGenerationRequest() && strings.TrimSpace(requestModel) != "" +func isGrokVideoCreateEndpoint(endpoint service.GrokMediaEndpoint) bool { + switch endpoint { + case service.GrokMediaEndpointVideosGenerations, + service.GrokMediaEndpointVideosEdits, + service.GrokMediaEndpointVideosExtensions: + return true + default: + return false + } +} + +// shouldRecordGrokMediaUsage gates usage writes for immediate (image) generation. +// Async video create never bills here — status polling does on official +// status=done with video.url (docs.x.ai Video Generation). +// Status/content polls, empty model, and failed generations with zero billable +// image units never bill via this helper. +func shouldRecordGrokMediaUsage(endpoint service.GrokMediaEndpoint, requestModel string, result *service.OpenAIForwardResult) bool { + if result == nil { + return false + } + if isGrokVideoCreateEndpoint(endpoint) || endpoint.IsVideoLookupRequest() { + return false + } + if !endpoint.IsGenerationRequest() || strings.TrimSpace(requestModel) == "" { + return false + } + return result.ImageCount > 0 +} + +// prepareGrokVideoCompletionBilling claims one-shot billing for official done+video.url +// observations (status poll or content download). Duration/model prefer status body; +// resolution uses create-time request (status response does not document resolution). +func prepareGrokVideoCompletionBilling( + ctx context.Context, + h *OpenAIGatewayHandler, + reqLog *zap.Logger, + apiKey *service.APIKey, + subject middleware2.AuthSubject, + taskRequestID string, + statusResult *service.OpenAIForwardResult, +) *service.OpenAIForwardResult { + if h == nil || h.gatewayService == nil || apiKey == nil || statusResult == nil { + return nil + } + // Forward already set VideoCount only when status=done && video.url (official). + if statusResult.VideoCount <= 0 { + return nil + } + taskRequestID = strings.TrimSpace(firstNonEmptyString(taskRequestID, statusResult.ResponseID)) + if taskRequestID == "" { + return nil + } + // Load create-time snapshot before claim so we can fail-closed without burning the claim + // when Redis lost pending and status cannot price the job. + pending, loadErr := h.gatewayService.LoadGrokVideoPendingBilling(ctx, taskRequestID, subject.UserID, apiKey.ID) + if loadErr != nil { + reqLog.Warn("grok_media.video_pending_billing_load_failed", zap.String("request_id", taskRequestID), zap.Error(loadErr)) + } + if pending == nil { + // Status omits resolution; without pending we would silently default to 480p and underbill. + // Allow billing only when official status carries duration (still may default resolution). + if statusResult.VideoDurationSeconds <= 0 { + reqLog.Error("grok_media.video_billing_skipped_missing_pending", + zap.String("request_id", taskRequestID), + zap.String("reason", "no create-time snapshot and status has no video.duration"), + ) + return nil + } + reqLog.Error("grok_media.video_billing_without_pending", + zap.String("request_id", taskRequestID), + zap.Int("status_duration_seconds", statusResult.VideoDurationSeconds), + zap.String("note", "resolution falls back to default 480p; investigate pending store failures"), + ) + } + claimed, err := h.gatewayService.ClaimGrokVideoBilling(ctx, taskRequestID, subject.UserID, apiKey.ID) + if err != nil { + reqLog.Warn("grok_media.video_billing_claim_failed", zap.String("request_id", taskRequestID), zap.Error(err)) + return nil + } + if !claimed { + reqLog.Debug("grok_media.video_billing_already_claimed", zap.String("request_id", taskRequestID)) + return nil + } + // Re-merge with pending: resolution is request-only; model/duration fill gaps. + merged := *statusResult + if pending != nil { + if strings.TrimSpace(merged.Model) == "" { + merged.Model = firstNonEmptyString(pending.BillingModel, pending.Model, pending.OriginalModel) + } + if strings.TrimSpace(merged.BillingModel) == "" { + merged.BillingModel = firstNonEmptyString(pending.BillingModel, pending.Model, merged.Model) + } + if strings.TrimSpace(merged.UpstreamModel) == "" { + merged.UpstreamModel = pending.UpstreamModel + } + // Official status omits resolution — always prefer create request. + if strings.TrimSpace(pending.VideoResolution) != "" { + merged.VideoResolution = pending.VideoResolution + } + if merged.VideoDurationSeconds <= 0 { + merged.VideoDurationSeconds = pending.VideoDurationSeconds + } + if strings.TrimSpace(merged.ResponseID) == "" { + merged.ResponseID = taskRequestID + } + } + if strings.TrimSpace(merged.Model) == "" { + merged.Model = "grok-imagine-video" + } + if strings.TrimSpace(merged.BillingModel) == "" { + merged.BillingModel = merged.Model + } + // Always force durable task id so usage_billing_dedup survives multi-poll + + // context-local request ids (do not prefer empty-only fill). + merged.RequestID = service.StableGrokVideoBillingRequestID(firstNonEmptyString(merged.ResponseID, taskRequestID)) + merged.ResponseID = firstNonEmptyString(merged.ResponseID, taskRequestID) + merged.VideoCount = 1 + // Pure video: do not keep legacy ImageCount (avoids image-path heuristics). + merged.ImageCount = 0 + // Official default resolution is 480p when the create request omitted it. + merged.VideoResolution = service.NormalizeVideoBillingResolutionOrDefault(merged.VideoResolution) + // Official default duration is 8s when neither status nor create provided it. + merged.VideoDurationSeconds = service.NormalizeVideoBillingDurationSecondsOrDefault(merged.VideoDurationSeconds) + // E2E latency for async video: create accept → this discovery of done+url. + // Bill on discovery (status/content), not after further client polls; duration + // must not be only the single discovery hop (~hundreds of ms). + if pending != nil { + if e2e := service.GrokVideoE2EDuration(pending.CreatedAt, time.Now()); e2e > 0 { + merged.Duration = e2e + } + } + return &merged +} + +func firstNonEmptyString(values ...string) string { + for _, v := range values { + if s := strings.TrimSpace(v); s != "" { + return s + } + } + return "" } func recordGrokMediaUsage( @@ -493,6 +672,18 @@ func recordGrokMediaUsage( OriginalModel: clientRequestedModel(c, requestModel), ChannelMappedModel: requestModel, } + // Async video: force durable task request id and release claim if billing fails. + videoTaskID := "" + if result != nil && result.VideoCount > 0 { + videoTaskID = strings.TrimSpace(firstNonEmptyString(requestID, result.ResponseID)) + if stable := service.StableGrokVideoBillingRequestID(firstNonEmptyString(result.ResponseID, requestID)); stable != "" { + result.RequestID = stable + } + // Prefer task id hash for payload fingerprint stability across status/content. + if len(body) == 0 && videoTaskID != "" { + payloadForHash = []byte(videoTaskID) + } + } h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) { if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{ Result: result, @@ -510,6 +701,14 @@ func recordGrokMediaUsage( SessionID: sessionID, ChannelUsageFields: channelUsageFields, }); err != nil { + if videoTaskID != "" { + if releaseErr := h.gatewayService.ReleaseGrokVideoBilling(ctx, videoTaskID, subject.UserID, apiKey.ID); releaseErr != nil { + reqLog.Warn("grok_media.video_billing_claim_release_failed", + zap.String("request_id", videoTaskID), + zap.Error(releaseErr), + ) + } + } logger.L().With( zap.String("component", "handler.openai_gateway.grok_media"), zap.Int64("user_id", subject.UserID), diff --git a/backend/internal/handler/grok_media_test.go b/backend/internal/handler/grok_media_test.go index fbd2af820e..5d5deec22e 100644 --- a/backend/internal/handler/grok_media_test.go +++ b/backend/internal/handler/grok_media_test.go @@ -3,6 +3,7 @@ package handler import ( "context" "errors" + "strings" "testing" "github.com/Wei-Shaw/sub2api/internal/service" @@ -41,13 +42,13 @@ func TestShouldRecordGrokMediaUsage(t *testing.T) { want: true, }, { - name: "video generation records usage", + name: "video generation defers usage until status", endpoint: service.GrokMediaEndpointVideosGenerations, model: "grok-imagine-video-1.5", - want: true, + want: false, }, { - name: "video status skips empty model usage", + name: "video status skips immediate helper (status path claims separately)", endpoint: service.GrokMediaEndpointVideoStatus, model: "", want: false, @@ -68,7 +69,18 @@ func TestShouldRecordGrokMediaUsage(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - require.Equal(t, tt.want, shouldRecordGrokMediaUsage(tt.endpoint, tt.model)) + // Nil result must never bill. + require.False(t, shouldRecordGrokMediaUsage(tt.endpoint, tt.model, nil)) + // Immediate helper only bills image generation (async video bills on status). + result := &service.OpenAIForwardResult{ImageCount: 1, VideoCount: 0} + if tt.endpoint.IsGenerationRequest() && !isGrokVideoCreateEndpoint(tt.endpoint) && strings.TrimSpace(tt.model) != "" { + require.Equal(t, tt.want, shouldRecordGrokMediaUsage(tt.endpoint, tt.model, result)) + } else { + require.False(t, shouldRecordGrokMediaUsage(tt.endpoint, tt.model, result)) + } + // Zero billable units never bill even for generation + model. + empty := &service.OpenAIForwardResult{} + require.False(t, shouldRecordGrokMediaUsage(tt.endpoint, tt.model, empty)) }) } } diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index b787d0f976..723bf901ce 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -2415,7 +2415,9 @@ func (h *OpenAIGatewayHandler) submitUsageRecordTask(parent context.Context, tas } func (h *OpenAIGatewayHandler) submitOpenAIUsageRecordTask(parent context.Context, result *service.OpenAIForwardResult, task service.UsageRecordTask) { - if result != nil && result.ImageCount > 0 { + // Money-critical bills never drop on pool overflow: media, search surcharge, voice. + if result != nil && (result.ImageCount > 0 || result.VideoCount > 0 || + result.SearchCount > 0 || result.WebSearchCalls > 0 || result.AudioUsage != nil) { h.submitMandatoryUsageRecordTask(parent, task) return } diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index 543cda1e55..bf3061d5e2 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -17,6 +17,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" pkghttputil "github.com/Wei-Shaw/sub2api/internal/pkg/httputil" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" coderws "github.com/coder/websocket" @@ -643,6 +644,9 @@ func TestResolveOpenAIMessagesDispatchMappedModel(t *testing.T) { }) t.Run("grok_group_maps_claude_cli_model_to_grok_default", func(t *testing.T) { + original := xai.RuntimeModelMappingOptions() + t.Cleanup(func() { xai.SetRuntimeModelMappingOptions(original) }) + xai.SetRuntimeModelMappingOptions(xai.ModelMappingOptions{EnableCrossClientMap: true}) apiKey := &service.APIKey{ Group: &service.Group{ Platform: service.PlatformGrok, diff --git a/backend/internal/handler/openai_images.go b/backend/internal/handler/openai_images.go index d5d16d1e71..da41bf6aa3 100644 --- a/backend/internal/handler/openai_images.go +++ b/backend/internal/handler/openai_images.go @@ -144,7 +144,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) { } sessionHash := h.gatewayService.GenerateExplicitSessionHash(c, body) - requestCtx := service.WithOpenAIImageGenerationIntent(c.Request.Context()) + requestCtx := service.WithOpenAIImagesEndpoint(service.WithOpenAIImageGenerationIntent(c.Request.Context())) maxAccountSwitches := h.maxAccountSwitches switchCount := 0 diff --git a/backend/internal/handler/usage_handler_request_type_test.go b/backend/internal/handler/usage_handler_request_type_test.go index 8dc9a8b442..1e480d2e34 100644 --- a/backend/internal/handler/usage_handler_request_type_test.go +++ b/backend/internal/handler/usage_handler_request_type_test.go @@ -224,6 +224,8 @@ func TestUserUsageListKeepsUserBillingAndIPWithoutAdminCostFields(t *testing.T) require.NotContains(t, body, "account_rate_multiplier") require.NotContains(t, body, "account_stats_cost") require.NotContains(t, body, "upstream_model") + require.NotContains(t, body, "upstream_response_model") + require.NotContains(t, body, "upstream_model_mismatch") require.NotContains(t, body, "billing_tier") require.NotContains(t, body, "channel_id") require.NotContains(t, body, `"account":`) diff --git a/backend/internal/handler/usage_record_submit_task_test.go b/backend/internal/handler/usage_record_submit_task_test.go index ebe5c3df82..66b3aca0d3 100644 --- a/backend/internal/handler/usage_record_submit_task_test.go +++ b/backend/internal/handler/usage_record_submit_task_test.go @@ -189,3 +189,33 @@ func TestOpenAIGatewayHandlerSubmitOpenAIUsageRecordTask_ImageResultUsesMandator require.True(t, called.Load(), "image usage task must be mandatory when async submit is dropped") } + +func TestOpenAIGatewayHandlerSubmitOpenAIUsageRecordTask_SearchCountUsesMandatoryFallback(t *testing.T) { + pool := service.NewUsageRecordWorkerPoolWithOptions(service.UsageRecordWorkerPoolOptions{ + WorkerCount: 1, + QueueSize: 1, + TaskTimeout: time.Second, + OverflowPolicy: "drop", + OverflowSamplePercent: 0, + AutoScaleEnabled: false, + }) + t.Cleanup(pool.Stop) + h := &OpenAIGatewayHandler{usageRecordWorkerPool: pool} + + block := make(chan struct{}) + release := make(chan struct{}) + pool.Submit(func(ctx context.Context) { + close(block) + <-release + }) + <-block + pool.Submit(func(ctx context.Context) {}) + + var called atomic.Bool + h.submitOpenAIUsageRecordTask(context.Background(), &service.OpenAIForwardResult{SearchCount: 3}, func(ctx context.Context) { + called.Store(true) + }) + close(release) + + require.True(t, called.Load(), "search surcharge usage task must be mandatory when async submit is dropped") +} diff --git a/backend/internal/pkg/antigravity/claude_types.go b/backend/internal/pkg/antigravity/claude_types.go index d8732d26ed..c0335c4156 100644 --- a/backend/internal/pkg/antigravity/claude_types.go +++ b/backend/internal/pkg/antigravity/claude_types.go @@ -178,6 +178,11 @@ var geminiModels = []modelDef{ {ID: "gemini-3.1-pro-high", DisplayName: "Gemini 3.1 Pro High", CreatedAt: "2026-02-19T00:00:00Z", IsReasoning: true}, {ID: "gemini-3.1-flash-image", DisplayName: "Gemini 3.1 Flash Image", CreatedAt: "2026-02-19T00:00:00Z"}, {ID: "gemini-3.1-flash-image-preview", DisplayName: "Gemini 3.1 Flash Image Preview", CreatedAt: "2026-02-19T00:00:00Z"}, + {ID: "gemini-3.6-flash", DisplayName: "Gemini 3.6 Flash", CreatedAt: "2026-07-21T00:00:00Z"}, + {ID: "gemini-3.6-flash-high", DisplayName: "Gemini 3.6 Flash High", CreatedAt: "2026-07-21T00:00:00Z", IsReasoning: true}, + {ID: "gemini-3.6-flash-low", DisplayName: "Gemini 3.6 Flash Low", CreatedAt: "2026-07-21T00:00:00Z", IsReasoning: true}, + {ID: "gemini-3.6-flash-medium", DisplayName: "Gemini 3.6 Flash Medium", CreatedAt: "2026-07-21T00:00:00Z", IsReasoning: true}, + {ID: "gemini-3.6-flash-tiered", DisplayName: "Gemini 3.6 Flash", CreatedAt: "2026-07-21T00:00:00Z", IsReasoning: true}, {ID: "gemini-3-pro-preview", DisplayName: "Gemini 3 Pro Preview", CreatedAt: "2025-06-01T00:00:00Z", IsReasoning: true}, {ID: "gemini-3-pro-image", DisplayName: "Gemini 3 Pro Image", CreatedAt: "2025-06-01T00:00:00Z"}, } diff --git a/backend/internal/pkg/antigravity/claude_types_test.go b/backend/internal/pkg/antigravity/claude_types_test.go index bb45c2f0c3..65c8078c33 100644 --- a/backend/internal/pkg/antigravity/claude_types_test.go +++ b/backend/internal/pkg/antigravity/claude_types_test.go @@ -20,6 +20,11 @@ func TestDefaultModels_ContainsNewAndLegacyImageModels(t *testing.T) { "gemini-3.1-flash-image", "gemini-3.1-flash-image-preview", "gemini-3-pro-image", // legacy compatibility + "gemini-3.6-flash", + "gemini-3.6-flash-high", + "gemini-3.6-flash-low", + "gemini-3.6-flash-medium", + "gemini-3.6-flash-tiered", } for _, id := range requiredIDs { diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic_invalid_blocks_test.go b/backend/internal/pkg/apicompat/responses_to_anthropic_invalid_blocks_test.go new file mode 100644 index 0000000000..ed410f6851 --- /dev/null +++ b/backend/internal/pkg/apicompat/responses_to_anthropic_invalid_blocks_test.go @@ -0,0 +1,207 @@ +package apicompat + +import ( + "encoding/json" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +// anthropicInboundBlockTypes 是 Anthropic Messages 请求体里合法的 content block +// 类型。转换结果只能落在这个集合内——发出集合外的类型,上游一律回 +// 400 "Request body format invalid"(见 issue #5329)。 +var anthropicInboundBlockTypes = map[string]bool{ + "text": true, + "image": true, + "document": true, + "tool_use": true, + "tool_result": true, + "thinking": true, + "redacted_thinking": true, +} + +func responsesToAnthropicMessages(t *testing.T, input string) []AnthropicMessage { + t.Helper() + var req ResponsesRequest + require.NoError(t, json.Unmarshal([]byte(`{"model":"glm-5.2","input":`+input+`}`), &req)) + out, err := ResponsesToAnthropicRequest(&req) + require.NoError(t, err) + return out.Messages +} + +// requireAnthropicMessagesAreSendable 断言消息序列不含 Anthropic 会拒收的形态: +// 未知 block 类型、空内容消息、纯空白 text 块。 +func requireAnthropicMessagesAreSendable(t *testing.T, messages []AnthropicMessage) { + t.Helper() + for i, m := range messages { + raw := strings.TrimSpace(string(m.Content)) + require.NotContains(t, []string{"", "null", `""`, "[]"}, raw, + "messages[%d] 内容为空,Anthropic 拒收空内容消息", i) + + var s string + if err := json.Unmarshal(m.Content, &s); err == nil { + require.NotEmpty(t, strings.TrimSpace(s), "messages[%d] 字符串内容不能全为空白", i) + continue + } + blocks := parseContentBlocks(m.Content) + require.NotEmpty(t, blocks, "messages[%d] 解析不出任何 block", i) + for j, b := range blocks { + require.True(t, anthropicInboundBlockTypes[b.Type], + "messages[%d].content[%d] 是 Anthropic 不认识的 block 类型 %q", i, j, b.Type) + if b.Type == "text" { + require.NotEmpty(t, strings.TrimSpace(b.Text), + "messages[%d].content[%d] 是空白 text 块,Anthropic 拒收", i, j) + } + } + } +} + +// issue #5329:工具执行后的下一轮,Codex 会把 reasoning item 一起回放。 +// 该 item 带 content 数组时,reasoning_text 块以前会被原样塞进 Anthropic 请求体。 +func TestResponsesToAnthropic_ReasoningItemWithContentIsDropped(t *testing.T) { + messages := responsesToAnthropicMessages(t, `[ + {"type":"message","role":"user","content":[{"type":"input_text","text":"run a shell command"}]}, + {"type":"reasoning","id":"rs_1","summary":[],"content":[{"type":"reasoning_text","text":"let me think"}]} + ]`) + + requireAnthropicMessagesAreSendable(t, messages) + require.Len(t, messages, 1) + require.NotContains(t, string(messages[0].Content), "reasoning_text") + require.NotContains(t, string(messages[0].Content), "let me think") +} + +// Codex 的常见 reasoning 形态(只有 summary + encrypted_content)本来就会被丢弃, +// 这条守卫确保行为没有被改变。 +func TestResponsesToAnthropic_ReasoningItemSummaryOnlyStillDropped(t *testing.T) { + messages := responsesToAnthropicMessages(t, `[ + {"type":"message","role":"user","content":[{"type":"input_text","text":"hi"}]}, + {"type":"reasoning","id":"rs_1","summary":[{"type":"summary_text","text":"s"}],"encrypted_content":"gAAAA"} + ]`) + + requireAnthropicMessagesAreSendable(t, messages) + require.Len(t, messages, 1) + require.NotContains(t, string(messages[0].Content), "gAAAA") +} + +// 未知 item type 的 content 以前会被逐字透传,把 Responses 专有分片带进上游请求。 +func TestResponsesToAnthropic_UnknownItemTypeContentIsSanitized(t *testing.T) { + messages := responsesToAnthropicMessages(t, `[ + {"type":"web_search_call","id":"ws_1","content":[{"type":"web_search_result","text":"payload"}]} + ]`) + + requireAnthropicMessagesAreSendable(t, messages) + require.Empty(t, messages, "整条内容都无法映射时不应发出消息") +} + +// 未知 item type 里夹带的可识别文本仍然保留,不做无谓丢弃。 +func TestResponsesToAnthropic_UnknownItemTypeKeepsRecognizableText(t *testing.T) { + messages := responsesToAnthropicMessages(t, `[ + {"type":"some_future_item","content":[ + {"type":"input_text","text":"keep me"}, + {"type":"reasoning_text","text":"drop me"} + ]} + ]`) + + requireAnthropicMessagesAreSendable(t, messages) + require.Len(t, messages, 1) + require.Contains(t, string(messages[0].Content), "keep me") + require.NotContains(t, string(messages[0].Content), "drop me") +} + +// user 消息的分片全部不可识别时,以前会退化成 content:"",Anthropic 拒收空内容消息。 +func TestResponsesToAnthropic_UserMessageWithOnlyUnknownPartsIsDropped(t *testing.T) { + messages := responsesToAnthropicMessages(t, `[ + {"type":"message","role":"user","content":[{"type":"input_file","file_id":"file_1"}]} + ]`) + + requireAnthropicMessagesAreSendable(t, messages) + require.Empty(t, messages) +} + +// assistant 侧同理:以前会退化成单个空 text 块,Anthropic 同样拒收。 +func TestResponsesToAnthropic_AssistantMessageWithOnlyUnknownPartsIsDropped(t *testing.T) { + messages := responsesToAnthropicMessages(t, `[ + {"type":"message","role":"user","content":[{"type":"input_text","text":"hi"}]}, + {"type":"message","role":"assistant","content":[{"type":"refusal","refusal":"no"}]} + ]`) + + requireAnthropicMessagesAreSendable(t, messages) + require.Len(t, messages, 1) + require.Equal(t, "user", messages[0].Role) +} + +// 完整的 Codex 工具续接回放:tool_use / tool_result 配对必须保持不变, +// 同时整个序列满足可发送不变式。 +func TestResponsesToAnthropic_CodexToolRoundStaysIntactAndSendable(t *testing.T) { + messages := responsesToAnthropicMessages(t, `[ + {"type":"message","role":"user","content":[{"type":"input_text","text":"run ls"}]}, + {"type":"reasoning","id":"rs_1","summary":[],"content":[{"type":"reasoning_text","text":"plan"}],"encrypted_content":"gAAAA"}, + {"type":"function_call","id":"fc_1","call_id":"call_1","name":"shell","arguments":"{\"cmd\":\"ls\"}"}, + {"type":"function_call_output","call_id":"call_1","output":"file1"}, + {"type":"message","role":"user","content":[{"type":"input_text","text":"continue"}]} + ]`) + + requireAnthropicMessagesAreSendable(t, messages) + + var sawToolUse, sawToolResult bool + for _, m := range messages { + for _, b := range parseContentBlocks(m.Content) { + switch b.Type { + case "tool_use": + sawToolUse = true + require.Equal(t, "call_1", b.ID) + require.Equal(t, "shell", b.Name) + case "tool_result": + sawToolResult = true + require.Equal(t, "call_1", b.ToolUseID) + } + } + } + require.True(t, sawToolUse, "function_call 必须转成 tool_use") + require.True(t, sawToolResult, "function_call_output 必须转成 tool_result") + require.NotContains(t, string(mustMarshal(t, messages)), "reasoning_text") + require.NotContains(t, string(mustMarshal(t, messages)), "gAAAA") +} + +func mustMarshal(t *testing.T, v any) []byte { + t.Helper() + b, err := json.Marshal(v) + require.NoError(t, err) + return b +} + +func TestAnthropicContentIsEmpty(t *testing.T) { + cases := []struct { + raw string + want bool + }{ + {``, true}, + {`""`, true}, + {`null`, true}, + {`[]`, true}, + {` [] `, true}, + {`"hi"`, false}, + {`[{"type":"text","text":"hi"}]`, false}, + } + for _, tc := range cases { + require.Equal(t, tc.want, anthropicContentIsEmpty(json.RawMessage(tc.raw)), "raw=%q", tc.raw) + } +} + +func TestAnthropicContentIsOnlyBlankText(t *testing.T) { + cases := []struct { + raw string + want bool + }{ + {`[{"type":"text","text":""}]`, true}, + {`[{"type":"text","text":" "}]`, true}, + {`[{"type":"text","text":""},{"type":"text","text":" "}]`, true}, + {`[{"type":"text","text":"hi"}]`, false}, + {`[{"type":"text","text":""},{"type":"image","source":{}}]`, false}, + {`[]`, false}, + } + for _, tc := range cases { + require.Equal(t, tc.want, anthropicContentIsOnlyBlankText(json.RawMessage(tc.raw)), "raw=%q", tc.raw) + } +} diff --git a/backend/internal/pkg/apicompat/responses_to_anthropic_request.go b/backend/internal/pkg/apicompat/responses_to_anthropic_request.go index 7a28b44abc..51ffe8baaf 100644 --- a/backend/internal/pkg/apicompat/responses_to_anthropic_request.go +++ b/backend/internal/pkg/apicompat/responses_to_anthropic_request.go @@ -164,11 +164,24 @@ func convertResponsesInputToAnthropic(instructions string, inputRaw json.RawMess Content: blockJSON, }) + case item.Type == "reasoning": + // Anthropic 无法摄入 OpenAI 的 reasoning:encrypted_content 是不透明的, + // 而 thinking 块的重放需要 Anthropic 自己签发的 signature,无法伪造。 + // Codex 常见形态(只带 summary + encrypted_content)本来就会被丢弃, + // 这里让带 content 数组的形态保持同样行为——否则 reasoning_text 块会被 + // 原样塞进 Anthropic 请求体,上游直接回 400。 + case item.Role == "user": content, err := convertResponsesUserToAnthropicContent(item.Content) if err != nil { return nil, nil, err } + // 内容里只有网关不认识的分片时,sanitize 会得到空串。Anthropic 拒收 + // 空内容消息("all messages must have non-empty content"),整条丢掉 + // 比发一条必然 400 的消息更可用。 + if anthropicContentIsEmpty(content) { + continue + } messages = append(messages, AnthropicMessage{ Role: "user", Content: content, @@ -179,19 +192,35 @@ func convertResponsesInputToAnthropic(instructions string, inputRaw json.RawMess if err != nil { return nil, nil, err } + // 同上:分片全不认识时会退化成单个空 text 块,而 Anthropic 拒收 + // 空文本块("text content blocks must contain non-whitespace text")。 + if anthropicContentIsEmpty(content) || anthropicContentIsOnlyBlankText(content) { + continue + } messages = append(messages, AnthropicMessage{ Role: "assistant", Content: content, }) default: - // Unknown role/type — attempt as user message - if item.Content != nil { - messages = append(messages, AnthropicMessage{ - Role: "user", - Content: item.Content, - }) + // 未知 role/type —— 尽量当作 user 消息保留其中的文本/图片。 + // 必须走与真实 user 消息同一套白名单转换:直接透传 item.Content 会把 + // Responses 专有的分片类型(reasoning_text、web_search_call 的载荷等) + // 原样发给 Anthropic,上游只会回 400 把整轮打挂。 + if item.Content == nil { + continue } + content, err := convertResponsesUserToAnthropicContent(item.Content) + if err != nil { + return nil, nil, err + } + if anthropicContentIsEmpty(content) { + continue + } + messages = append(messages, AnthropicMessage{ + Role: "user", + Content: content, + }) } } @@ -393,6 +422,32 @@ func extractTextFromContent(raw json.RawMessage) string { // convertResponsesUserToAnthropicContent converts a Responses user message // content field into Anthropic content blocks JSON. +// anthropicContentIsEmpty 判断转换结果是否为"空内容"。 +// convertResponsesUserToAnthropicContent 在没有任何可识别分片时返回 JSON 空串, +// 而 Anthropic 拒收空内容消息。 +func anthropicContentIsEmpty(content json.RawMessage) bool { + trimmed := strings.TrimSpace(string(content)) + switch trimmed { + case "", "null", `""`, "[]": + return true + } + return false +} + +// anthropicContentIsOnlyBlankText 判断内容是否只由空白 text 块组成。 +func anthropicContentIsOnlyBlankText(content json.RawMessage) bool { + blocks := parseContentBlocks(content) + if len(blocks) == 0 { + return false + } + for _, b := range blocks { + if b.Type != "text" || strings.TrimSpace(b.Text) != "" { + return false + } + } + return true +} + func convertResponsesUserToAnthropicContent(raw json.RawMessage) (json.RawMessage, error) { if len(raw) == 0 { return json.Marshal("") // empty string content diff --git a/backend/internal/pkg/ctxkey/ctxkey.go b/backend/internal/pkg/ctxkey/ctxkey.go index 2ec3e0fdbe..c34ef589a0 100644 --- a/backend/internal/pkg/ctxkey/ctxkey.go +++ b/backend/internal/pkg/ctxkey/ctxkey.go @@ -50,6 +50,12 @@ const ( // OpenAIImageGenerationIntent 标识 OpenAI 请求会触发生图能力(用于图片能力维度限流) OpenAIImageGenerationIntent Key = "ctx_openai_image_generation_intent" + // OpenAIImagesEndpoint 标识请求是从 /v1/images/* 入站的。 + // 与 OpenAIImageGenerationIntent 的区别:后者只表示"这次请求会生图", + // /v1/responses 带图片模型时也会置位;本 key 只在专用生图端点置位, + // 用于区分"用错端点"与"端点用对了但账号没能力"。 + OpenAIImagesEndpoint Key = "ctx_openai_images_endpoint" + // Group 认证后的分组信息,由 API Key 认证中间件设置 Group Key = "ctx_group" diff --git a/backend/internal/pkg/proxyutil/dialer.go b/backend/internal/pkg/proxyutil/dialer.go index e437cae342..3e135e3eb9 100644 --- a/backend/internal/pkg/proxyutil/dialer.go +++ b/backend/internal/pkg/proxyutil/dialer.go @@ -16,10 +16,29 @@ import ( "net/http" "net/url" "strings" + "time" "golang.org/x/net/proxy" ) +const ( + // socks5DialTimeout 限制到 SOCKS5 代理自身的 TCP 建连耗时。 + socks5DialTimeout = 10 * time.Second + // socks5DialKeepAlive 与 Go 默认 keepalive 探测间隔保持一致。 + socks5DialKeepAlive = 30 * time.Second +) + +// socks5ForwardDialer 是 SOCKS5 dialer 的底层拨号器。 +// +// proxy.FromURL 的默认 forward dialer 是 proxy.Direct(零值 net.Dialer,无超时), +// 代理地址不可达时会一直卡到内核 TCP 重传耗尽(Linux 约 130 秒)。SOCKS5 分支会 +// 覆盖 Transport.DialContext,因此调用方在 Transport 上设置的建连超时对这条路径 +// 无效,必须在这里补上。 +var socks5ForwardDialer = &net.Dialer{ + Timeout: socks5DialTimeout, + KeepAlive: socks5DialKeepAlive, +} + // ConfigureTransportProxy 根据代理 URL 配置 Transport // // 支持的协议: @@ -45,7 +64,7 @@ func ConfigureTransportProxy(transport *http.Transport, proxyURL *url.URL) error return nil case "socks5", "socks5h": - dialer, err := proxy.FromURL(proxyURL, proxy.Direct) + dialer, err := proxy.FromURL(proxyURL, socks5ForwardDialer) if err != nil { return fmt.Errorf("create socks5 dialer: %w", err) } diff --git a/backend/internal/pkg/proxyutil/dialer_timeout_test.go b/backend/internal/pkg/proxyutil/dialer_timeout_test.go new file mode 100644 index 0000000000..212bf0adac --- /dev/null +++ b/backend/internal/pkg/proxyutil/dialer_timeout_test.go @@ -0,0 +1,58 @@ +package proxyutil + +import ( + "context" + "errors" + "net" + "net/http" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +var errStub = errors.New("stub dial") + +// 回归:SOCKS5 分支覆盖了调用方在 Transport 上设置的 DialContext, +// 底层 forward dialer 必须自带建连超时。proxy.Direct 是零值 net.Dialer, +// 代理不可达时会一直卡到内核 TCP 重传耗尽(Linux 约 130 秒)。 +func TestSOCKS5ForwardDialerHasBoundedTimeout(t *testing.T) { + require.Greater(t, socks5ForwardDialer.Timeout, time.Duration(0)) + require.Equal(t, socks5DialTimeout, socks5ForwardDialer.Timeout) + require.Equal(t, socks5DialKeepAlive, socks5ForwardDialer.KeepAlive) +} + +func TestConfigureTransportProxySOCKS5SetsDialContext(t *testing.T) { + for _, scheme := range []string{"socks5", "socks5h"} { + t.Run(scheme, func(t *testing.T) { + proxyURL, err := url.Parse(scheme + "://127.0.0.1:1080") + require.NoError(t, err) + + transport := &http.Transport{} + require.NoError(t, ConfigureTransportProxy(transport, proxyURL)) + require.NotNil(t, transport.DialContext) + require.Nil(t, transport.Proxy, "SOCKS5 不应设置 Transport.Proxy") + }) + } +} + +// HTTP 代理走 Transport.Proxy,不得覆盖调用方设置的 DialContext。 +func TestConfigureTransportProxyHTTPPreservesDialContext(t *testing.T) { + proxyURL, err := url.Parse("http://127.0.0.1:8080") + require.NoError(t, err) + + called := false + transport := &http.Transport{} + transport.DialContext = func(_ context.Context, _, _ string) (net.Conn, error) { + called = true + return nil, errStub + } + + require.NoError(t, ConfigureTransportProxy(transport, proxyURL)) + require.NotNil(t, transport.Proxy) + require.NotNil(t, transport.DialContext) + + _, _ = transport.DialContext(context.Background(), "tcp", "127.0.0.1:1") + require.True(t, called, "HTTP 代理分支不应替换调用方的 DialContext") +} diff --git a/backend/internal/pkg/redissession/store.go b/backend/internal/pkg/redissession/store.go new file mode 100644 index 0000000000..3909fceecc --- /dev/null +++ b/backend/internal/pkg/redissession/store.go @@ -0,0 +1,125 @@ +// Package redissession provides a multi-instance OAuth session backend. +package redissession + +import ( + "context" + "encoding/json" + "errors" + "strings" + "time" + + "github.com/redis/go-redis/v9" +) + +var ErrNotConfigured = errors.New("redis session store not configured") + +// Store persists JSON sessions and single-use markers under one namespace. +type Store struct { + rdb *redis.Client + prefix string + ttl time.Duration +} + +func New(rdb *redis.Client, prefix string, ttl time.Duration) *Store { + if ttl <= 0 { + ttl = 30 * time.Minute + } + prefix = strings.TrimSpace(prefix) + if prefix == "" { + prefix = "oauth:session" + } + if !strings.HasSuffix(prefix, ":") { + prefix += ":" + } + return &Store{rdb: rdb, prefix: prefix, ttl: ttl} +} + +func (s *Store) dataKey(id string) string { return s.prefix + strings.TrimSpace(id) } +func (s *Store) usedKey(id string) string { return s.prefix + "used:" + strings.TrimSpace(id) } + +func (s *Store) Set(ctx context.Context, id string, value any) error { + if s == nil || s.rdb == nil { + return ErrNotConfigured + } + id = strings.TrimSpace(id) + if id == "" { + return errors.New("session id is required") + } + if ctx == nil { + ctx = context.Background() + } + raw, err := json.Marshal(value) + if err != nil { + return err + } + return s.rdb.Set(ctx, s.dataKey(id), raw, s.ttl).Err() +} + +func (s *Store) Get(ctx context.Context, id string, dest any) (bool, error) { + if s == nil || s.rdb == nil { + return false, ErrNotConfigured + } + id = strings.TrimSpace(id) + if id == "" { + return false, nil + } + if ctx == nil { + ctx = context.Background() + } + raw, err := s.rdb.Get(ctx, s.dataKey(id)).Bytes() + if errors.Is(err, redis.Nil) { + return false, nil + } + if err != nil { + return false, err + } + if err := json.Unmarshal(raw, dest); err != nil { + return false, err + } + return true, nil +} + +func (s *Store) Delete(ctx context.Context, id string) error { + if s == nil || s.rdb == nil { + return ErrNotConfigured + } + id = strings.TrimSpace(id) + if id == "" { + return nil + } + if ctx == nil { + ctx = context.Background() + } + return s.rdb.Del(ctx, s.dataKey(id), s.usedKey(id)).Err() +} + +// TryConsume returns true only for the first claim while the session exists. +func (s *Store) TryConsume(ctx context.Context, id string) (bool, error) { + if s == nil || s.rdb == nil { + return false, ErrNotConfigured + } + id = strings.TrimSpace(id) + if id == "" { + return false, nil + } + if ctx == nil { + ctx = context.Background() + } + ttl := s.ttl + if remaining, err := s.rdb.TTL(ctx, s.dataKey(id)).Result(); err == nil && remaining > 0 { + ttl = remaining + } + ok, err := s.rdb.SetNX(ctx, s.usedKey(id), "1", ttl).Result() + if err != nil || !ok { + return ok, err + } + exists, err := s.rdb.Exists(ctx, s.dataKey(id)).Result() + if err != nil { + return false, err + } + if exists == 0 { + _ = s.rdb.Del(ctx, s.usedKey(id)).Err() + return false, nil + } + return true, nil +} diff --git a/backend/internal/pkg/redissession/store_test.go b/backend/internal/pkg/redissession/store_test.go new file mode 100644 index 0000000000..ad5832fb84 --- /dev/null +++ b/backend/internal/pkg/redissession/store_test.go @@ -0,0 +1,40 @@ +//go:build unit + +package redissession + +import ( + "context" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func TestStoreRoundTripAndSingleUse(t *testing.T) { + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { _ = rdb.Close() }) + store := New(rdb, "oauth:test", time.Minute) + ctx := context.Background() + + require.NoError(t, store.Set(ctx, "sid", map[string]string{"state": "state"})) + var got map[string]string + ok, err := store.Get(ctx, "sid", &got) + require.NoError(t, err) + require.True(t, ok) + require.Equal(t, "state", got["state"]) + + ok, err = store.TryConsume(ctx, "sid") + require.NoError(t, err) + require.True(t, ok) + ok, err = store.TryConsume(ctx, "sid") + require.NoError(t, err) + require.False(t, ok) + + require.NoError(t, store.Delete(ctx, "sid")) + ok, err = store.Get(ctx, "sid", &got) + require.NoError(t, err) + require.False(t, ok) +} diff --git a/backend/internal/pkg/usagestats/usage_log_types.go b/backend/internal/pkg/usagestats/usage_log_types.go index 959907ca77..a100b2ccaa 100644 --- a/backend/internal/pkg/usagestats/usage_log_types.go +++ b/backend/internal/pkg/usagestats/usage_log_types.go @@ -274,13 +274,14 @@ type UsageLogFilters struct { RequestID string Model string // ModelFilterSource controls how Model is matched. Empty preserves raw usage_logs.model semantics. - ModelFilterSource string - RequestType *int16 - Stream *bool - BillingType *int8 - BillingMode string - StartTime *time.Time - EndTime *time.Time + ModelFilterSource string + RequestType *int16 + Stream *bool + BillingType *int8 + BillingMode string + UpstreamModelMismatch *bool + StartTime *time.Time + EndTime *time.Time // ExactTotal requests exact COUNT(*) for pagination. Default false for fast large-table paging. ExactTotal bool } diff --git a/backend/internal/pkg/xai/billing.go b/backend/internal/pkg/xai/billing.go index 3c620a9884..44e4772d44 100644 --- a/backend/internal/pkg/xai/billing.go +++ b/backend/internal/pkg/xai/billing.go @@ -20,7 +20,9 @@ const ( // one bump here covers OAuth traffic and billing probes together. // Keep in sync with https://x.ai/cli/stable. CLIClientVersion = "0.2.114" - CLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)" + // billingCLIUserAgent is the legacy pager/shell UA used by billing probes. + // Distinct from CLIUserAgent() in cli_identity.go (workspace-style UA). + billingCLIUserAgent = "grok-pager/" + CLIClientVersion + " grok-shell/" + CLIClientVersion + " (macos; aarch64)" BillingWeeklyPath = "/billing?format=credits" BillingMonthlyPath = "/billing" @@ -43,14 +45,22 @@ type BillingProductUsage struct { } // BillingConfig is the nested config object from /v1/billing responses. +// Weekly (`?format=credits`) and monthly (`/billing`) share this shape; absolute +// money fields typically appear on the credits (prepaid/on-demand) or monthly +// (limit/used) responses. type BillingConfig struct { - CurrentPeriod *BillingPeriod `json:"currentPeriod,omitempty"` - CreditUsagePercent *float64 `json:"creditUsagePercent,omitempty"` - ProductUsage []BillingProductUsage `json:"productUsage,omitempty"` - MonthlyLimit json.RawMessage `json:"monthlyLimit,omitempty"` - Used json.RawMessage `json:"used,omitempty"` - BillingPeriodStart string `json:"billingPeriodStart,omitempty"` - BillingPeriodEnd string `json:"billingPeriodEnd,omitempty"` + CurrentPeriod *BillingPeriod `json:"currentPeriod,omitempty"` + CreditUsagePercent *float64 `json:"creditUsagePercent,omitempty"` + ProductUsage []BillingProductUsage `json:"productUsage,omitempty"` + MonthlyLimit json.RawMessage `json:"monthlyLimit,omitempty"` + Used json.RawMessage `json:"used,omitempty"` + OnDemandCap json.RawMessage `json:"onDemandCap,omitempty"` + OnDemandUsed json.RawMessage `json:"onDemandUsed,omitempty"` + PrepaidBalance json.RawMessage `json:"prepaidBalance,omitempty"` + IsUnifiedBillingUser bool `json:"isUnifiedBillingUser,omitempty"` + TopUpMethod string `json:"topUpMethod,omitempty"` + BillingPeriodStart string `json:"billingPeriodStart,omitempty"` + BillingPeriodEnd string `json:"billingPeriodEnd,omitempty"` } // BillingPayload is the top-level body from /v1/billing. @@ -65,6 +75,8 @@ type BillingProductSummary struct { } // BillingSummary is the merged weekly + monthly billing view. +// Cents fields remain the authoritative monthly numbers; dollar fields are the +// operator-facing absolute money view (prepaid / on-demand / monthly $). type BillingSummary struct { PeriodType string `json:"period_type,omitempty"` // weekly | monthly | unknown UsagePercent *float64 `json:"usage_percent,omitempty"` @@ -77,17 +89,26 @@ type BillingSummary struct { BillingPeriodStart string `json:"billing_period_start,omitempty"` BillingPeriodEnd string `json:"billing_period_end,omitempty"` UsedPercent *float64 `json:"used_percent,omitempty"` - Plan string `json:"plan,omitempty"` // SuperGrok | SuperGrok Heavy | "" - StatusCode int `json:"status_code,omitempty"` - WeeklyStatusCode int `json:"weekly_status_code,omitempty"` - MonthlyStatusCode int `json:"monthly_status_code,omitempty"` - Source string `json:"source,omitempty"` - FetchedAt string `json:"fetched_at,omitempty"` - UpdatedAt string `json:"updated_at,omitempty"` - WeeklyUpdatedAt string `json:"weekly_updated_at,omitempty"` - MonthlyUpdatedAt string `json:"monthly_updated_at,omitempty"` - Partial bool `json:"partial,omitempty"` - FailedWindows []string `json:"failed_windows,omitempty"` + // Absolute money (USD). Prepaid/on-demand come from credits probe as dollars. + // MonthlyLimit/MonthlyUsed are cents/100 for consistent $ display. + PrepaidBalance *float64 `json:"prepaid_balance,omitempty"` + MonthlyLimit *float64 `json:"monthly_limit,omitempty"` + MonthlyUsed *float64 `json:"monthly_used,omitempty"` + OnDemandCap *float64 `json:"on_demand_cap,omitempty"` + OnDemandUsed *float64 `json:"on_demand_used,omitempty"` + TopUpMethod string `json:"top_up_method,omitempty"` + IsUnifiedBillingUser bool `json:"is_unified_billing_user,omitempty"` + Plan string `json:"plan,omitempty"` // SuperGrok | SuperGrok Heavy | "" + StatusCode int `json:"status_code,omitempty"` + WeeklyStatusCode int `json:"weekly_status_code,omitempty"` + MonthlyStatusCode int `json:"monthly_status_code,omitempty"` + Source string `json:"source,omitempty"` + FetchedAt string `json:"fetched_at,omitempty"` + UpdatedAt string `json:"updated_at,omitempty"` + WeeklyUpdatedAt string `json:"weekly_updated_at,omitempty"` + MonthlyUpdatedAt string `json:"monthly_updated_at,omitempty"` + Partial bool `json:"partial,omitempty"` + FailedWindows []string `json:"failed_windows,omitempty"` } // BuildBillingURL builds weekly or monthly billing URL against the CLI chat proxy. @@ -127,7 +148,7 @@ func ApplyCLIBillingHeaders(req *http.Request, accessToken string) { req.Header.Set("Content-Type", "application/json") req.Header.Set(CLITokenAuthHeader, CLITokenAuthValue) req.Header.Set(CLIClientVersionHeader, CLIClientVersion) - req.Header.Set("User-Agent", CLIUserAgent) + req.Header.Set("User-Agent", billingCLIUserAgent) } // ParseBillingPayload unmarshals a billing API response body. @@ -152,18 +173,15 @@ func BuildBillingSummary(config *BillingConfig) *BillingSummary { periodType := resolvePeriodType(period) creditUsage := cloneFloat(config.CreditUsagePercent) + // Weekly period bounds must not fall back to monthly billing period ends — + // that would park accounts on a multi-week horizon when weekly UsagePercent + // is high (scheduler seven_day uses PeriodEnd). periodStart := "" periodEnd := "" if period != nil { periodStart = strings.TrimSpace(period.Start) periodEnd = strings.TrimSpace(period.End) } - if periodStart == "" { - periodStart = strings.TrimSpace(config.BillingPeriodStart) - } - if periodEnd == "" { - periodEnd = strings.TrimSpace(config.BillingPeriodEnd) - } products := make([]BillingProductSummary, 0, len(config.ProductUsage)) for _, item := range config.ProductUsage { @@ -179,6 +197,11 @@ func BuildBillingSummary(config *BillingConfig) *BillingSummary { monthlyLimit := parseCentValue(config.MonthlyLimit) used := parseCentValue(config.Used) + // Absolute money on credits responses is dollar-denominated ({"val": 12}). + // Monthly limit/used are cents (same as MonthlyLimitCents / UsedCents). + prepaid := parseCentValue(config.PrepaidBalance) + onDemandCap := parseCentValue(config.OnDemandCap) + onDemandUsed := parseCentValue(config.OnDemandUsed) billingStart := strings.TrimSpace(config.BillingPeriodStart) billingEnd := strings.TrimSpace(config.BillingPeriodEnd) @@ -198,7 +221,7 @@ func BuildBillingSummary(config *BillingConfig) *BillingSummary { usedPercent = &v } - hasWeekly := creditUsage != nil || periodType == "weekly" || len(products) > 0 + hasWeekly := creditUsage != nil || periodType == "weekly" || len(products) > 0 || prepaid != nil || onDemandCap != nil || onDemandUsed != nil hasMonthly := monthlyLimit != nil || used != nil || (!hasWeekly && billingEnd != "") if !hasWeekly && !hasMonthly { return nil @@ -228,6 +251,24 @@ func BuildBillingSummary(config *BillingConfig) *BillingSummary { summary.BillingPeriodEnd = billingEnd } summary.UsedPercent = usedPercent + summary.PrepaidBalance = prepaid + if onDemandCap != nil { + summary.OnDemandCap = onDemandCap + } + if onDemandUsed != nil { + summary.OnDemandUsed = onDemandUsed + } + // Expose monthly cents as dollars for UI absolute rows. + if monthlyLimit != nil { + v := *monthlyLimit / 100 + summary.MonthlyLimit = &v + } + if used != nil { + v := *used / 100 + summary.MonthlyUsed = &v + } + summary.TopUpMethod = strings.TrimSpace(config.TopUpMethod) + summary.IsUnifiedBillingUser = config.IsUnifiedBillingUser summary.Plan = resolvePlan(monthlyLimit) return summary } @@ -257,6 +298,22 @@ func MergeBillingProbeResult(previous, weekly, monthly *BillingSummary, weeklyOK out.PeriodStart = weekly.PeriodStart out.PeriodEnd = weekly.PeriodEnd out.ProductUsage = weekly.ProductUsage + // Absolute prepaid / on-demand usually ride the credits (weekly) response. + if weekly.PrepaidBalance != nil { + out.PrepaidBalance = weekly.PrepaidBalance + } + if weekly.OnDemandCap != nil { + out.OnDemandCap = weekly.OnDemandCap + } + if weekly.OnDemandUsed != nil { + out.OnDemandUsed = weekly.OnDemandUsed + } + if weekly.TopUpMethod != "" { + out.TopUpMethod = weekly.TopUpMethod + } + if weekly.IsUnifiedBillingUser { + out.IsUnifiedBillingUser = true + } out.WeeklyUpdatedAt = now } if monthlyOK && monthly != nil { @@ -269,6 +326,15 @@ func MergeBillingProbeResult(previous, weekly, monthly *BillingSummary, weeklyOK out.BillingPeriodStart = monthly.BillingPeriodStart out.BillingPeriodEnd = monthly.BillingPeriodEnd out.UsedPercent = monthly.UsedPercent + out.MonthlyLimit = monthly.MonthlyLimit + out.MonthlyUsed = monthly.MonthlyUsed + // Monthly probe may also carry on-demand cap when credits omitted it. + if monthly.OnDemandCap != nil && out.OnDemandCap == nil { + out.OnDemandCap = monthly.OnDemandCap + } + if monthly.OnDemandUsed != nil && out.OnDemandUsed == nil { + out.OnDemandUsed = monthly.OnDemandUsed + } out.Plan = monthly.Plan out.MonthlyUpdatedAt = now } diff --git a/backend/internal/pkg/xai/billing_test.go b/backend/internal/pkg/xai/billing_test.go index b508d5494f..3dbfc89bf2 100644 --- a/backend/internal/pkg/xai/billing_test.go +++ b/backend/internal/pkg/xai/billing_test.go @@ -48,7 +48,11 @@ func TestBuildBillingSummaryWeeklyAndMonthly(t *testing.T) { "config": { "currentPeriod": {"type":"WEEKLY","start":"2026-07-09T03:25:00Z","end":"2026-07-16T03:25:00Z"}, "creditUsagePercent": 2.0, - "productUsage": [{"product":"Api","usagePercent":2.0}] + "productUsage": [{"product":"Api","usagePercent":2.0}], + "prepaidBalance": {"val": 12}, + "onDemandCap": {"val": 100}, + "onDemandUsed": {"val": 5}, + "isUnifiedBillingUser": true } }`) monthlyBody := []byte(`{ @@ -72,10 +76,16 @@ func TestBuildBillingSummaryWeeklyAndMonthly(t *testing.T) { require.Equal(t, "weekly", weekly.PeriodType) require.InDelta(t, 2.0, *weekly.UsagePercent, 1e-9) require.Equal(t, "Api", weekly.ProductUsage[0].Product) + require.InDelta(t, 12, *weekly.PrepaidBalance, 1e-9) + require.InDelta(t, 100, *weekly.OnDemandCap, 1e-9) + require.InDelta(t, 5, *weekly.OnDemandUsed, 1e-9) + require.True(t, weekly.IsUnifiedBillingUser) require.Equal(t, "SuperGrok", monthly.Plan) require.InDelta(t, 15000, *monthly.MonthlyLimitCents, 1e-9) require.InDelta(t, 78, *monthly.UsedCents, 1e-9) require.InDelta(t, 0.52, *monthly.UsedPercent, 1e-2) + require.InDelta(t, 150, *monthly.MonthlyLimit, 1e-9) + require.InDelta(t, 0.78, *monthly.MonthlyUsed, 1e-9) merged := MergeBillingProbeResult(nil, weekly, monthly, true, true) require.Equal(t, "weekly", merged.PeriodType) @@ -83,6 +93,10 @@ func TestBuildBillingSummaryWeeklyAndMonthly(t *testing.T) { require.Equal(t, "SuperGrok", merged.Plan) require.InDelta(t, 15000, *merged.MonthlyLimitCents, 1e-9) require.Equal(t, "2026-08-01T00:00:00Z", merged.BillingPeriodEnd) + require.InDelta(t, 12, *merged.PrepaidBalance, 1e-9) + require.InDelta(t, 100, *merged.OnDemandCap, 1e-9) + require.InDelta(t, 150, *merged.MonthlyLimit, 1e-9) + require.InDelta(t, 0.78, *merged.MonthlyUsed, 1e-9) } func TestParseCentValueBareNumber(t *testing.T) { @@ -105,6 +119,27 @@ func TestBuildBillingSummaryMonthlyOnlyKeepsWeeklyUsageEmpty(t *testing.T) { require.InDelta(t, 50, *summary.UsedPercent, 1e-9) } +func TestBuildBillingSummaryWeeklyDoesNotInheritMonthlyPeriodEnd(t *testing.T) { + t.Parallel() + // Weekly usage without currentPeriod.end must not copy billingPeriodEnd (monthly). + payload, err := ParseBillingPayload([]byte(`{ + "config": { + "creditUsagePercent": 95.0, + "productUsage": [{"product":"Api","usagePercent":95.0}], + "billingPeriodStart": "2026-07-01T00:00:00Z", + "billingPeriodEnd": "2026-08-01T00:00:00Z", + "monthlyLimit": {"val": 15000}, + "used": {"val": 1000} + } + }`)) + require.NoError(t, err) + summary := BuildBillingSummary(payload.Config) + require.NotNil(t, summary) + require.Equal(t, "weekly", summary.PeriodType) + require.Equal(t, "", summary.PeriodEnd) + require.Equal(t, "2026-08-01T00:00:00Z", summary.BillingPeriodEnd) +} + func TestMergeBillingProbeResultRetainsFailedWindow(t *testing.T) { t.Parallel() previous := &BillingSummary{ diff --git a/backend/internal/pkg/xai/cli_identity.go b/backend/internal/pkg/xai/cli_identity.go new file mode 100644 index 0000000000..480e2e40af --- /dev/null +++ b/backend/internal/pkg/xai/cli_identity.go @@ -0,0 +1,79 @@ +package xai + +import ( + "net/http" + "os" + "strings" + + "golang.org/x/mod/semver" +) + +// Fixed Grok Build / CLI-chat-proxy client identity. +// These values are intentionally pinned in-binary (not scraped from live CLI). +// Operators may bump the version via XAI_GROK_CLI_VERSION without a release. +const ( + // CLIProxyHost is the hostname that requires the official CLI identity headers. + CLIProxyHost = "cli-chat-proxy.grok.com" + + // CLIStableVersion is the known-good minimum client version accepted by cli-chat-proxy. + CLIStableVersion = "0.2.93" + + // CLIVersionEnv is the optional operator override for CLIStableVersion. + CLIVersionEnv = "XAI_GROK_CLI_VERSION" + + // CLITokenAuth is required by cli-chat-proxy for Grok Build OAuth tokens. + CLITokenAuth = "xai-grok-cli" + + // CLIClientIdentifier is the x-grok-client-identifier value used by Grok shell/CLI. + CLIClientIdentifier = "grok-shell" + + // CLIClientMode is used by billing / quota probes on the CLI surface. + CLIClientMode = "cli" +) + +// ResolveCLIVersion returns a supported CLI client version. +// Empty or invalid overrides fall back to CLIClientVersion (the pinned +// preferred client pin in billing.go). CLIStableVersion is only the minimum +// accepted by IsSupportedCLIVersion, not the default identity we advertise. +func ResolveCLIVersion() string { + version := strings.TrimSpace(os.Getenv(CLIVersionEnv)) + if !IsSupportedCLIVersion(version) { + return CLIClientVersion + } + return version +} + +// IsSupportedCLIVersion reports whether version is a valid semver string at or +// above CLIStableVersion (prereleases below a higher release are rejected when +// they compare less than the stable pin). +func IsSupportedCLIVersion(version string) bool { + canonical := "v" + version + minimum := "v" + CLIStableVersion + return semver.IsValid(canonical) && + semver.Canonical(canonical) == canonical && + semver.Compare(canonical, minimum) >= 0 +} + +// CLIUserAgent builds the workspace-style User-Agent for a CLI client version. +func CLIUserAgent(version string) string { + if strings.TrimSpace(version) == "" { + version = CLIClientVersion + } + return "xai-grok-workspace/" + version +} + +// ApplyCLIProxyHeaders stamps the fixed Grok CLI identity when the request +// targets cli-chat-proxy. Direct api.x.ai traffic is left unchanged. +func ApplyCLIProxyHeaders(req *http.Request) { + if req == nil || req.URL == nil || !strings.EqualFold(strings.TrimSpace(req.URL.Hostname()), CLIProxyHost) { + return + } + if req.Header == nil { + req.Header = make(http.Header) + } + version := ResolveCLIVersion() + req.Header.Set("X-XAI-Token-Auth", CLITokenAuth) + req.Header.Set("x-grok-client-version", version) + req.Header.Set("x-grok-client-identifier", CLIClientIdentifier) + req.Header.Set("User-Agent", CLIUserAgent(version)) +} diff --git a/backend/internal/pkg/xai/cli_identity_test.go b/backend/internal/pkg/xai/cli_identity_test.go new file mode 100644 index 0000000000..521df7483d --- /dev/null +++ b/backend/internal/pkg/xai/cli_identity_test.go @@ -0,0 +1,67 @@ +package xai + +import ( + "net/http" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestResolveCLIVersionDefaultsToPinnedClientVersion(t *testing.T) { + t.Setenv(CLIVersionEnv, "") + // Default advertise pin is CLIClientVersion; CLIStableVersion is only the floor. + require.Equal(t, CLIClientVersion, ResolveCLIVersion()) + require.True(t, IsSupportedCLIVersion(CLIClientVersion)) + require.True(t, IsSupportedCLIVersion(CLIStableVersion)) +} + +func TestResolveCLIVersionAcceptsValidOverride(t *testing.T) { + t.Setenv(CLIVersionEnv, "0.2.95-alpha.1") + require.Equal(t, "0.2.95-alpha.1", ResolveCLIVersion()) +} + +func TestResolveCLIVersionRejectsUnsafeOrTooOld(t *testing.T) { + for _, version := range []string{ + "0.2.92", + "0.2.93-beta.1", + "0.2.95\r\nX-Injected: true", + "0.2.093", + "0.3", + "1", + } { + t.Run(version, func(t *testing.T) { + t.Setenv(CLIVersionEnv, version) + require.Equal(t, CLIClientVersion, ResolveCLIVersion()) + }) + } +} + +func TestApplyCLIProxyHeaders(t *testing.T) { + t.Setenv(CLIVersionEnv, "") + + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "sub2api-grok/1.0") + + ApplyCLIProxyHeaders(req) + + require.Equal(t, CLIClientVersion, req.Header.Get("x-grok-client-version")) + require.Equal(t, CLIClientIdentifier, req.Header.Get("x-grok-client-identifier")) + require.Equal(t, CLITokenAuth, req.Header.Get("X-XAI-Token-Auth")) + require.Equal(t, CLIUserAgent(CLIClientVersion), req.Header.Get("User-Agent")) +} + +func TestApplyCLIProxyHeadersLeavesAPIHostUnchanged(t *testing.T) { + t.Setenv(CLIVersionEnv, "0.2.95") + + req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "sub2api-grok/1.0") + + ApplyCLIProxyHeaders(req) + + require.Empty(t, req.Header.Get("x-grok-client-version")) + require.Empty(t, req.Header.Get("x-grok-client-identifier")) + require.Empty(t, req.Header.Get("X-XAI-Token-Auth")) + require.Equal(t, "sub2api-grok/1.0", req.Header.Get("User-Agent")) +} diff --git a/backend/internal/pkg/xai/models.go b/backend/internal/pkg/xai/models.go index 3a65f32c51..c25f489317 100644 --- a/backend/internal/pkg/xai/models.go +++ b/backend/internal/pkg/xai/models.go @@ -1,28 +1,130 @@ package xai +import ( + "strings" + "sync/atomic" +) + +// runtimeMappingOpts holds operator-configured defaults applied when Grok +// accounts leave credentials.model_mapping empty. Updated from settings. +var runtimeMappingOpts atomic.Value // ModelMappingOptions +var runtimeMappingVersion atomic.Uint64 + +func init() { + runtimeMappingOpts.Store(ModelMappingOptions{}) + runtimeMappingVersion.Store(1) +} + +// SetRuntimeModelMappingOptions updates process-wide defaults used by +// DefaultModelMapping (e.g. after settings load). Safe for concurrent use. +func SetRuntimeModelMappingOptions(opts ModelMappingOptions) { + runtimeMappingOpts.Store(opts) + runtimeMappingVersion.Add(1) +} + +// RuntimeModelMappingVersion changes whenever runtime mapping options change. +// Account-level caches include it so settings updates take effect without a restart. +func RuntimeModelMappingVersion() uint64 { + return runtimeMappingVersion.Load() +} + +// RuntimeModelMappingOptions returns the last options set via SetRuntimeModelMappingOptions. +func RuntimeModelMappingOptions() ModelMappingOptions { + if v := runtimeMappingOpts.Load(); v != nil { + if opts, ok := v.(ModelMappingOptions); ok { + return opts + } + } + return ModelMappingOptions{} +} + // Model describes an xAI model in OpenAI-compatible /models shape. type Model struct { ID string `json:"id"` Object string `json:"object"` + Type string `json:"type,omitempty"` Created int64 `json:"created,omitempty"` OwnedBy string `json:"owned_by"` DisplayName string `json:"display_name,omitempty"` } +// DefaultTextModel is the built-in fallback for empty model fields and Grok +// text aliases (e.g. "grok", "grok-latest"). Operators may override the runtime +// default via settings key grok_default_text_model. +const DefaultTextModel = "grok-4.5" + +// Official Imagine model IDs (https://docs.x.ai/docs/models). +const ( + DefaultImagineImageQualityModel = "grok-imagine-image-quality" + DefaultImagineImageFastModel = "grok-imagine-image" + DefaultImagineVideoModel = "grok-imagine-video" + DefaultImagineVideo15LegacyModel = "grok-imagine-video-1.5" + DefaultImagineVideo15Model = "grok-imagine-video-1.5-preview" +) + +// ModelMappingOptions controls optional expansions of the default mapping. +// Cross-client wildcards (gpt-*/claude-*) default ON via settings +// grok_cross_client_model_map_enabled so Codex/Claude clients keep working +// against Grok groups (map to DefaultText / grok-4.5). Operators may disable. +type ModelMappingOptions struct { + // DefaultText is the target for empty models and optional cross-client maps. + // Empty → DefaultTextModel (grok-4.5). + DefaultText string + // EnableCrossClientMap merges gpt-*/codex-*/o*/claude-* → DefaultText. + EnableCrossClientMap bool +} + +func (o ModelMappingOptions) defaultText() string { + if t := strings.TrimSpace(o.DefaultText); t != "" { + return t + } + return DefaultTextModel +} + var defaultModels = []Model{ - {ID: "grok-4.5", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.5"}, - {ID: "grok-4.3", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.3"}, - {ID: "grok-build-0.1", Object: "model", OwnedBy: "xai", DisplayName: "Grok Build 0.1"}, - {ID: "grok-composer-2.5-fast", Object: "model", OwnedBy: "xai", DisplayName: "Grok Composer 2.5 Fast"}, - {ID: "grok-4.20-0309-reasoning", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Reasoning"}, - {ID: "grok-4.20-0309-non-reasoning", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Non Reasoning"}, - {ID: "grok-4.20-multi-agent-0309", Object: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Multi Agent"}, - {ID: "grok-imagine", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine"}, - {ID: "grok-imagine-image", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image"}, - {ID: "grok-imagine-image-quality", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image Quality"}, - {ID: "grok-imagine-edit", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Edit"}, - {ID: "grok-imagine-video", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video"}, - {ID: "grok-imagine-video-1.5", Object: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5"}, + // Text + {ID: "grok-4.5", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.5"}, + {ID: "grok-4.3", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.3"}, + {ID: "grok-3-mini", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 3 Mini"}, + {ID: "grok-3-mini-fast", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 3 Mini Fast"}, + {ID: "grok-build-0.1", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Build 0.1"}, + {ID: "grok-composer-2.5-fast", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Composer 2.5 Fast"}, + {ID: "grok-4.20-0309-reasoning", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Reasoning"}, + {ID: "grok-4.20-0309-non-reasoning", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Non Reasoning"}, + {ID: "grok-4.20-multi-agent-0309", Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok 4.20 Multi Agent"}, + // Imagine + {ID: DefaultImagineImageQualityModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image Quality"}, + {ID: DefaultImagineImageFastModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Image"}, + {ID: DefaultImagineVideoModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video"}, + {ID: DefaultImagineVideo15Model, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5 Preview"}, + {ID: DefaultImagineVideo15LegacyModel, Object: "model", Type: "model", OwnedBy: "xai", DisplayName: "Grok Imagine Video 1.5 Legacy"}, +} + +// grokTextResponsesModelAliases is the source of truth for Grok text models +// accepted by the Responses path: client-facing / undated aliases → canonical +// upstream ID. Used by DefaultModelMapping and IsGrokTextResponsesModelID. +var grokTextResponsesModelAliases = map[string]string{ + "grok": DefaultTextModel, + "grok-latest": DefaultTextModel, + "grok-4.5": DefaultTextModel, + "grok-4.5-latest": DefaultTextModel, + "grok-4.3": "grok-4.3", + "grok-4.3-latest": "grok-4.3", + "grok-3-mini": "grok-3-mini", + "grok-3-mini-fast": "grok-3-mini-fast", + "grok-build": "grok-build-0.1", + "grok-build-latest": DefaultTextModel, + "grok-build-0.1": "grok-build-0.1", + "grok-composer-2.5-fast": "grok-composer-2.5-fast", + "grok-composer": "grok-composer-2.5-fast", + "composer-2.5": "grok-composer-2.5-fast", + "grok-4.20-reasoning": "grok-4.20-0309-reasoning", + "grok-4.20-0309-reasoning": "grok-4.20-0309-reasoning", + "grok-4.20-non-reasoning": "grok-4.20-0309-non-reasoning", + "grok-4.20-0309-non-reasoning": "grok-4.20-0309-non-reasoning", + "grok-4.20-multi-agent": "grok-4.20-multi-agent-0309", + "grok-4.20-multi-agent-latest": "grok-4.20-multi-agent-0309", + "grok-4.20-multi-agent-0309": "grok-4.20-multi-agent-0309", } func DefaultModels() []Model { @@ -40,19 +142,163 @@ func DefaultModelIDs() []string { return ids } +// DefaultModelMapping returns native Grok/Imagine identity + aliases, using +// runtime options (default text model / optional cross-client wildcards). +// Does NOT enable gpt-*/claude-* unless SetRuntimeModelMappingOptions enables them. func DefaultModelMapping() map[string]string { - mapping := make(map[string]string, len(defaultModels)+5) + return ModelMappingWithOptions(RuntimeModelMappingOptions()) +} + +// ModelMappingWithOptions builds the default Grok mapping with optional +// cross-client wildcards and a configurable default text model. +func ModelMappingWithOptions(opts ModelMappingOptions) map[string]string { + defaultText := opts.defaultText() + mapping := make(map[string]string, len(defaultModels)+len(grokTextResponsesModelAliases)+48) for _, model := range defaultModels { mapping[model.ID] = model.ID } - mapping["grok"] = "grok-4.5" - mapping["grok-latest"] = "grok-4.5" - mapping["grok-4.5-latest"] = "grok-4.5" - mapping["grok-build"] = "grok-build-0.1" - mapping["grok-build-latest"] = "grok-4.5" - mapping["grok-composer"] = "grok-composer-2.5-fast" - mapping["composer-2.5"] = "grok-composer-2.5-fast" - mapping["grok-4.20-reasoning"] = "grok-4.20-0309-reasoning" - mapping["grok-4.20-non-reasoning"] = "grok-4.20-0309-non-reasoning" + for alias, canonical := range grokTextResponsesModelAliases { + // Remap aliases that pointed at DefaultTextModel constant to runtime default. + if canonical == DefaultTextModel { + mapping[alias] = defaultText + } else { + mapping[alias] = canonical + } + } + // Imagine aliases / legacy IDs → official catalog. + mapping["grok-imagine"] = DefaultImagineImageQualityModel + mapping["grok-imagine-1"] = DefaultImagineImageQualityModel + // Backward-compatible client alias; xAI exposes image editing through the + // image-quality model rather than a separate grok-imagine-edit model. + mapping["grok-imagine-edit"] = DefaultImagineImageQualityModel + mapping["grok-imagine-image"] = DefaultImagineImageFastModel + mapping["grok-imagine-image-quality"] = DefaultImagineImageQualityModel + // Keep official IDs as identity so client-requested model strings are not + // rewritten on the wire (pricing still canonicalizes 1.5* via CanonicalImagineVideoModel). + mapping["grok-imagine-video"] = DefaultImagineVideoModel + mapping["grok-imagine-video-1.5"] = DefaultImagineVideo15LegacyModel + mapping["grok-imagine-video-1.5-preview"] = DefaultImagineVideo15Model + // Informal alias only: + mapping["grok-video-1.5"] = DefaultImagineVideo15Model + + if opts.EnableCrossClientMap { + // Codex / OpenAI Responses client defaults (wildcard patterns). + mapping["gpt-*"] = defaultText + mapping["codex-*"] = defaultText + mapping["o1*"] = defaultText + mapping["o3*"] = defaultText + mapping["o4*"] = defaultText + // Claude Code defaults when operators intentionally enable bridging. + mapping["claude-*"] = defaultText + } + addGrokProviderPrefixedMappings(mapping) return mapping } + +func addGrokProviderPrefixedMappings(mapping map[string]string) { + snapshot := make(map[string]string, len(mapping)) + for key, value := range mapping { + snapshot[key] = value + } + for key, value := range snapshot { + if !isGrokNativeOrAlias(key) { + continue + } + for _, prefix := range []string{"xai/", "x-ai/", "grok/"} { + mapping[prefix+key] = value + } + } +} + +func isGrokNativeOrAlias(model string) bool { + model = strings.ToLower(strings.TrimSpace(model)) + if model == "" || strings.Contains(model, "*") { + return false + } + return strings.HasPrefix(model, "grok") || + strings.HasPrefix(model, "imagine") || + strings.HasPrefix(model, "composer") +} + +// StripGrokProviderPrefix removes common provider prefixes accepted for +// xAI/Grok models, returning the native model ID. +func StripGrokProviderPrefix(model string) string { + trimmed := strings.TrimSpace(model) + lower := strings.ToLower(trimmed) + for _, prefix := range []string{"xai/", "x-ai/", "grok/"} { + if strings.HasPrefix(lower, prefix) { + return strings.TrimSpace(trimmed[len(prefix):]) + } + } + return trimmed +} + +// IsGrokModelID reports whether model looks like a native Grok/xAI model id +// (including aliases). Claude/OpenAI model names return false. +func IsGrokModelID(model string) bool { + normalized := strings.ToLower(StripGrokProviderPrefix(model)) + if normalized == "" { + return false + } + if strings.HasPrefix(normalized, "grok") { + return true + } + if strings.HasPrefix(normalized, "imagine") { + return true + } + return false +} + +// IsGrokTextResponsesModelID reports whether model is a known Grok text model +// for the Responses API. Imagine image/video and unknown custom ids return false. +func IsGrokTextResponsesModelID(model string) bool { + normalized := strings.ToLower(StripGrokProviderPrefix(model)) + _, ok := grokTextResponsesModelAliases[normalized] + return ok +} + +// ResolveGrokTextResponsesModelID canonicalizes a Grok text alias before upstream. +// empty or bare aliases that resolve via DefaultTextModel use defaultText when set. +func ResolveGrokTextResponsesModelID(model string, defaultText ...string) string { + fallback := DefaultTextModel + if len(defaultText) > 0 && strings.TrimSpace(defaultText[0]) != "" { + fallback = strings.TrimSpace(defaultText[0]) + } + trimmed := strings.TrimSpace(model) + if trimmed == "" { + return fallback + } + normalized := strings.ToLower(StripGrokProviderPrefix(trimmed)) + if canonical, ok := grokTextResponsesModelAliases[normalized]; ok { + if canonical == DefaultTextModel { + return fallback + } + return canonical + } + return StripGrokProviderPrefix(trimmed) +} + +// ResolveDefaultTextModel returns defaultText (or DefaultTextModel) when model is empty. +func ResolveDefaultTextModel(model string, defaultText ...string) string { + if trimmed := strings.TrimSpace(model); trimmed != "" { + return trimmed + } + if len(defaultText) > 0 && strings.TrimSpace(defaultText[0]) != "" { + return strings.TrimSpace(defaultText[0]) + } + return DefaultTextModel +} + +// CanonicalImagineVideoModel normalizes video model ids for pricing tables. +// Legacy "grok-imagine-video-1.5" shares the 1.5 price family with preview. +func CanonicalImagineVideoModel(model string) string { + m := strings.ToLower(StripGrokProviderPrefix(model)) + switch { + case m == "" || m == DefaultImagineVideoModel || m == "grok-imagine-video-preview": + return DefaultImagineVideoModel + case strings.HasPrefix(m, "grok-imagine-video-1.5") || m == "grok-video-1.5": + return DefaultImagineVideo15Model + default: + return m + } +} diff --git a/backend/internal/pkg/xai/models_test.go b/backend/internal/pkg/xai/models_test.go new file mode 100644 index 0000000000..98aef08cfc --- /dev/null +++ b/backend/internal/pkg/xai/models_test.go @@ -0,0 +1,65 @@ +package xai + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestDefaultModelMappingExcludesCrossClientWildcards(t *testing.T) { + original := RuntimeModelMappingOptions() + t.Cleanup(func() { SetRuntimeModelMappingOptions(original) }) + SetRuntimeModelMappingOptions(ModelMappingOptions{}) + mapping := DefaultModelMapping() + + require.Equal(t, "grok-4.5", mapping["grok"]) + require.Equal(t, "grok-4.5", mapping["grok-latest"]) + require.Equal(t, "grok-build-0.1", mapping["grok-build"]) + require.Equal(t, DefaultTextModel, mapping["grok-build-latest"]) + require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-edit"]) + require.Equal(t, DefaultImagineVideo15LegacyModel, mapping["grok-imagine-video-1.5"]) + require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"]) + require.Equal(t, "grok-4.5", mapping["xai/grok"]) + + // Cross-vendor wildcards must stay opt-in. + _, hasGPT := mapping["gpt-*"] + _, hasClaude := mapping["claude-*"] + require.False(t, hasGPT) + require.False(t, hasClaude) +} + +func TestModelMappingWithOptionsCrossClient(t *testing.T) { + t.Parallel() + mapping := ModelMappingWithOptions(ModelMappingOptions{ + DefaultText: "grok-4.3", + EnableCrossClientMap: true, + }) + require.Equal(t, "grok-4.3", mapping["grok"]) + require.Equal(t, "grok-4.3", mapping["gpt-*"]) + require.Equal(t, "grok-4.3", mapping["claude-*"]) + require.Equal(t, "grok-4.3", mapping["codex-*"]) +} + +func TestCanonicalImagineVideoModel(t *testing.T) { + t.Parallel() + require.Equal(t, DefaultImagineVideoModel, CanonicalImagineVideoModel("grok-imagine-video")) + require.Equal(t, DefaultImagineVideo15Model, CanonicalImagineVideoModel("grok-imagine-video-1.5")) + require.Equal(t, DefaultImagineVideo15Model, CanonicalImagineVideoModel("grok-imagine-video-1.5-preview")) + require.Equal(t, DefaultImagineVideo15Model, CanonicalImagineVideoModel("xai/grok-video-1.5")) + require.Equal(t, "grok-imagine-video-2", CanonicalImagineVideoModel("grok-imagine-video-2")) +} + +func TestIsGrokModelID(t *testing.T) { + t.Parallel() + require.True(t, IsGrokModelID("grok-4.5")) + require.True(t, IsGrokModelID("x-ai/grok-4.3")) + require.False(t, IsGrokModelID("gpt-5")) + require.False(t, IsGrokModelID("claude-sonnet-4")) +} + +func TestResolveGrokTextResponsesModelID(t *testing.T) { + t.Parallel() + require.Equal(t, "grok-4.5", ResolveGrokTextResponsesModelID("")) + require.Equal(t, "grok-4.3", ResolveGrokTextResponsesModelID("grok", "grok-4.3")) + require.Equal(t, "grok-4.20-multi-agent-0309", ResolveGrokTextResponsesModelID("grok-4.20-multi-agent")) +} diff --git a/backend/internal/pkg/xai/oauth.go b/backend/internal/pkg/xai/oauth.go index afa44d06a0..a59d65ddf6 100644 --- a/backend/internal/pkg/xai/oauth.go +++ b/backend/internal/pkg/xai/oauth.go @@ -1,33 +1,40 @@ package xai import ( + "context" "crypto/rand" "crypto/sha256" "encoding/base64" "encoding/hex" "errors" "fmt" + "log/slog" "net/url" "os" "strings" "sync" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/redissession" "github.com/Wei-Shaw/sub2api/internal/util/logredact" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" + "github.com/redis/go-redis/v9" ) const ( - OAuthIssuer = "https://auth.x.ai" - DiscoveryURL = OAuthIssuer + "/.well-known/openid-configuration" - DefaultAuthorizeURL = OAuthIssuer + "/oauth2/authorize" - DefaultTokenURL = OAuthIssuer + "/oauth2/token" - DefaultBaseURL = "https://api.x.ai/v1" - DefaultCLIBaseURL = "https://cli-chat-proxy.grok.com/v1" - DefaultClientID = "b1a00492-073a-47ea-816f-4c329264a828" - DefaultScope = "openid profile email offline_access grok-cli:access api:access" - DefaultRedirectURI = "http://127.0.0.1:56121/callback" - SessionTTL = 30 * time.Minute + OAuthIssuer = "https://auth.x.ai" + DiscoveryURL = OAuthIssuer + "/.well-known/openid-configuration" + DefaultAuthorizeURL = OAuthIssuer + "/oauth2/authorize" + DefaultTokenURL = OAuthIssuer + "/oauth2/token" + DefaultBaseURL = "https://api.x.ai/v1" + DefaultCLIBaseURL = "https://cli-chat-proxy.grok.com/v1" + DefaultUSEast1BaseURL = "https://us-east-1.api.x.ai/v1" + DefaultUSWest2BaseURL = "https://us-west-2.api.x.ai/v1" + DefaultEUWest1BaseURL = "https://eu-west-1.api.x.ai/v1" + DefaultClientID = "b1a00492-073a-47ea-816f-4c329264a828" + DefaultScope = "openid profile email offline_access grok-cli:access api:access" + DefaultRedirectURI = "http://127.0.0.1:56121/callback" + SessionTTL = 30 * time.Minute EnvAuthorizeURL = "XAI_OAUTH_AUTHORIZE_URL" EnvTokenURL = "XAI_OAUTH_TOKEN_URL" @@ -56,32 +63,110 @@ type OAuthSession struct { ProxyURL string `json:"proxy_url,omitempty"` RedirectURI string `json:"redirect_uri"` CreatedAt time.Time `json:"created_at"` + + mu sync.Mutex + consumed bool } -// SessionStore manages xAI OAuth sessions in memory. +func (s *OAuthSession) TryConsume() bool { + if s == nil { + return false + } + s.mu.Lock() + defer s.mu.Unlock() + if s.consumed { + return false + } + s.consumed = true + return true +} + +// SessionStore manages xAI OAuth sessions with an optional Redis backend. type SessionStore struct { - mu sync.RWMutex - sessions map[string]*OAuthSession - stopOnce sync.Once - stopCh chan struct{} + mu sync.RWMutex + sessions map[string]*OAuthSession + localOnly map[string]struct{} + stopOnce sync.Once + stopCh chan struct{} + remote *redissession.Store +} + +type oauthSessionDTO struct { + State string `json:"state"` + CodeVerifier string `json:"code_verifier"` + CodeChallenge string `json:"code_challenge"` + ClientID string `json:"client_id,omitempty"` + Scope string `json:"scope,omitempty"` + ProxyURL string `json:"proxy_url,omitempty"` + RedirectURI string `json:"redirect_uri"` + CreatedAt time.Time `json:"created_at"` } func NewSessionStore() *SessionStore { store := &SessionStore{ - sessions: make(map[string]*OAuthSession), - stopCh: make(chan struct{}), + sessions: make(map[string]*OAuthSession), + localOnly: make(map[string]struct{}), + stopCh: make(chan struct{}), } go store.cleanup() return store } +func NewRedisSessionStore(rdb *redis.Client) *SessionStore { + store := NewSessionStore() + if rdb != nil { + store.remote = redissession.New(rdb, "oauth:session:xai", SessionTTL) + } + return store +} + func (s *SessionStore) Set(sessionID string, session *OAuthSession) { + if session == nil { + return + } + var remoteErr error + if s != nil && s.remote != nil { + remoteErr = s.remote.Set(context.Background(), sessionID, oauthSessionDTO{ + State: session.State, CodeVerifier: session.CodeVerifier, CodeChallenge: session.CodeChallenge, + ClientID: session.ClientID, Scope: session.Scope, ProxyURL: session.ProxyURL, + RedirectURI: session.RedirectURI, CreatedAt: session.CreatedAt, + }) + } s.mu.Lock() defer s.mu.Unlock() s.sessions[sessionID] = session + if remoteErr != nil { + s.localOnly[sessionID] = struct{}{} + slog.Warn("xai oauth session Redis write failed; using process-local fallback", "error", remoteErr) + } else { + delete(s.localOnly, sessionID) + } } func (s *SessionStore) Get(sessionID string) (*OAuthSession, bool) { + if s.isLocalOnly(sessionID) { + return s.getMemory(sessionID) + } + if s != nil && s.remote != nil { + var dto oauthSessionDTO + ok, err := s.remote.Get(context.Background(), sessionID, &dto) + if err != nil || !ok || time.Since(dto.CreatedAt) > SessionTTL { + return nil, false + } + session := &OAuthSession{ + State: dto.State, CodeVerifier: dto.CodeVerifier, CodeChallenge: dto.CodeChallenge, + ClientID: dto.ClientID, Scope: dto.Scope, ProxyURL: dto.ProxyURL, + RedirectURI: dto.RedirectURI, CreatedAt: dto.CreatedAt, + } + s.mu.Lock() + s.sessions[sessionID] = session + s.mu.Unlock() + return session, true + } + return s.getMemory(sessionID) +} + +func (s *SessionStore) getMemory(sessionID string) (*OAuthSession, bool) { s.mu.RLock() defer s.mu.RUnlock() session, ok := s.sessions[sessionID] @@ -95,9 +180,39 @@ func (s *SessionStore) Get(sessionID string) (*OAuthSession, bool) { } func (s *SessionStore) Delete(sessionID string) { + if s != nil && s.remote != nil { + _ = s.remote.Delete(context.Background(), sessionID) + } s.mu.Lock() defer s.mu.Unlock() delete(s.sessions, sessionID) + delete(s.localOnly, sessionID) +} + +func (s *SessionStore) TryConsumeSession(sessionID string) bool { + if s == nil { + return false + } + if s.isLocalOnly(sessionID) { + return s.tryConsumeMemory(sessionID) + } + if s.remote != nil { + ok, err := s.remote.TryConsume(context.Background(), sessionID) + return err == nil && ok + } + return s.tryConsumeMemory(sessionID) +} + +func (s *SessionStore) isLocalOnly(sessionID string) bool { + s.mu.RLock() + defer s.mu.RUnlock() + _, ok := s.localOnly[sessionID] + return ok +} + +func (s *SessionStore) tryConsumeMemory(sessionID string) bool { + session, ok := s.getMemory(sessionID) + return ok && session.TryConsume() } func (s *SessionStore) Stop() { @@ -118,6 +233,7 @@ func (s *SessionStore) cleanup() { for id, session := range s.sessions { if time.Since(session.CreatedAt) > SessionTTL { delete(s.sessions, id) + delete(s.localOnly, id) } } s.mu.Unlock() diff --git a/backend/internal/pkg/xai/oauth_redis_fallback_test.go b/backend/internal/pkg/xai/oauth_redis_fallback_test.go new file mode 100644 index 0000000000..40023acd68 --- /dev/null +++ b/backend/internal/pkg/xai/oauth_redis_fallback_test.go @@ -0,0 +1,35 @@ +//go:build unit + +package xai + +import ( + "context" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func TestSessionStoreRedisFallbackIsLimitedToFailedWrites(t *testing.T) { + mr := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: mr.Addr(), MaxRetries: -1}) + t.Cleanup(func() { _ = client.Close() }) + store := NewRedisSessionStore(client) + defer store.Stop() + session := func(state string) *OAuthSession { return &OAuthSession{State: state, CreatedAt: time.Now()} } + + store.Set("remote", session("remote")) + require.NoError(t, store.remote.Delete(context.Background(), "remote")) + _, ok := store.Get("remote") + require.False(t, ok, "a remote miss must not revive the stale local copy") + + mr.Close() + store.Set("local-only", session("local")) + got, ok := store.Get("local-only") + require.True(t, ok) + require.Equal(t, "local", got.State) + require.True(t, store.TryConsumeSession("local-only")) + require.False(t, store.TryConsumeSession("local-only")) +} diff --git a/backend/internal/pkg/xai/oauth_test.go b/backend/internal/pkg/xai/oauth_test.go index 53c8983103..3919f4576b 100644 --- a/backend/internal/pkg/xai/oauth_test.go +++ b/backend/internal/pkg/xai/oauth_test.go @@ -255,6 +255,14 @@ func TestBuildResponsesURLWithValidatorUsesCallerPolicy(t *testing.T) { require.Equal(t, "http://grok.example.test/v1/responses", target) } +func TestValidateTrustedBaseURLAcceptsOfficialRegionalHosts(t *testing.T) { + for _, raw := range []string{DefaultUSEast1BaseURL, DefaultUSWest2BaseURL, DefaultEUWest1BaseURL} { + got, err := ValidateTrustedBaseURL(raw) + require.NoError(t, err, raw) + require.Equal(t, raw, got) + } +} + func TestBuildResponsesURLPreservesUnsafeOverrideCustomPath(t *testing.T) { t.Setenv(EnvAllowUnsafeURLOverrides, "true") @@ -338,8 +346,9 @@ func TestRuntimeSanityReportsInvalidOverridesWithoutSecrets(t *testing.T) { } func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) { - t.Parallel() - + original := RuntimeModelMappingOptions() + t.Cleanup(func() { SetRuntimeModelMappingOptions(original) }) + SetRuntimeModelMappingOptions(ModelMappingOptions{}) mapping := DefaultModelMapping() require.Equal(t, "grok-4.5", mapping["grok"]) require.Equal(t, "grok-4.5", mapping["grok-latest"]) @@ -352,10 +361,13 @@ func TestDefaultModelMappingIncludesGrokAliases(t *testing.T) { require.Equal(t, "grok-4.20-0309-reasoning", mapping["grok-4.20-reasoning"]) require.Equal(t, "grok-4.20-0309-non-reasoning", mapping["grok-4.20-non-reasoning"]) require.Equal(t, "grok-4.20-multi-agent-0309", mapping["grok-4.20-multi-agent-0309"]) - require.Equal(t, "grok-imagine", mapping["grok-imagine"]) - require.Equal(t, "grok-imagine-image", mapping["grok-imagine-image"]) - require.Equal(t, "grok-imagine-image-quality", mapping["grok-imagine-image-quality"]) - require.Equal(t, "grok-imagine-edit", mapping["grok-imagine-edit"]) - require.Equal(t, "grok-imagine-video", mapping["grok-imagine-video"]) - require.Equal(t, "grok-imagine-video-1.5", mapping["grok-imagine-video-1.5"]) + require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine"]) + require.Equal(t, DefaultImagineImageFastModel, mapping["grok-imagine-image"]) + require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-image-quality"]) + require.Equal(t, DefaultImagineImageQualityModel, mapping["grok-imagine-edit"]) + require.Equal(t, DefaultImagineVideoModel, mapping["grok-imagine-video"]) + require.Equal(t, DefaultImagineVideo15LegacyModel, mapping["grok-imagine-video-1.5"]) + require.Equal(t, DefaultImagineVideo15Model, mapping["grok-imagine-video-1.5-preview"]) + _, hasGPT := mapping["gpt-*"] + require.False(t, hasGPT, "cross-client wildcards must be opt-in") } diff --git a/backend/internal/pkg/xai/quota.go b/backend/internal/pkg/xai/quota.go index fe269d1d84..11133a05b3 100644 --- a/backend/internal/pkg/xai/quota.go +++ b/backend/internal/pkg/xai/quota.go @@ -7,10 +7,14 @@ import ( "time" ) -const GrokFreeRolling24hTokenLimit int64 = 1_000_000 +// GrokFreeRolling24hTokenLimit is the operator soft-gate nominal Free allowance +// (rolling 24h). Soft-gate default matches this; upstream header limits may +// still report historical 1M/2M Free snapshots. +const GrokFreeRolling24hTokenLimit int64 = 500_000 var grokFreeRolling24hTokenLimits = map[int64]struct{}{ GrokFreeRolling24hTokenLimit: {}, + 1_000_000: {}, // Observed Free limit variants. 2_000_000: {}, // Legacy Free limit observed before July 2026. } @@ -61,11 +65,27 @@ var quotaHeaderAllowlist = []string{ "x-ratelimit-limit-tokens", "x-ratelimit-remaining-tokens", "x-ratelimit-reset-tokens", + "x-rate-limit-limit-requests", + "x-rate-limit-remaining-requests", + "x-rate-limit-reset-requests", + "x-rate-limit-limit-tokens", + "x-rate-limit-remaining-tokens", + "x-rate-limit-reset-tokens", "retry-after", "x-subscription-tier", "xai-subscription-tier", + "x-xai-subscription-tier", + "x-xai-user-tier", + "xai-user-tier", + "xai-tier", + "x-user-tier", + "x-plan-tier", + "x-subscription-plan", "x-entitlement-status", "xai-entitlement-status", + "x-xai-entitlement-status", + "x-xai-user-entitlement-status", + "x-user-entitlement-status", } func ParseQuotaHeaders(headers http.Header, statusCode int) *QuotaSnapshot { @@ -95,8 +115,24 @@ func parseQuotaHeaders(headers http.Header, statusCode int, source string, keepE if retryAfter := parseRetryAfter(headers.Get("retry-after")); retryAfter != nil { snapshot.RetryAfterSeconds = retryAfter } - snapshot.SubscriptionTier = firstHeader(headers, "xai-subscription-tier", "x-subscription-tier") - snapshot.EntitlementStatus = firstHeader(headers, "xai-entitlement-status", "x-entitlement-status") + snapshot.SubscriptionTier = firstHeader(headers, + "xai-subscription-tier", + "x-subscription-tier", + "x-xai-subscription-tier", + "x-xai-user-tier", + "xai-user-tier", + "xai-tier", + "x-user-tier", + "x-plan-tier", + "x-subscription-plan", + ) + snapshot.EntitlementStatus = firstHeader(headers, + "xai-entitlement-status", + "x-entitlement-status", + "x-xai-entitlement-status", + "x-xai-user-entitlement-status", + "x-user-entitlement-status", + ) for _, name := range quotaHeaderAllowlist { if value := strings.TrimSpace(headers.Get(name)); value != "" { @@ -121,11 +157,23 @@ func parseQuotaHeaders(headers http.Header, statusCode int, source string, keepE } func parseQuotaWindow(headers http.Header, dimension string) *QuotaWindow { + limitHeader := firstHeader(headers, + "x-ratelimit-limit-"+dimension, + "x-rate-limit-limit-"+dimension, + ) + remainingHeader := firstHeader(headers, + "x-ratelimit-remaining-"+dimension, + "x-rate-limit-remaining-"+dimension, + ) + resetHeader := firstHeader(headers, + "x-ratelimit-reset-"+dimension, + "x-rate-limit-reset-"+dimension, + ) window := &QuotaWindow{ - Limit: parseInt64Ptr(headers.Get("x-ratelimit-limit-" + dimension)), - Remaining: parseInt64Ptr(headers.Get("x-ratelimit-remaining-" + dimension)), + Limit: parseInt64Ptr(limitHeader), + Remaining: parseInt64Ptr(remainingHeader), } - if reset := parseResetHeader(headers.Get("x-ratelimit-reset-" + dimension)); reset != nil { + if reset := parseResetHeader(resetHeader); reset != nil { window.ResetUnix = reset window.ResetAt = time.Unix(*reset, 0).UTC().Format(time.RFC3339) } @@ -141,11 +189,27 @@ func parseResetHeader(raw string) *int64 { return nil } if value, err := strconv.ParseInt(raw, 10, 64); err == nil { - if value > 1_000_000_000_000 { + // xAI (and OpenAI-compatible upstreams) may express the reset as a + // millisecond epoch, a second epoch, or a *relative* number of seconds + // until reset (e.g. "60"). Disambiguate by magnitude, mirroring the + // Kiro reset parser, so a relative "60" is not misread as 1970-01-01. + switch { + case value >= 1_000_000_000_000: // milliseconds epoch → seconds value = value / 1000 + case value >= 1_000_000_000: // already a plausible unix-seconds epoch (>= 2001-09) + // keep as-is + default: // relative seconds from now + value = time.Now().Unix() + value } return &value } + if duration, err := time.ParseDuration(raw); err == nil && duration > 0 { + if duration < time.Second { + duration = time.Second + } + value := time.Now().Add(duration).Unix() + return &value + } if t, err := time.Parse(time.RFC3339, raw); err == nil { value := t.Unix() return &value diff --git a/backend/internal/pkg/xai/quota_test.go b/backend/internal/pkg/xai/quota_test.go index 983a87e84f..40dfde4c09 100644 --- a/backend/internal/pkg/xai/quota_test.go +++ b/backend/internal/pkg/xai/quota_test.go @@ -5,6 +5,7 @@ package xai import ( "net/http" "testing" + "time" "github.com/stretchr/testify/require" ) @@ -41,6 +42,104 @@ func TestParseQuotaHeaders(t *testing.T) { require.NotContains(t, snapshot.Headers, "authorization") } +func TestParseQuotaHeadersAcceptsXAITierAliases(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-xai-user-tier", "supergrok-heavy") + headers.Set("x-xai-user-entitlement-status", "enabled") + + snapshot := ParseQuotaHeaders(headers, http.StatusOK) + require.NotNil(t, snapshot) + require.True(t, snapshot.HeadersObserved) + require.Equal(t, "supergrok-heavy", snapshot.SubscriptionTier) + require.Equal(t, "enabled", snapshot.EntitlementStatus) + require.Equal(t, "supergrok-heavy", snapshot.Headers["x-xai-user-tier"]) + require.Equal(t, "enabled", snapshot.Headers["x-xai-user-entitlement-status"]) +} + +func TestParseQuotaHeadersAcceptsRateLimitAliases(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-rate-limit-limit-tokens", "500000") + headers.Set("x-rate-limit-remaining-tokens", "100") + headers.Set("x-rate-limit-reset-tokens", "1893456000") + + snapshot := ParseQuotaHeaders(headers, http.StatusOK) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Tokens) + require.Equal(t, int64(500000), *snapshot.Tokens.Limit) + require.Equal(t, int64(100), *snapshot.Tokens.Remaining) + require.Equal(t, int64(1893456000), *snapshot.Tokens.ResetUnix) + require.Contains(t, snapshot.Headers, "x-rate-limit-limit-tokens") +} + +func TestParseResetHeaderRelativeSecondsNotMisreadAsEpoch(t *testing.T) { + t.Parallel() + + headers := http.Header{} + // xAI may return the reset window as a relative number of seconds ("60"). + // It must resolve to ~now+60s, NOT 1970-01-01 (epoch 60). + headers.Set("x-ratelimit-reset-requests", "60") + headers.Set("x-ratelimit-remaining-requests", "0") + + before := time.Now().Unix() + snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Requests) + require.NotNil(t, snapshot.Requests.ResetUnix) + got := *snapshot.Requests.ResetUnix + require.GreaterOrEqual(t, got, before+59) + require.LessOrEqual(t, got, time.Now().Unix()+61) +} + +func TestParseResetHeaderDurationWindow(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-ratelimit-reset-requests", "6m0s") + headers.Set("x-ratelimit-remaining-requests", "0") + + before := time.Now().Unix() + snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Requests) + require.NotNil(t, snapshot.Requests.ResetUnix) + got := *snapshot.Requests.ResetUnix + require.GreaterOrEqual(t, got, before+359) + require.LessOrEqual(t, got, time.Now().Unix()+361) +} + +func TestParseResetHeaderSubsecondDurationCeilsToFutureSecond(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-rate-limit-reset-tokens", "250ms") + headers.Set("x-rate-limit-remaining-tokens", "0") + + before := time.Now().Unix() + snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Tokens) + require.NotNil(t, snapshot.Tokens.ResetUnix) + require.GreaterOrEqual(t, *snapshot.Tokens.ResetUnix, before) + require.LessOrEqual(t, *snapshot.Tokens.ResetUnix, time.Now().Unix()+2) +} + +func TestParseResetHeaderMillisecondsEpochNormalized(t *testing.T) { + t.Parallel() + + headers := http.Header{} + headers.Set("x-ratelimit-reset-tokens", "1893456000000") // ms epoch + headers.Set("x-ratelimit-remaining-tokens", "0") + + snapshot := ParseQuotaHeaders(headers, http.StatusTooManyRequests) + require.NotNil(t, snapshot) + require.NotNil(t, snapshot.Tokens) + require.Equal(t, int64(1893456000), *snapshot.Tokens.ResetUnix) +} + func TestParseQuotaHeadersReturnsNilForMissingHeaders(t *testing.T) { t.Parallel() @@ -66,6 +165,8 @@ func TestIsGrokFreeRolling24hTokenLimit(t *testing.T) { t.Parallel() require.True(t, IsGrokFreeRolling24hTokenLimit(GrokFreeRolling24hTokenLimit)) + require.True(t, IsGrokFreeRolling24hTokenLimit(500_000)) + require.True(t, IsGrokFreeRolling24hTokenLimit(1_000_000), "observed Free limit variants remain classifiable") require.True(t, IsGrokFreeRolling24hTokenLimit(2_000_000), "legacy snapshots remain classifiable") require.False(t, IsGrokFreeRolling24hTokenLimit(3_000_000)) } diff --git a/backend/internal/pkg/xai/sso_device.go b/backend/internal/pkg/xai/sso_device.go index e533e394d2..825616b93c 100644 --- a/backend/internal/pkg/xai/sso_device.go +++ b/backend/internal/pkg/xai/sso_device.go @@ -8,6 +8,7 @@ import ( "fmt" "io" "net/http" + "net/http/cookiejar" "net/url" "sort" "strconv" @@ -25,6 +26,7 @@ const ( SSOConversionTimeout = 90 * time.Second ssoMaxAuthBody = 2 << 20 + ssoMaxTokenLength = 16 << 10 ssoDefaultUA = "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36" ssoDefaultTokenTTL = 6 * time.Hour ) @@ -51,7 +53,7 @@ type SSODeviceOptions struct { type ssoDeviceFlow struct { client SSODeviceHTTPClient userAgent string - cookies map[string]string + cookieJar http.CookieJar sleep func(context.Context, time.Duration) error } @@ -80,11 +82,16 @@ func ConvertSSOToBuild(ctx context.Context, ssoToken string, opts *SSODeviceOpti if sleep == nil { sleep = sleepContext } + jar, err := cookiejar.New(nil) + if err != nil { + return nil, err + } + seedSSOCookies(jar, ssoToken) flow := &ssoDeviceFlow{ client: client, userAgent: userAgent, - cookies: map[string]string{"sso": ssoToken, "sso-rw": ssoToken}, + cookieJar: jar, sleep: sleep, } return flow.convert(ctx) @@ -253,7 +260,7 @@ func (f *ssoDeviceFlow) do(ctx context.Context, method, endpoint string, form ur request.Header.Set("Accept", "application/json, text/html;q=0.9, */*;q=0.8") request.Header.Set("Accept-Language", "zh-CN,zh;q=0.9,en;q=0.8") request.Header.Set("User-Agent", f.userAgent) - if cookie := f.cookieHeader(); cookie != "" { + if cookie := f.cookieHeader(request.URL); cookie != "" { request.Header.Set("Cookie", cookie) } if currentForm != nil { @@ -264,7 +271,7 @@ func (f *ssoDeviceFlow) do(ctx context.Context, method, endpoint string, form ur if err != nil { return 0, currentURL, nil, err } - f.captureCookies(response) + f.captureCookies(request.URL, response) data, readErr := io.ReadAll(io.LimitReader(response.Body, ssoMaxAuthBody+1)) _ = response.Body.Close() if readErr != nil { @@ -298,30 +305,49 @@ func (f *ssoDeviceFlow) do(ctx context.Context, method, endpoint string, form ur return 0, currentURL, nil, errors.New("xAI OAuth redirected too many times") } -func (f *ssoDeviceFlow) captureCookies(response *http.Response) { +func seedSSOCookies(jar http.CookieJar, token string) { + if jar == nil { + return + } + for _, rawURL := range []string{SSOAccountsURL, OAuthIssuer + "/"} { + target, err := url.Parse(rawURL) + if err != nil { + continue + } + jar.SetCookies(target, []*http.Cookie{ + {Name: "sso", Value: token, Path: "/", Secure: true, HttpOnly: true}, + {Name: "sso-rw", Value: token, Path: "/", Secure: true, HttpOnly: true}, + }) + } +} + +func (f *ssoDeviceFlow) captureCookies(requestURL *url.URL, response *http.Response) { + if f == nil || f.cookieJar == nil || requestURL == nil || response == nil { + return + } + cookies := make([]*http.Cookie, 0) for _, cookie := range response.Cookies() { name := strings.TrimSpace(cookie.Name) value := strings.TrimSpace(cookie.Value) if name == "" || len(name) > 128 || len(value) > 16384 || strings.ContainsAny(name+value, "\r\n\x00") { continue } - if cookie.MaxAge < 0 { - delete(f.cookies, name) - continue - } - f.cookies[name] = value + cookie.Name = name + cookie.Value = value + cookies = append(cookies, cookie) } + f.cookieJar.SetCookies(requestURL, cookies) } -func (f *ssoDeviceFlow) cookieHeader() string { - keys := make([]string, 0, len(f.cookies)) - for key := range f.cookies { - keys = append(keys, key) +func (f *ssoDeviceFlow) cookieHeader(requestURL *url.URL) string { + if f == nil || f.cookieJar == nil || requestURL == nil { + return "" } - sort.Strings(keys) - parts := make([]string, 0, len(keys)) - for _, key := range keys { - parts = append(parts, key+"="+f.cookies[key]) + cookies := f.cookieJar.Cookies(requestURL) + sort.Slice(cookies, func(i, j int) bool { return cookies[i].Name < cookies[j].Name }) + parts := make([]string, 0, len(cookies)) + for _, cookie := range cookies { + parts = append(parts, cookie.Name+"="+cookie.Value) } return strings.Join(parts, "; ") } @@ -363,7 +389,11 @@ func NormalizeSSOToken(value string) string { } func sanitizeSSOToken(value string) string { - return strings.NewReplacer("\r", "", "\n", "", "\x00", "").Replace(strings.TrimSpace(value)) + value = strings.NewReplacer("\r", "", "\n", "", "\x00", "").Replace(strings.TrimSpace(value)) + if len(value) > ssoMaxTokenLength { + return "" + } + return value } func DecodeJWTClaims(token string) map[string]any { diff --git a/backend/internal/pkg/xai/sso_device_test.go b/backend/internal/pkg/xai/sso_device_test.go index 27dff15f4d..cd3bae3aa3 100644 --- a/backend/internal/pkg/xai/sso_device_test.go +++ b/backend/internal/pkg/xai/sso_device_test.go @@ -6,6 +6,7 @@ import ( "context" "io" "net/http" + "net/http/cookiejar" "net/url" "strings" "testing" @@ -25,7 +26,7 @@ func (c *ssoDeviceFakeClient) Do(req *http.Request) (*http.Response, error) { switch req.URL.String() { case SSOAccountsURL: require.Equal(c.t, http.MethodGet, req.Method) - return ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": {"session=web-session; Path=/"}}, `{}`), nil + return ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": {"session=web-session; Domain=x.ai; Path=/"}}, `{}`), nil case SSODeviceURL: require.Equal(c.t, http.MethodPost, req.Method) values := readSSODeviceForm(c.t, req) @@ -92,6 +93,33 @@ func TestNormalizeSSOTokenAcceptsCookieHeader(t *testing.T) { require.Equal(t, "token-1", NormalizeSSOToken("Cookie: foo=bar; sso=token-1; sso-rw=token-2")) require.Equal(t, "token-2", NormalizeSSOToken("sso-rw=token-2; foo=bar")) require.Equal(t, "raw-token", NormalizeSSOToken(" raw-token ; ignored=1")) + require.Empty(t, NormalizeSSOToken(strings.Repeat("x", ssoMaxTokenLength+1))) +} + +func TestSSODeviceCookieJarHonorsDomainAndPath(t *testing.T) { + jar, err := cookiejar.New(nil) + require.NoError(t, err) + flow := &ssoDeviceFlow{cookieJar: jar} + accountsURL, err := url.Parse("https://accounts.x.ai/") + require.NoError(t, err) + authURL, err := url.Parse("https://auth.x.ai/oauth2/device/verify") + require.NoError(t, err) + + flow.captureCookies(accountsURL, ssoDeviceResponse(http.StatusOK, http.Header{"Set-Cookie": { + "host-only=accounts; Path=/", + "shared=all-xai; Domain=x.ai; Path=/", + "narrow=oauth-only; Domain=x.ai; Path=/oauth2", + }}, "")) + + authCookies := flow.cookieHeader(authURL) + require.NotContains(t, authCookies, "host-only=accounts") + require.Contains(t, authCookies, "shared=all-xai") + require.Contains(t, authCookies, "narrow=oauth-only") + + accountsCookies := flow.cookieHeader(accountsURL) + require.Contains(t, accountsCookies, "host-only=accounts") + require.Contains(t, accountsCookies, "shared=all-xai") + require.NotContains(t, accountsCookies, "narrow=oauth-only") } func ssoDeviceResponse(status int, header http.Header, body string) *http.Response { diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index c2773996f4..a1c7cad801 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -190,7 +190,12 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p, + group.FieldVideoModelPrices, group.FieldWebSearchPricePerCall, + group.FieldSearchPricePer1k, + group.FieldAudioRealtimePricePerMin, + group.FieldAudioTtsPricePerMillionChars, + group.FieldAudioSttPricePerHour, group.FieldClaudeCodeOnly, group.FieldFallbackGroupID, group.FieldFallbackGroupIDOnInvalidRequest, @@ -969,7 +974,12 @@ func groupEntityToService(g *dbent.Group) *service.Group { VideoPrice480P: g.VideoPrice480p, VideoPrice720P: g.VideoPrice720p, VideoPrice1080P: g.VideoPrice1080p, + VideoModelPrices: service.NormalizeVideoModelPrices(g.VideoModelPrices), WebSearchPricePerCall: g.WebSearchPricePerCall, + SearchPricePer1k: g.SearchPricePer1k, + AudioRealtimePricePerMin: g.AudioRealtimePricePerMin, + AudioTTSPricePerMillionChars: g.AudioTtsPricePerMillionChars, + AudioSTTPricePerHour: g.AudioSttPricePerHour, DefaultValidityDays: g.DefaultValidityDays, ClaudeCodeOnly: g.ClaudeCodeOnly, FallbackGroupID: g.FallbackGroupID, diff --git a/backend/internal/repository/api_key_repo_messages_dispatch_unit_test.go b/backend/internal/repository/api_key_repo_messages_dispatch_unit_test.go index 4a462ab154..0e079d30d4 100644 --- a/backend/internal/repository/api_key_repo_messages_dispatch_unit_test.go +++ b/backend/internal/repository/api_key_repo_messages_dispatch_unit_test.go @@ -19,6 +19,9 @@ func TestGroupEntityToService_PreservesMessagesDispatchModelConfig(t *testing.T) RateMultiplier: 1, AllowMessagesDispatch: true, DefaultMappedModel: "gpt-5.4", + VideoModelPrices: map[string]map[string]float64{ + service.VideoPriceFamilyGrokImagineVideo15: {service.VideoBillingResolution720P: 0.14}, + }, MessagesDispatchModelConfig: service.OpenAIMessagesDispatchModelConfig{ OpusMappedModel: "gpt-5.4-nano", SonnetMappedModel: "gpt-5.3-codex", @@ -32,6 +35,7 @@ func TestGroupEntityToService_PreservesMessagesDispatchModelConfig(t *testing.T) got := groupEntityToService(group) require.NotNil(t, got) require.Equal(t, group.MessagesDispatchModelConfig, got.MessagesDispatchModelConfig) + require.Equal(t, group.VideoModelPrices, got.VideoModelPrices) } func TestAPIKeyRepository_GetByKeyForAuth_PreservesMessagesDispatchModelConfig_SQLite(t *testing.T) { diff --git a/backend/internal/repository/gateway_cache.go b/backend/internal/repository/gateway_cache.go index 6c345e1a78..dc3d4a980a 100644 --- a/backend/internal/repository/gateway_cache.go +++ b/backend/internal/repository/gateway_cache.go @@ -7,6 +7,7 @@ import ( "errors" "fmt" "strconv" + "strings" "time" "github.com/Wei-Shaw/sub2api/internal/service" @@ -64,6 +65,68 @@ func (c *gatewayCache) DeleteSessionAccountID(ctx context.Context, groupID int64 return c.rdb.Del(ctx, key).Err() } +const ( + grokVideoPendingBillingPrefix = "grok_video_pending:" + grokVideoBilledPrefix = "grok_video_billed:" +) + +func (c *gatewayCache) SetGrokVideoPendingBilling(ctx context.Context, key string, payload []byte, ttl time.Duration) error { + if c == nil || c.rdb == nil { + return errors.New("gateway cache unavailable") + } + key = strings.TrimSpace(key) + if key == "" || len(payload) == 0 { + return errors.New("invalid grok video pending billing payload") + } + if ttl <= 0 { + ttl = 24 * time.Hour + } + return c.rdb.Set(ctx, grokVideoPendingBillingPrefix+key, payload, ttl).Err() +} + +func (c *gatewayCache) GetGrokVideoPendingBilling(ctx context.Context, key string) ([]byte, error) { + if c == nil || c.rdb == nil { + return nil, errors.New("gateway cache unavailable") + } + key = strings.TrimSpace(key) + if key == "" { + return nil, errors.New("invalid grok video pending billing key") + } + val, err := c.rdb.Get(ctx, grokVideoPendingBillingPrefix+key).Bytes() + if err != nil { + if errors.Is(err, redis.Nil) { + return nil, nil + } + return nil, err + } + return val, nil +} + +func (c *gatewayCache) ClaimGrokVideoBilled(ctx context.Context, key string, ttl time.Duration) (bool, error) { + if c == nil || c.rdb == nil { + return false, errors.New("gateway cache unavailable") + } + key = strings.TrimSpace(key) + if key == "" { + return false, errors.New("invalid grok video billed key") + } + if ttl <= 0 { + ttl = 48 * time.Hour + } + return c.rdb.SetNX(ctx, grokVideoBilledPrefix+key, "1", ttl).Result() +} + +func (c *gatewayCache) ReleaseGrokVideoBilled(ctx context.Context, key string) error { + if c == nil || c.rdb == nil { + return errors.New("gateway cache unavailable") + } + key = strings.TrimSpace(key) + if key == "" { + return errors.New("invalid grok video billed key") + } + return c.rdb.Del(ctx, grokVideoBilledPrefix+key).Err() +} + // Compile-time assertion: gatewayCache must implement CyberSessionBlockStore. var _ service.CyberSessionBlockStore = (*gatewayCache)(nil) var _ service.LiveCallStore = (*gatewayCache)(nil) diff --git a/backend/internal/repository/grok_oauth_client.go b/backend/internal/repository/grok_oauth_client.go index 38f6cfb96e..c461122e7b 100644 --- a/backend/internal/repository/grok_oauth_client.go +++ b/backend/internal/repository/grok_oauth_client.go @@ -1,10 +1,15 @@ package repository import ( + "bytes" "context" + "encoding/json" "errors" + "fmt" + "io" "net/http" "net/url" + "os" "strings" "time" @@ -20,8 +25,24 @@ type grokOAuthClient struct { tokenURL string } +const ( + accountsBaseURL = "https://accounts.x.ai" + loginRPCEndpoint = accountsBaseURL + "/api/rpc" + turnstileWebsiteURL = accountsBaseURL + turnstileWebsiteKey = "0x4AAAAAAAhr9JGVDZbrZOo0" + yesCaptchaCreateTask = "https://api.yescaptcha.com/createTask" + yesCaptchaGetResult = "https://api.yescaptcha.com/getTaskResult" +) + func NewGrokOAuthClient() service.GrokOAuthClient { - return &grokOAuthClient{tokenURL: xai.EffectiveTokenURL()} + // Fail closed: never fall back to an unvalidated EffectiveTokenURL (env can + // point at an attacker host and steal code/refresh tokens). + tokenURL, err := xai.ValidatedTokenURL() + if err != nil || strings.TrimSpace(tokenURL) == "" { + // Official allowlisted endpoint only — never EffectiveTokenURL() (raw env). + tokenURL = xai.DefaultTokenURL + } + return &grokOAuthClient{tokenURL: tokenURL} } func (c *grokOAuthClient) ExchangeCode(ctx context.Context, code, codeVerifier, redirectURI, proxyURL, clientID string) (*xai.TokenResponse, error) { @@ -90,6 +111,31 @@ func (c *grokOAuthClient) RefreshToken(ctx context.Context, refreshToken, proxyU return &tokenResp, nil } +// LoginWithPassword authenticates against accounts.x.ai and returns an ephemeral SSO cookie. +// Password and SSO must never be written to account credentials or logs. +func (c *grokOAuthClient) LoginWithPassword(ctx context.Context, email, password, proxyURL string) (*service.GrokPasswordLoginResult, error) { + turnstileToken, err := solveTurnstile(ctx) + if err != nil { + return nil, err + } + httpClient, err := createGrokHTTPClient(proxyURL, true) + if err != nil { + return nil, infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CLIENT_INIT_FAILED", "create HTTP client: %v", err) + } + cookieSetterURL, err := createGrokPasswordSession(ctx, httpClient, strings.TrimSpace(email), password, turnstileToken) + if err != nil { + return nil, err + } + ssoToken, err := extractGrokSSOToken(ctx, httpClient, cookieSetterURL) + if err != nil { + return nil, err + } + return &service.GrokPasswordLoginResult{ + Email: strings.TrimSpace(email), + SSOToken: ssoToken, + }, nil +} + func (c *grokOAuthClient) ConvertSSOToBuild(ctx context.Context, ssoToken, proxyURL string) (*xai.TokenResponse, error) { client, err := createGrokSSOHTTPClient(proxyURL) if err != nil { @@ -167,6 +213,20 @@ func grokOAuthStatusError(code, message string, resp *req.Response) error { func grokOAuthHasExplicitEntitlementDenial(body string) bool { lower := strings.ToLower(body) + // Billing exhaustion is recoverable. xAI may include a generic + // access_denied code alongside the quota message, so it must win over the + // entitlement marker during token refresh. + for _, phrase := range []string{ + "spending limit", + "run out of credits", + "out of credits", + "credits exhausted", + "included free usage", + } { + if strings.Contains(lower, phrase) { + return false + } + } compact := strings.NewReplacer(" ", "", "\n", "", "\r", "", "\t", "").Replace(lower) for _, field := range []string{"error", "code", "reason"} { for _, value := range []string{"access_denied", "entitlement_denied", "subscription_required", "no_active_subscription"} { @@ -179,3 +239,199 @@ func grokOAuthHasExplicitEntitlementDenial(body string) bool { strings.Contains(lower, "subscription required") || strings.Contains(lower, "no active grok subscription") } + +func createGrokHTTPClient(proxyURL string, noRedirect bool) (*http.Client, error) { + transport := &http.Transport{} + if strings.TrimSpace(proxyURL) != "" { + parsed, err := url.Parse(proxyURL) + if err != nil { + return nil, err + } + transport.Proxy = http.ProxyURL(parsed) + } + client := &http.Client{Timeout: 120 * time.Second, Transport: transport} + if noRedirect { + client.CheckRedirect = func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + } + } + return client, nil +} + +func solveTurnstile(ctx context.Context) (string, error) { + clientKey := strings.TrimSpace(os.Getenv("YESCAPTCHA_CLIENT_KEY")) + if clientKey == "" { + clientKey = strings.TrimSpace(os.Getenv("YESCAPTCHA_API_KEY")) + } + if clientKey == "" { + return "", infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_CAPTCHA_KEY_REQUIRED", "yescaptcha client key is required for Grok password authorization") + } + createBody, err := json.Marshal(map[string]any{ + "clientKey": clientKey, + "task": map[string]any{ + "type": "TurnstileTaskProxyless", + "websiteURL": turnstileWebsiteURL, + "websiteKey": turnstileWebsiteKey, + }, + }) + if err != nil { + return "", infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_CAPTCHA_FAILED", "encode captcha create request failed: %v", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, yesCaptchaCreateTask, bytes.NewReader(createBody)) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "build captcha create request failed: %v", err) + } + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "create captcha task failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + var createResp struct { + ErrorID int `json:"errorId"` + TaskID string `json:"taskId"` + ErrorDescription string `json:"errorDescription"` + } + if err := json.NewDecoder(resp.Body).Decode(&createResp); err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "decode captcha create response failed: %v", err) + } + if createResp.ErrorID != 0 || strings.TrimSpace(createResp.TaskID) == "" { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "captcha create failed: %s", createResp.ErrorDescription) + } + deadline := time.Now().Add(90 * time.Second) + for time.Now().Before(deadline) { + select { + case <-ctx.Done(): + return "", ctx.Err() + case <-time.After(5 * time.Second): + } + body, err := json.Marshal(map[string]any{"clientKey": clientKey, "taskId": createResp.TaskID}) + if err != nil { + return "", infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_CAPTCHA_FAILED", "encode captcha poll request failed: %v", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, yesCaptchaGetResult, bytes.NewReader(body)) + if err != nil { + continue + } + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + continue + } + var pollResp struct { + ErrorID int `json:"errorId"` + Status string `json:"status"` + ErrorDescription string `json:"errorDescription"` + Solution struct { + Token string `json:"token"` + } `json:"solution"` + } + err = json.NewDecoder(resp.Body).Decode(&pollResp) + _ = resp.Body.Close() + if err != nil { + continue + } + if pollResp.ErrorID != 0 { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_CAPTCHA_FAILED", "captcha poll failed: %s", pollResp.ErrorDescription) + } + if pollResp.Status == "ready" && strings.TrimSpace(pollResp.Solution.Token) != "" { + return pollResp.Solution.Token, nil + } + } + return "", infraerrors.New(http.StatusGatewayTimeout, "GROK_OAUTH_CAPTCHA_TIMEOUT", "captcha solve timed out") +} + +func createGrokPasswordSession(ctx context.Context, client *http.Client, email, password, turnstileToken string) (string, error) { + payload, err := json.Marshal(map[string]any{ + "rpc": "createSession", + "req": map[string]any{ + "createSessionRequest": map[string]any{ + "credentials": map[string]any{ + "case": "emailAndPassword", + "value": map[string]any{ + "email": email, + "clearTextPassword": password, + }, + }, + }, + "turnstileToken": turnstileToken, + }, + }) + if err != nil { + return "", infraerrors.Newf(http.StatusInternalServerError, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "encode password login request failed: %v", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, loginRPCEndpoint, bytes.NewReader(payload)) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "build password login request failed: %v", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Origin", accountsBaseURL) + req.Header.Set("Referer", accountsBaseURL+"/sign-in?redirect=grok-com&email=true") + req.Header.Set("User-Agent", "Mozilla/5.0") + req.Header.Set("Accept", "*/*") + resp, err := client.Do(req) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "password login request failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + body, _ := io.ReadAll(resp.Body) + if resp.StatusCode != http.StatusOK { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "password login returned status %d: %s", resp.StatusCode, logredact.RedactText(string(body))) + } + var loginResp struct { + CookieSetterURL string `json:"cookieSetterUrl"` + Error string `json:"error"` + } + if err := json.Unmarshal(body, &loginResp); err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "decode password login response failed: %v", err) + } + if strings.TrimSpace(loginResp.Error) != "" { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "password login error: %s", logredact.RedactText(loginResp.Error)) + } + if strings.TrimSpace(loginResp.CookieSetterURL) == "" { + return "", infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "password login did not return cookieSetterUrl") + } + return loginResp.CookieSetterURL, nil +} + +func extractGrokSSOToken(ctx context.Context, client *http.Client, cookieSetterURL string) (string, error) { + safeURL, err := validateGrokCookieSetterURL(cookieSetterURL) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "invalid cookie setter url: %v", err) + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, safeURL.String(), nil) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "build cookie setter request: %v", err) + } + req.Header.Set("User-Agent", "Mozilla/5.0") + req.Header.Set("Accept", "text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8") + req.Header.Set("Referer", accountsBaseURL+"/") + resp, err := client.Do(req) + if err != nil { + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "follow cookie setter url failed: %v", err) + } + defer func() { _ = resp.Body.Close() }() + for _, cookie := range resp.Header.Values("Set-Cookie") { + if token, ok := strings.CutPrefix(cookie, "sso="); ok { + if idx := strings.Index(token, ";"); idx > 0 { + token = token[:idx] + } + return strings.TrimSpace(token), nil + } + } + return "", infraerrors.Newf(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "no sso cookie found in response (status=%d)", resp.StatusCode) +} + +func validateGrokCookieSetterURL(rawURL string) (*url.URL, error) { + parsed, err := url.Parse(strings.TrimSpace(rawURL)) + if err != nil { + return nil, err + } + if parsed.Scheme != "https" || !strings.EqualFold(parsed.Hostname(), "accounts.x.ai") { + return nil, fmt.Errorf("url must use https://accounts.x.ai") + } + if parsed.User != nil || parsed.Port() != "" || parsed.Fragment != "" || parsed.Opaque != "" { + return nil, fmt.Errorf("url contains disallowed authority or fragment components") + } + return parsed, nil +} diff --git a/backend/internal/repository/grok_oauth_client_test.go b/backend/internal/repository/grok_oauth_client_test.go index 690a8e2139..482c3065b0 100644 --- a/backend/internal/repository/grok_oauth_client_test.go +++ b/backend/internal/repository/grok_oauth_client_test.go @@ -47,6 +47,8 @@ func TestGrokOAuthClientExchangeAndRefreshUseFormFields(t *testing.T) { } })) defer server.Close() + // Tests inject a loopback token endpoint; allowlist requires unsafe override. + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") t.Setenv(xai.EnvTokenURL, server.URL) client := NewGrokOAuthClient() @@ -80,6 +82,7 @@ func TestGrokOAuthClientRefreshForbiddenClassifiesOnlyExplicitEntitlement(t *tes _, _ = w.Write([]byte(tt.body)) })) defer server.Close() + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") t.Setenv(xai.EnvTokenURL, server.URL) client := NewGrokOAuthClient() @@ -96,6 +99,7 @@ func TestGrokOAuthClientStatusErrorRedactsSensitiveResponseBody(t *testing.T) { _, _ = w.Write([]byte(`{"error":"invalid_grant","access_token":"access-secret","refresh_token":"refresh-secret","code_verifier":"verifier-secret"}`)) })) defer server.Close() + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") t.Setenv(xai.EnvTokenURL, server.URL) client := NewGrokOAuthClient() @@ -118,4 +122,15 @@ func TestGrokOAuthEntitlementDenialRequiresExplicitEvidence(t *testing.T) { require.True(t, grokOAuthHasExplicitEntitlementDenial(`{"message":"no active Grok subscription"}`)) require.False(t, grokOAuthHasExplicitEntitlementDenial(`{"error":"forbidden","message":"request forbidden"}`)) require.False(t, grokOAuthHasExplicitEntitlementDenial(`403 Forbidden`)) + require.False(t, grokOAuthHasExplicitEntitlementDenial(`{"error":"access_denied","message":"You have run out of credits"}`)) + require.False(t, grokOAuthHasExplicitEntitlementDenial(`{"code":"subscription_required","message":"included free usage exhausted"}`)) +} + +func TestNewGrokOAuthClient_UnvalidatedTokenURLFallsBackToDefault(t *testing.T) { + // Without unsafe overrides, a random env TokenURL must not be used (fail-closed). + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "") + t.Setenv(xai.EnvTokenURL, "https://evil.example/oauth/token") + + client := NewGrokOAuthClient().(*grokOAuthClient) + require.Equal(t, xai.DefaultTokenURL, client.tokenURL) } diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 31724f7c55..4cbc27909d 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -82,7 +82,12 @@ func createGroupRecord(ctx context.Context, client *dbent.Client, groupIn *servi SetNillableVideoPrice480p(groupIn.VideoPrice480P). SetNillableVideoPrice720p(groupIn.VideoPrice720P). SetNillableVideoPrice1080p(groupIn.VideoPrice1080P). + SetVideoModelPrices(service.NormalizeVideoModelPrices(groupIn.VideoModelPrices)). SetNillableWebSearchPricePerCall(groupIn.WebSearchPricePerCall). + SetNillableSearchPricePer1k(groupIn.SearchPricePer1k). + SetNillableAudioRealtimePricePerMin(groupIn.AudioRealtimePricePerMin). + SetNillableAudioTtsPricePerMillionChars(groupIn.AudioTTSPricePerMillionChars). + SetNillableAudioSttPricePerHour(groupIn.AudioSTTPricePerHour). SetDefaultValidityDays(groupIn.DefaultValidityDays). SetClaudeCodeOnly(groupIn.ClaudeCodeOnly). SetNillableFallbackGroupID(groupIn.FallbackGroupID). @@ -254,6 +259,7 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er SetNillableVideoPrice480p(groupIn.VideoPrice480P). SetNillableVideoPrice720p(groupIn.VideoPrice720P). SetNillableVideoPrice1080p(groupIn.VideoPrice1080P). + SetVideoModelPrices(service.NormalizeVideoModelPrices(groupIn.VideoModelPrices)). SetDefaultValidityDays(groupIn.DefaultValidityDays). SetClaudeCodeOnly(groupIn.ClaudeCodeOnly). SetModelRoutingEnabled(groupIn.ModelRoutingEnabled). @@ -327,6 +333,26 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er } else { builder = builder.ClearWebSearchPricePerCall() } + if groupIn.SearchPricePer1k != nil { + builder = builder.SetSearchPricePer1k(*groupIn.SearchPricePer1k) + } else { + builder = builder.ClearSearchPricePer1k() + } + if groupIn.AudioRealtimePricePerMin != nil { + builder = builder.SetAudioRealtimePricePerMin(*groupIn.AudioRealtimePricePerMin) + } else { + builder = builder.ClearAudioRealtimePricePerMin() + } + if groupIn.AudioTTSPricePerMillionChars != nil { + builder = builder.SetAudioTtsPricePerMillionChars(*groupIn.AudioTTSPricePerMillionChars) + } else { + builder = builder.ClearAudioTtsPricePerMillionChars() + } + if groupIn.AudioSTTPricePerHour != nil { + builder = builder.SetAudioSttPricePerHour(*groupIn.AudioSTTPricePerHour) + } else { + builder = builder.ClearAudioSttPricePerHour() + } // 处理 FallbackGroupID:nil 时清除,否则设置 if groupIn.FallbackGroupID != nil { diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go index 68a986648a..904ff09eb9 100644 --- a/backend/internal/repository/http_upstream.go +++ b/backend/internal/repository/http_upstream.go @@ -23,6 +23,7 @@ import ( "github.com/andybalholm/brotli" "github.com/klauspost/compress/zstd" + "golang.org/x/mod/semver" "golang.org/x/net/http2" "github.com/Wei-Shaw/sub2api/internal/config" @@ -33,7 +34,6 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" - "golang.org/x/mod/semver" ) // 默认配置常量 @@ -55,6 +55,18 @@ const ( // defaultResponseHeaderTimeout: 默认等待响应头超时时间(5分钟) // LLM 请求可能排队较久,需要较长超时 defaultResponseHeaderTimeout = 300 * time.Second + // defaultUpstreamDialTimeout: 默认 TCP/DNS 建连超时(10秒) + // Transport 不设置 DialContext 时会退化为零值 net.Dialer(无超时),建连阶段 + // 只能依赖内核默认 TCP 重传(Linux 约 130 秒)。ResponseHeaderTimeout 只约束 + // 连接建立之后等待响应头的阶段,覆盖不到 DNS 解析与 TCP 握手。 + // 上游域名被解析到 443 不可达的 IP 时(DNS 污染/路由异常),单个账号就要卡满 + // 内核超时;而多账号故障转移是串行的,一次请求会阻塞数分钟且不写中间错误。 + defaultUpstreamDialTimeout = 10 * time.Second + // defaultUpstreamDialKeepAlive: TCP keepalive 探测间隔,与 Go 默认值保持一致 + defaultUpstreamDialKeepAlive = 30 * time.Second + // defaultUpstreamTLSHandshakeTimeout: TLS 握手超时(10秒) + // 与建连超时同量级,避免 TCP 已连通但对端不推进握手时无限等待 + defaultUpstreamTLSHandshakeTimeout = 10 * time.Second // defaultMaxUpstreamClients: 默认最大客户端缓存数量 // 超出后会淘汰最久未使用的客户端 defaultMaxUpstreamClients = 5000 @@ -73,12 +85,12 @@ const ( openAIHTTP2PingTimeout = 15 * time.Second // The Grok CLI proxy rejects requests that do not identify a supported - // client version. Keep a known-good stable version in the binary while - // allowing operators to bump it without waiting for a Sub2API release. - grokCLIProxyHost = "cli-chat-proxy.grok.com" + // client version. Host/env/version pins live in package xai so service, + // billing, and transport layers advertise the same identity. + grokCLIProxyHost = xai.CLIProxyHost grokOfficialAPIHost = "api.x.ai" - grokCLIStableVersion = xai.CLIClientVersion - grokCLIVersionOverride = "XAI_GROK_CLI_VERSION" + grokCLIStableVersion = xai.CLIClientVersion // preferred pin (not the minimum floor) + grokCLIVersionOverride = xai.CLIVersionEnv grokFallbackBodyLimit = 64 << 10 ) @@ -438,6 +450,11 @@ type prefixedReadCloser struct { // the final shared transport boundary. Keying this behavior to the exact CLI // proxy host keeps direct api.x.ai traffic unchanged and automatically covers // Responses, Chat Completions, media, quota probes, and account tests. +// +// Operator overrides must be >= CLIClientVersion (the preferred pin). Package +// xai.IsSupportedCLIVersion uses a lower floor (CLIStableVersion) for general +// validation; transport is stricter so we never silently advertise an older pin +// than the binary default. func applyGrokCLIProxyHeaders(req *http.Request) { if req == nil || req.URL == nil || !strings.EqualFold(strings.TrimSpace(req.URL.Hostname()), grokCLIProxyHost) { return @@ -449,14 +466,15 @@ func applyGrokCLIProxyHeaders(req *http.Request) { if !isSupportedGrokCLIVersion(version) { version = grokCLIStableVersion } - req.Header.Set("X-XAI-Token-Auth", "xai-grok-cli") + req.Header.Set("X-XAI-Token-Auth", xai.CLITokenAuth) req.Header.Set("x-grok-client-version", version) - req.Header.Set("User-Agent", "xai-grok-workspace/"+version) + req.Header.Set("x-grok-client-identifier", xai.CLIClientIdentifier) + req.Header.Set("User-Agent", xai.CLIUserAgent(version)) } func isSupportedGrokCLIVersion(version string) bool { canonical := "v" + version - minimum := "v" + grokCLIStableVersion + minimum := "v" + xai.CLIClientVersion return semver.IsValid(canonical) && semver.Canonical(canonical) == canonical && semver.Compare(canonical, minimum) >= 0 @@ -1246,6 +1264,17 @@ func defaultPoolSettings(cfg *config.Config) poolSettings { } } +// newUpstreamDialer 构建上游 Transport 的 TCP dialer。 +// +// 必须显式提供:http.Transport 的 DialContext 为 nil 时使用零值 net.Dialer, +// 建连没有任何超时上限,只能等内核 TCP 重传耗尽(Linux 约 130 秒)。 +func newUpstreamDialer() *net.Dialer { + return &net.Dialer{ + Timeout: defaultUpstreamDialTimeout, + KeepAlive: defaultUpstreamDialKeepAlive, + } +} + // buildUpstreamTransport 构建上游请求的 Transport // 使用配置文件中的连接池参数,支持生产环境调优 // @@ -1258,6 +1287,8 @@ func defaultPoolSettings(cfg *config.Config) poolSettings { // - error: 代理配置错误 // // Transport 参数说明: +// - DialContext: DNS 解析 + TCP 建连超时(不设置则无上限,退化为内核默认重传) +// - TLSHandshakeTimeout: TLS 握手超时 // - MaxIdleConns: 所有主机的最大空闲连接总数 // - MaxIdleConnsPerHost: 每主机最大空闲连接数(影响连接复用率) // - MaxConnsPerHost: 每主机最大连接数(达到后新请求等待) @@ -1265,6 +1296,8 @@ func defaultPoolSettings(cfg *config.Config) poolSettings { // - ResponseHeaderTimeout: 等待响应头超时(不影响流式传输) func buildUpstreamTransport(settings poolSettings, proxyURL *url.URL, protocolMode string) (*http.Transport, error) { transport := &http.Transport{ + DialContext: newUpstreamDialer().DialContext, + TLSHandshakeTimeout: defaultUpstreamTLSHandshakeTimeout, MaxIdleConns: settings.maxIdleConns, MaxIdleConnsPerHost: settings.maxIdleConnsPerHost, MaxConnsPerHost: settings.maxConnsPerHost, diff --git a/backend/internal/repository/http_upstream_dial_timeout_test.go b/backend/internal/repository/http_upstream_dial_timeout_test.go new file mode 100644 index 0000000000..00d6ce1a41 --- /dev/null +++ b/backend/internal/repository/http_upstream_dial_timeout_test.go @@ -0,0 +1,75 @@ +package repository + +import ( + "context" + "net" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// 回归:上游 Transport 必须显式配置建连超时。 +// +// http.Transport.DialContext 为 nil 时 Go 使用零值 net.Dialer(Timeout=0), +// DNS 解析与 TCP 握手没有任何上限,只能等内核重传耗尽(Linux 约 130 秒)。 +// ResponseHeaderTimeout 只覆盖连接建立之后的阶段,管不到建连。 +// 上游域名被解析到不可达 IP 时,串行的多账号故障转移会把一次请求拖到数分钟。 +func TestBuildUpstreamTransportSetsDialTimeout(t *testing.T) { + settings := defaultPoolSettings(nil) + + transport, err := buildUpstreamTransport(settings, nil, upstreamProtocolModeDefault) + require.NoError(t, err) + require.NotNil(t, transport.DialContext, "DialContext 缺失会退化为无超时的零值 dialer") + require.Equal(t, defaultUpstreamTLSHandshakeTimeout, transport.TLSHandshakeTimeout) +} + +func TestNewUpstreamDialerHasBoundedTimeout(t *testing.T) { + dialer := newUpstreamDialer() + + require.Greater(t, dialer.Timeout, time.Duration(0), "建连超时必须有上限") + require.Equal(t, defaultUpstreamDialTimeout, dialer.Timeout) + require.Equal(t, defaultUpstreamDialKeepAlive, dialer.KeepAlive) +} + +// 建连超时对 HTTP 代理同样生效:Transport.Proxy 走的仍是 DialContext, +// 代理地址不可达时必须快速失败而不是挂满内核超时。 +func TestBuildUpstreamTransportKeepsDialTimeoutWithHTTPProxy(t *testing.T) { + proxyURL, err := url.Parse("http://127.0.0.1:1080") + require.NoError(t, err) + + transport, err := buildUpstreamTransport(defaultPoolSettings(nil), proxyURL, upstreamProtocolModeDefault) + require.NoError(t, err) + require.NotNil(t, transport.Proxy) + require.NotNil(t, transport.DialContext) +} + +// SOCKS5 分支会覆盖 Transport.DialContext,覆盖后仍必须是有超时的拨号器。 +func TestBuildUpstreamTransportKeepsDialContextWithSOCKS5Proxy(t *testing.T) { + proxyURL, err := url.Parse("socks5h://127.0.0.1:1080") + require.NoError(t, err) + + transport, err := buildUpstreamTransport(defaultPoolSettings(nil), proxyURL, upstreamProtocolModeDefault) + require.NoError(t, err) + require.NotNil(t, transport.DialContext) +} + +// Timeout 字段确实被 net.Dialer 用于建连:拨一个已被 close 的本地监听端口, +// 断言 Dialer 走的是自己的超时路径而不是无限等待。 +// (不依赖外网可达性,CI 中确定性执行。) +func TestUpstreamDialerRespectsContextCancellation(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + addr := listener.Addr().String() + require.NoError(t, listener.Close()) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + conn, err := newUpstreamDialer().DialContext(ctx, "tcp", addr) + if conn != nil { + _ = conn.Close() + } + require.Error(t, err, "已取消的 context 必须立即中止拨号") +} diff --git a/backend/internal/repository/migrations_runner.go b/backend/internal/repository/migrations_runner.go index cb29270b39..169a01675a 100644 --- a/backend/internal/repository/migrations_runner.go +++ b/backend/internal/repository/migrations_runner.go @@ -57,6 +57,8 @@ const schedulerOutboxPendingDedupKeyMigration = "153_scheduler_outbox_pending_de const schedulerOutboxPendingDedupKeyIndex = "idx_scheduler_outbox_pending_dedup_key" const latestAPIKeyIPIndexMigration = "174_add_usage_logs_api_key_latest_ip_index_notx.sql" const latestAPIKeyIPIndex = "idx_usage_logs_api_key_latest_ip" +const usageLogsUpstreamModelMismatchIndexMigration = "195_add_usage_log_upstream_model_mismatch_index_notx.sql" +const usageLogsUpstreamModelMismatchIndex = "idx_usage_logs_upstream_model_mismatch_created_at" type migrationChecksumCompatibilityRule struct { fileChecksum string @@ -84,6 +86,11 @@ var migrationChecksumCompatibilityRules = map[string]migrationChecksumCompatibil // 195 originally seeded mode=v2; flipped to v1 (safe default / opt-in v2). Existing DBs // that already applied the v2 seed keep their row and the historical checksum. "195_channel_monitor_mode.sql": newMigrationChecksumCompatibilityRule("13f3792f3e3e53ee96e26415c884cf8062c77172824b54fcc9a8c0c2b1f185ec", "4c74fe33ef2274cc72e1bb49671e651274532c034b29f5b2982c2a4c88d101a6"), + // 220 originally cleared video prices for all non-grok platforms (including composite); + // composite is now preserved because it may route to Grok accounts. + "220_clear_non_grok_video_generation_config.sql": newMigrationChecksumCompatibilityRule("85e320b9ec64f2d3fcd8cf705b2b4e76a7b49f7a57140c14bff97f32691c818b", "3da48c8fdffe6390325f43d08b8e353e0a365df43d44a78dbbe655d0deb18402"), + "219_group_search_price_per_1k.sql": newMigrationChecksumCompatibilityRule("e86786ebcc3b14206fd2d321380a4e50e80cdadbfcf4962c639255e6a14008db", "df6ffd71b97e30ec2c8fe7b95e15783042dea58c553e32701ee7c42a5619af80"), + "218_group_audio_voice_pricing.sql": newMigrationChecksumCompatibilityRule("40ee9f3a2af0e0a5e99dabc878fd0fe98be1011f26bcfcefcac7197f7081f0e7", "c2a5e5b4ffd6968ad1c10593289fbc11192cdea19fec3ed9bce3a84eff9a8351"), } // ApplyMigrations 将嵌入的 SQL 迁移文件应用到指定的数据库。 @@ -286,6 +293,8 @@ func prepareNonTransactionalMigration(ctx context.Context, db migrationConnectio return dropInvalidIndexIfPresent(ctx, db, schedulerOutboxPendingDedupKeyIndex) case latestAPIKeyIPIndexMigration: return dropInvalidIndexIfPresent(ctx, db, latestAPIKeyIPIndex) + case usageLogsUpstreamModelMismatchIndexMigration: + return dropInvalidIndexIfPresent(ctx, db, usageLogsUpstreamModelMismatchIndex) default: return nil } diff --git a/backend/internal/repository/migrations_runner_notx_test.go b/backend/internal/repository/migrations_runner_notx_test.go index 4ab6274b92..55588bb043 100644 --- a/backend/internal/repository/migrations_runner_notx_test.go +++ b/backend/internal/repository/migrations_runner_notx_test.go @@ -155,6 +155,42 @@ CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_api_key_latest_ip require.NoError(t, mock.ExpectationsWereMet()) } +func TestApplyMigrationsFS_NonTransactionalMigration_UsageModelMismatchIndexDropsInvalidIndexBeforeRetry(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer func() { _ = db.Close() }() + + prepareMigrationsBootstrapExpectations(mock) + mock.ExpectQuery("SELECT checksum FROM schema_migrations WHERE filename = \\$1"). + WithArgs(usageLogsUpstreamModelMismatchIndexMigration). + WillReturnError(sql.ErrNoRows) + mock.ExpectQuery("SELECT EXISTS \\("). + WithArgs(usageLogsUpstreamModelMismatchIndex). + WillReturnRows(sqlmock.NewRows([]string{"exists"}).AddRow(true)) + mock.ExpectExec("DROP INDEX CONCURRENTLY IF EXISTS idx_usage_logs_upstream_model_mismatch_created_at"). + WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec("CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_upstream_model_mismatch_created_at"). + WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec("INSERT INTO schema_migrations \\(filename, checksum\\) VALUES \\(\\$1, \\$2\\)"). + WithArgs(usageLogsUpstreamModelMismatchIndexMigration, sqlmock.AnyArg()). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("SELECT pg_advisory_unlock\\(\\$1\\)"). + WithArgs(migrationsAdvisoryLockID). + WillReturnResult(sqlmock.NewResult(0, 1)) + + fsys := fstest.MapFS{ + usageLogsUpstreamModelMismatchIndexMigration: &fstest.MapFile{Data: []byte(` +CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_upstream_model_mismatch_created_at + ON usage_logs (created_at DESC, id DESC) + WHERE upstream_model_mismatch IS TRUE; +`)}, + } + + err = applyMigrationsFS(context.Background(), db, fsys) + require.NoError(t, err) + require.NoError(t, mock.ExpectationsWereMet()) +} + func TestApplyMigrationsFS_PaymentOrdersOutTradeNoUniqueMigration_FailsFastOnDuplicatePrecheck(t *testing.T) { db, mock, err := sqlmock.New() require.NoError(t, err) diff --git a/backend/internal/repository/migrations_schema_integration_test.go b/backend/internal/repository/migrations_schema_integration_test.go index b393bc15b6..e4af36ae25 100644 --- a/backend/internal/repository/migrations_schema_integration_test.go +++ b/backend/internal/repository/migrations_schema_integration_test.go @@ -76,6 +76,24 @@ func TestMigrationsRunner_IsIdempotent_AndSchemaIsUpToDate(t *testing.T) { requireColumn(t, tx, "usage_logs", "video_count", "integer", 0, false) requireColumn(t, tx, "usage_logs", "video_resolution", "character varying", 10, true) requireColumn(t, tx, "usage_logs", "video_duration_seconds", "integer", 0, true) + requireColumn(t, tx, "usage_logs", "upstream_response_model", "character varying", 200, true) + requireColumn(t, tx, "usage_logs", "upstream_model_mismatch", "boolean", 0, true) + requireIndex(t, tx, "usage_logs", usageLogsUpstreamModelMismatchIndex) + + var mismatchIndexDef string + require.NoError(t, tx.QueryRowContext(context.Background(), ` +SELECT pg_get_indexdef(i.indexrelid) +FROM pg_class idx +JOIN pg_index i ON i.indexrelid = idx.oid +JOIN pg_class tbl ON tbl.oid = i.indrelid +JOIN pg_namespace ns ON ns.oid = tbl.relnamespace +WHERE ns.nspname = 'public' + AND tbl.relname = 'usage_logs' + AND idx.relname = $1 +`, usageLogsUpstreamModelMismatchIndex).Scan(&mismatchIndexDef)) + require.Contains(t, mismatchIndexDef, "created_at DESC") + require.Contains(t, mismatchIndexDef, "id DESC") + require.Contains(t, mismatchIndexDef, "WHERE (upstream_model_mismatch IS TRUE)") requireConstraintDefinitionContains( t, tx, diff --git a/backend/internal/repository/usage_log_repo_insert.go b/backend/internal/repository/usage_log_repo_insert.go index 809b0652b5..6d932c0d43 100644 --- a/backend/internal/repository/usage_log_repo_insert.go +++ b/backend/internal/repository/usage_log_repo_insert.go @@ -31,6 +31,8 @@ var usageLogInsertArgTypes = [...]string{ "text", // model "text", // requested_model "text", // upstream_model + "text", // upstream_response_model + "boolean", // upstream_model_mismatch "bigint", // group_id "bigint", // subscription_id "integer", // input_tokens @@ -227,6 +229,8 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor, model, requested_model, upstream_model, + upstream_response_model, + upstream_model_mismatch, group_id, subscription_id, input_tokens, @@ -278,12 +282,12 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor, session_id, created_at ) VALUES ( - $1, $2, $3, $4, $5, $6, $7, - $8, $9, - $10, $11, $12, $13, - $14, $15, $16, $17, - $18, $19, $20, $21, $22, $23, - $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54, $55, $56, $57 + $1, $2, $3, $4, $5, $6, $7, $8, $9, + $10, $11, + $12, $13, $14, $15, + $16, $17, $18, $19, + $20, $21, $22, $23, $24, $25, + $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54, $55, $56, $57, $58, $59 ) ON CONFLICT (request_id, api_key_id) DO NOTHING RETURNING id, created_at @@ -682,6 +686,8 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage model, requested_model, upstream_model, + upstream_response_model, + upstream_model_mismatch, group_id, subscription_id, input_tokens, @@ -734,9 +740,9 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage created_at ) AS (VALUES `) - // Each batch row prepends the synthetic input_index before the 57 + // Each batch row prepends the synthetic input_index before the 59 // usage-log column values. - args := make([]any, 0, len(keys)*58) + args := make([]any, 0, len(keys)*60) argPos := 1 for idx, key := range keys { if idx > 0 { @@ -772,6 +778,8 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage model, requested_model, upstream_model, + upstream_response_model, + upstream_model_mismatch, group_id, subscription_id, input_tokens, @@ -831,6 +839,8 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage model, requested_model, upstream_model, + upstream_response_model, + upstream_model_mismatch, group_id, subscription_id, input_tokens, @@ -930,6 +940,8 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( model, requested_model, upstream_model, + upstream_response_model, + upstream_model_mismatch, group_id, subscription_id, input_tokens, @@ -982,7 +994,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( created_at ) AS (VALUES `) - args := make([]any, 0, len(preparedList)*57) + args := make([]any, 0, len(preparedList)*59) argPos := 1 for idx, prepared := range preparedList { if idx > 0 { @@ -1015,6 +1027,8 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( model, requested_model, upstream_model, + upstream_response_model, + upstream_model_mismatch, group_id, subscription_id, input_tokens, @@ -1074,6 +1088,8 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) ( model, requested_model, upstream_model, + upstream_response_model, + upstream_model_mismatch, group_id, subscription_id, input_tokens, @@ -1141,6 +1157,8 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared model, requested_model, upstream_model, + upstream_response_model, + upstream_model_mismatch, group_id, subscription_id, input_tokens, @@ -1192,12 +1210,12 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared session_id, created_at ) VALUES ( - $1, $2, $3, $4, $5, $6, $7, - $8, $9, - $10, $11, $12, $13, - $14, $15, $16, $17, - $18, $19, $20, $21, $22, $23, - $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54, $55, $56, $57 + $1, $2, $3, $4, $5, $6, $7, $8, $9, + $10, $11, + $12, $13, $14, $15, + $16, $17, $18, $19, + $20, $21, $22, $23, $24, $25, + $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54, $55, $56, $57, $58, $59 ) ON CONFLICT (request_id, api_key_id) DO NOTHING `, prepared.args...) @@ -1244,6 +1262,8 @@ func prepareUsageLogInsert(log *service.UsageLog) usageLogInsertPrepared { requestedModel = strings.TrimSpace(log.Model) } upstreamModel := nullString(log.UpstreamModel) + upstreamResponseModel := nullString(log.UpstreamResponseModel) + upstreamModelMismatch := nullBool(log.UpstreamModelMismatch) var requestIDArg any if requestID != "" { @@ -1263,6 +1283,8 @@ func prepareUsageLogInsert(log *service.UsageLog) usageLogInsertPrepared { log.Model, nullString(&requestedModel), upstreamModel, + upstreamResponseModel, + upstreamModelMismatch, groupID, subscriptionID, log.InputTokens, diff --git a/backend/internal/repository/usage_log_repo_query.go b/backend/internal/repository/usage_log_repo_query.go index 084b79ad97..f2fa6e30b0 100644 --- a/backend/internal/repository/usage_log_repo_query.go +++ b/backend/internal/repository/usage_log_repo_query.go @@ -19,7 +19,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/service" ) -const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, image_input_tokens, image_input_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, session_id, created_at" +const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, upstream_response_model, upstream_model_mismatch, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, image_input_tokens, image_input_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, session_id, created_at" func (r *usageLogRepository) GetByID(ctx context.Context, id int64) (log *service.UsageLog, err error) { query := "SELECT " + usageLogSelectColumns + " FROM usage_logs WHERE id = $1" @@ -127,6 +127,9 @@ func (r *usageLogRepository) ListWithFilters(ctx context.Context, params paginat args = append(args, int16(*filters.BillingType)) } conditions, args = appendUsageLogBillingModeWhereCondition(conditions, args, filters.BillingMode) + if filters.UpstreamModelMismatch != nil { + conditions = append(conditions, upstreamModelMismatchCondition("upstream_model_mismatch", *filters.UpstreamModelMismatch)) + } if filters.StartTime != nil { conditions = append(conditions, fmt.Sprintf("created_at >= $%d", len(args)+1)) args = append(args, *filters.StartTime) @@ -157,6 +160,13 @@ func (r *usageLogRepository) ListWithFilters(ctx context.Context, params paginat return logs, page, nil } +func upstreamModelMismatchCondition(column string, mismatch bool) string { + if mismatch { + return column + " IS TRUE" + } + return column + " IS FALSE" +} + func shouldUseFastUsageLogTotal(filters UsageLogFilters) bool { if filters.ExactTotal { return false @@ -437,6 +447,8 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e model string requestedModel sql.NullString upstreamModel sql.NullString + upstreamResponseModel sql.NullString + upstreamModelMismatch sql.NullBool groupID sql.NullInt64 subscriptionID sql.NullInt64 inputTokens int @@ -498,6 +510,8 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e &model, &requestedModel, &upstreamModel, + &upstreamResponseModel, + &upstreamModelMismatch, &groupID, &subscriptionID, &inputTokens, @@ -651,6 +665,13 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e if upstreamModel.Valid { log.UpstreamModel = &upstreamModel.String } + if upstreamResponseModel.Valid { + log.UpstreamResponseModel = &upstreamResponseModel.String + } + if upstreamModelMismatch.Valid { + value := upstreamModelMismatch.Bool + log.UpstreamModelMismatch = &value + } if channelID.Valid { value := channelID.Int64 log.ChannelID = &value @@ -703,6 +724,13 @@ func nullString(v *string) sql.NullString { return sql.NullString{String: *v, Valid: true} } +func nullBool(v *bool) sql.NullBool { + if v == nil { + return sql.NullBool{} + } + return sql.NullBool{Bool: *v, Valid: true} +} + func nullStringIntMapJSON(v map[string]int) any { if len(v) == 0 { return nil diff --git a/backend/internal/repository/usage_log_repo_request_type_test.go b/backend/internal/repository/usage_log_repo_request_type_test.go index 047d0bf65f..a2b82dfa60 100644 --- a/backend/internal/repository/usage_log_repo_request_type_test.go +++ b/backend/internal/repository/usage_log_repo_request_type_test.go @@ -48,6 +48,8 @@ func TestUsageLogRepositoryCreateSyncRequestTypeAndLegacyFields(t *testing.T) { log.Model, log.RequestedModel, sqlmock.AnyArg(), // upstream_model + sqlmock.AnyArg(), // upstream_response_model + sqlmock.AnyArg(), // upstream_model_mismatch sqlmock.AnyArg(), // group_id sqlmock.AnyArg(), // subscription_id log.InputTokens, @@ -137,9 +139,11 @@ func TestUsageLogRepositoryCreate_PersistsServiceTier(t *testing.T) { log.RequestID, log.Model, log.RequestedModel, - sqlmock.AnyArg(), - sqlmock.AnyArg(), - sqlmock.AnyArg(), + sqlmock.AnyArg(), // upstream_model + sqlmock.AnyArg(), // upstream_response_model + sqlmock.AnyArg(), // upstream_model_mismatch + sqlmock.AnyArg(), // group_id + sqlmock.AnyArg(), // subscription_id log.InputTokens, log.OutputTokens, log.CacheCreationTokens, @@ -211,8 +215,8 @@ func TestBuildUsageLogBestEffortInsertQuery_IncludesRequestedModelColumn(t *test query, args := buildUsageLogBestEffortInsertQuery([]usageLogInsertPrepared{prepared}) require.Contains(t, query, "INSERT INTO usage_logs (") - require.Contains(t, query, "\n\t\t\tmodel,\n\t\t\trequested_model,\n\t\t\tupstream_model,") - require.Contains(t, query, "\n\t\t\trequest_id,\n\t\t\tmodel,\n\t\t\trequested_model,\n\t\t\tupstream_model,") + require.Contains(t, query, "\n\t\t\tmodel,\n\t\t\trequested_model,\n\t\t\tupstream_model,\n\t\t\tupstream_response_model,\n\t\t\tupstream_model_mismatch,") + require.Contains(t, query, "\n\t\t\trequest_id,\n\t\t\tmodel,\n\t\t\trequested_model,\n\t\t\tupstream_model,\n\t\t\tupstream_response_model,\n\t\t\tupstream_model_mismatch,") require.Len(t, args, len(prepared.args)) require.Equal(t, prepared.args[5], args[5]) } @@ -273,11 +277,11 @@ func TestPrepareUsageLogInsert_PersistsImageSizeMetadata(t *testing.T) { CreatedAt: time.Date(2025, 1, 6, 12, 0, 0, 0, time.UTC), }) - require.Equal(t, sql.NullString{String: imageSize, Valid: true}, prepared.args[36]) - require.Equal(t, sql.NullString{String: inputSize, Valid: true}, prepared.args[37]) - require.Equal(t, sql.NullString{String: outputSize, Valid: true}, prepared.args[38]) - require.Equal(t, sql.NullString{String: source, Valid: true}, prepared.args[39]) - breakdownJSON, ok := prepared.args[40].(string) + require.Equal(t, sql.NullString{String: imageSize, Valid: true}, prepared.args[38]) + require.Equal(t, sql.NullString{String: inputSize, Valid: true}, prepared.args[39]) + require.Equal(t, sql.NullString{String: outputSize, Valid: true}, prepared.args[40]) + require.Equal(t, sql.NullString{String: source, Valid: true}, prepared.args[41]) + breakdownJSON, ok := prepared.args[42].(string) require.True(t, ok) require.JSONEq(t, `{"1K":1,"4K":1}`, breakdownJSON) } @@ -809,6 +813,8 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { "gpt-image-2", sql.NullString{Valid: true, String: "gpt-image-2"}, sql.NullString{}, + sql.NullString{}, + sql.NullBool{}, sql.NullInt64{}, sql.NullInt64{}, 0, 0, 0, 0, 0, 0, @@ -872,6 +878,8 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { "gpt-5", // model sql.NullString{Valid: true, String: "gpt-5"}, // requested_model sql.NullString{}, // upstream_model + sql.NullString{}, // upstream_response_model + sql.NullBool{}, // upstream_model_mismatch sql.NullInt64{}, // group_id sql.NullInt64{}, // subscription_id 1, // input_tokens @@ -942,6 +950,8 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { "gpt-5", sql.NullString{Valid: true, String: "gpt-5"}, sql.NullString{}, + sql.NullString{}, + sql.NullBool{}, sql.NullInt64{}, sql.NullInt64{}, 1, 2, 3, 4, 5, 6, @@ -1000,6 +1010,8 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) { "gpt-5.4", sql.NullString{Valid: true, String: "gpt-5.4"}, sql.NullString{}, + sql.NullString{}, + sql.NullBool{}, sql.NullInt64{}, sql.NullInt64{}, 1, 2, 3, 4, 5, 6, diff --git a/backend/internal/repository/usage_log_repo_stats.go b/backend/internal/repository/usage_log_repo_stats.go index aeccd49b30..53d60b7f4a 100644 --- a/backend/internal/repository/usage_log_repo_stats.go +++ b/backend/internal/repository/usage_log_repo_stats.go @@ -683,6 +683,9 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us args = append(args, int16(*filters.BillingType)) } conditions, args = appendUsageLogBillingModeWhereCondition(conditions, args, filters.BillingMode) + if filters.UpstreamModelMismatch != nil { + conditions = append(conditions, upstreamModelMismatchCondition("upstream_model_mismatch", *filters.UpstreamModelMismatch)) + } if filters.StartTime != nil { conditions = append(conditions, fmt.Sprintf("created_at >= $%d", len(args)+1)) args = append(args, *filters.StartTime) diff --git a/backend/internal/repository/usage_log_repo_stats_integration_test.go b/backend/internal/repository/usage_log_repo_stats_integration_test.go index 44052d07cc..e9025b1bd1 100644 --- a/backend/internal/repository/usage_log_repo_stats_integration_test.go +++ b/backend/internal/repository/usage_log_repo_stats_integration_test.go @@ -4,6 +4,7 @@ package repository import ( "context" + "strings" "testing" "time" @@ -12,6 +13,66 @@ import ( "github.com/stretchr/testify/require" ) +func TestUsageLog_UpstreamModelMismatchFilterAndPartialIndex(t *testing.T) { + ctx := context.Background() + tx := testEntTx(t) + client := tx.Client() + repo := newUsageLogRepositoryWithSQL(client, tx) + + user := mustCreateUser(t, client, &service.User{Email: "model-audit@test.com"}) + apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: user.ID, Key: "sk-model-audit", Name: "model-audit"}) + account := mustCreateAccount(t, client, &service.Account{Name: "model-audit-account"}) + now := time.Now().UTC() + responseModel := "gpt-5.4" + for _, mismatch := range []bool{true, false} { + mismatchValue := mismatch + _, err := repo.Create(ctx, &service.UsageLog{ + UserID: user.ID, APIKeyID: apiKey.ID, AccountID: account.ID, + Model: "gpt-5.5", InputTokens: 1, OutputTokens: 1, + UpstreamResponseModel: &responseModel, UpstreamModelMismatch: &mismatchValue, + CreatedAt: now, + }) + require.NoError(t, err) + } + + start := now.Add(-time.Hour) + end := now.Add(time.Hour) + trueValue := true + stats, err := repo.GetStatsWithFilters(ctx, usagestats.UsageLogFilters{ + UserID: user.ID, StartTime: &start, EndTime: &end, UpstreamModelMismatch: &trueValue, + }) + require.NoError(t, err) + require.Equal(t, int64(1), stats.TotalRequests) + + trend, err := repo.GetUsageTrendWithUsageFilters(ctx, start, end, "hour", usagestats.UsageLogFilters{ + UserID: user.ID, UpstreamModelMismatch: &trueValue, + }) + require.NoError(t, err) + require.Len(t, trend, 1) + require.Equal(t, int64(1), trend[0].Requests) + + _, err = tx.ExecContext(ctx, "SET LOCAL enable_seqscan = off") + require.NoError(t, err) + rows, err := tx.QueryContext(ctx, ` +EXPLAIN (COSTS OFF) +SELECT id +FROM usage_logs +WHERE upstream_model_mismatch IS TRUE +ORDER BY created_at DESC, id DESC +LIMIT 100 +`) + require.NoError(t, err) + defer func() { require.NoError(t, rows.Close()) }() + var planLines []string + for rows.Next() { + var line string + require.NoError(t, rows.Scan(&line)) + planLines = append(planLines, line) + } + require.NoError(t, rows.Err()) + require.Contains(t, strings.Join(planLines, "\n"), usageLogsUpstreamModelMismatchIndex) +} + func TestUsageLog_GetStatsWithFilters_AggregatesAndEndpoints(t *testing.T) { ctx := context.Background() tx := testEntTx(t) diff --git a/backend/internal/repository/usage_log_repo_trend.go b/backend/internal/repository/usage_log_repo_trend.go index 76863911c6..6d8bd79eaa 100644 --- a/backend/internal/repository/usage_log_repo_trend.go +++ b/backend/internal/repository/usage_log_repo_trend.go @@ -265,20 +265,20 @@ func (r *usageLogRepository) GetUserUsageTrendByUserID(ctx context.Context, user // GetUserModelStats 获取指定用户的模型统计 func (r *usageLogRepository) GetUserModelStats(ctx context.Context, userID int64, startTime, endTime time.Time) (results []ModelStat, err error) { - return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, 0, 0, 0, "", nil, nil, nil, usagestats.ModelSourceRequested, "") + return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, 0, 0, 0, "", nil, nil, nil, usagestats.ModelSourceRequested, "", nil) } // GetUsageTrendWithFilters returns usage trend data with optional filters func (r *usageLogRepository) GetUsageTrendWithFilters(ctx context.Context, startTime, endTime time.Time, granularity string, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8) (results []TrendDataPoint, err error) { - return r.getUsageTrendWithFilters(ctx, startTime, endTime, granularity, userID, apiKeyID, accountID, groupID, model, "", requestType, stream, billingType, "") + return r.getUsageTrendWithFilters(ctx, startTime, endTime, granularity, userID, apiKeyID, accountID, groupID, model, "", requestType, stream, billingType, "", nil) } func (r *usageLogRepository) GetUsageTrendWithUsageFilters(ctx context.Context, startTime, endTime time.Time, granularity string, filters UsageLogFilters) (results []TrendDataPoint, err error) { - return r.getUsageTrendWithFilters(ctx, startTime, endTime, granularity, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.ModelFilterSource, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode) + return r.getUsageTrendWithFilters(ctx, startTime, endTime, granularity, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.ModelFilterSource, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode, filters.UpstreamModelMismatch) } -func (r *usageLogRepository) getUsageTrendWithFilters(ctx context.Context, startTime, endTime time.Time, granularity string, userID, apiKeyID, accountID, groupID int64, model string, modelSource string, requestType *int16, stream *bool, billingType *int8, billingMode string) (results []TrendDataPoint, err error) { - if shouldUsePreaggregatedTrend(granularity, userID, apiKeyID, accountID, groupID, model, requestType, stream, billingType, billingMode) { +func (r *usageLogRepository) getUsageTrendWithFilters(ctx context.Context, startTime, endTime time.Time, granularity string, userID, apiKeyID, accountID, groupID int64, model string, modelSource string, requestType *int16, stream *bool, billingType *int8, billingMode string, upstreamModelMismatch *bool) (results []TrendDataPoint, err error) { + if shouldUsePreaggregatedTrend(granularity, userID, apiKeyID, accountID, groupID, model, requestType, stream, billingType, billingMode, upstreamModelMismatch) { aggregated, aggregatedErr := r.getUsageTrendFromAggregates(ctx, startTime, endTime, granularity) if aggregatedErr == nil && len(aggregated) > 0 { return aggregated, nil @@ -326,6 +326,9 @@ func (r *usageLogRepository) getUsageTrendWithFilters(ctx context.Context, start args = append(args, int16(*billingType)) } query, args = appendUsageLogBillingModeQueryFilter(query, args, billingMode, "") + if upstreamModelMismatch != nil { + query += " AND " + upstreamModelMismatchCondition("upstream_model_mismatch", *upstreamModelMismatch) + } query += " GROUP BY date ORDER BY date ASC" rows, err := r.sql.QueryContext(ctx, query, args...) @@ -348,7 +351,7 @@ func (r *usageLogRepository) getUsageTrendWithFilters(ctx context.Context, start return results, nil } -func shouldUsePreaggregatedTrend(granularity string, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8, billingMode string) bool { +func shouldUsePreaggregatedTrend(granularity string, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8, billingMode string, upstreamModelMismatch *bool) bool { if granularity != "day" && granularity != "hour" { return false } @@ -360,7 +363,8 @@ func shouldUsePreaggregatedTrend(granularity string, userID, apiKeyID, accountID requestType == nil && stream == nil && billingType == nil && - billingMode == "" + billingMode == "" && + upstreamModelMismatch == nil } func (r *usageLogRepository) getUsageTrendFromAggregates(ctx context.Context, startTime, endTime time.Time, granularity string) (results []TrendDataPoint, err error) { @@ -425,20 +429,20 @@ func (r *usageLogRepository) getUsageTrendFromAggregates(ctx context.Context, st // GetModelStatsWithFilters returns model statistics with optional filters func (r *usageLogRepository) GetModelStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8) (results []ModelStat, err error) { - return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, "", requestType, stream, billingType, usagestats.ModelSourceRequested, "") + return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, "", requestType, stream, billingType, usagestats.ModelSourceRequested, "", nil) } // GetModelStatsWithFiltersBySource returns model statistics with optional filters and model source dimension. // source: requested | upstream | mapping. func (r *usageLogRepository) GetModelStatsWithFiltersBySource(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8, source string) (results []ModelStat, err error) { - return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, "", requestType, stream, billingType, source, "") + return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, "", requestType, stream, billingType, source, "", nil) } func (r *usageLogRepository) GetModelStatsWithUsageFiltersBySource(ctx context.Context, startTime, endTime time.Time, filters UsageLogFilters, source string) (results []ModelStat, err error) { - return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType, source, filters.BillingMode) + return r.getModelStatsWithFiltersBySource(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType, source, filters.BillingMode, filters.UpstreamModelMismatch) } -func (r *usageLogRepository) getModelStatsWithFiltersBySource(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8, source string, billingMode string) (results []ModelStat, err error) { +func (r *usageLogRepository) getModelStatsWithFiltersBySource(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8, source string, billingMode string, upstreamModelMismatch *bool) (results []ModelStat, err error) { actualCostExpr := "COALESCE(SUM(actual_cost), 0) as actual_cost" // 当仅按 account_id 聚合时,实际费用使用账号倍率(total_cost * account_rate_multiplier)。 if accountID > 0 && userID == 0 && apiKeyID == 0 { @@ -490,6 +494,9 @@ func (r *usageLogRepository) getModelStatsWithFiltersBySource(ctx context.Contex args = append(args, int16(*billingType)) } query, args = appendUsageLogBillingModeQueryFilter(query, args, billingMode, "") + if upstreamModelMismatch != nil { + query += " AND " + upstreamModelMismatchCondition("upstream_model_mismatch", *upstreamModelMismatch) + } query += fmt.Sprintf(" GROUP BY %s ORDER BY total_tokens DESC", modelExpr) rows, err := r.sql.QueryContext(ctx, query, args...) @@ -514,14 +521,14 @@ func (r *usageLogRepository) getModelStatsWithFiltersBySource(ctx context.Contex // GetGroupStatsWithFilters returns group usage statistics with optional filters func (r *usageLogRepository) GetGroupStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8) (results []usagestats.GroupStat, err error) { - return r.getGroupStatsWithFilters(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, "", requestType, stream, billingType, "") + return r.getGroupStatsWithFilters(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, "", requestType, stream, billingType, "", nil) } func (r *usageLogRepository) GetGroupStatsWithUsageFilters(ctx context.Context, startTime, endTime time.Time, filters UsageLogFilters) (results []usagestats.GroupStat, err error) { - return r.getGroupStatsWithFilters(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode) + return r.getGroupStatsWithFilters(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType, filters.BillingMode, filters.UpstreamModelMismatch) } -func (r *usageLogRepository) getGroupStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8, billingMode string) (results []usagestats.GroupStat, err error) { +func (r *usageLogRepository) getGroupStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, model string, requestType *int16, stream *bool, billingType *int8, billingMode string, upstreamModelMismatch *bool) (results []usagestats.GroupStat, err error) { query := ` SELECT COALESCE(ul.group_id, 0) as group_id, @@ -564,6 +571,9 @@ func (r *usageLogRepository) getGroupStatsWithFilters(ctx context.Context, start args = append(args, int16(*billingType)) } query, args = appendUsageLogBillingModeQueryFilter(query, args, billingMode, "ul") + if upstreamModelMismatch != nil { + query += " AND " + upstreamModelMismatchCondition("ul.upstream_model_mismatch", *upstreamModelMismatch) + } query += " GROUP BY ul.group_id, g.name ORDER BY total_tokens DESC" rows, err := r.sql.QueryContext(ctx, query, args...) diff --git a/backend/internal/repository/usage_log_session_id_unit_test.go b/backend/internal/repository/usage_log_session_id_unit_test.go index 7804dce306..64c1c1bab2 100644 --- a/backend/internal/repository/usage_log_session_id_unit_test.go +++ b/backend/internal/repository/usage_log_session_id_unit_test.go @@ -32,7 +32,7 @@ func newSessionIDUsageLog(sessionID *string) *service.UsageLog { // arg slice / arg-type table so the five INSERT column lists stay in sync. session_id // is the penultimate arg (created_at is always last). func TestPrepareUsageLogInsert_SessionIDArgWiring(t *testing.T) { - require.Len(t, usageLogInsertArgTypes, 57, "arg-type table must include session_id") + require.Len(t, usageLogInsertArgTypes, 59, "arg-type table must include session_id") sessionID := "sess-persisted-123" prepared := prepareUsageLogInsert(newSessionIDUsageLog(&sessionID)) diff --git a/backend/internal/repository/user_subscription_repo.go b/backend/internal/repository/user_subscription_repo.go index 8047d7d78b..0af3eda985 100644 --- a/backend/internal/repository/user_subscription_repo.go +++ b/backend/internal/repository/user_subscription_repo.go @@ -368,7 +368,7 @@ func (r *userSubscriptionRepository) UpdateNotes(ctx context.Context, subscripti return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) } -func (r *userSubscriptionRepository) ActivateWindows(ctx context.Context, id int64, start time.Time) error { +func (r *userSubscriptionRepository) ActivateWindows(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error { client := clientFromContext(ctx, r.client) n, err := client.UserSubscription.Update(). Where( @@ -377,24 +377,24 @@ func (r *userSubscriptionRepository) ActivateWindows(ctx context.Context, id int usersubscription.WeeklyWindowStartIsNil(), usersubscription.MonthlyWindowStartIsNil(), ). - SetDailyWindowStart(start). - SetWeeklyWindowStart(start). - SetMonthlyWindowStart(start). + SetDailyWindowStart(dailyStart). + SetWeeklyWindowStart(periodicStart). + SetMonthlyWindowStart(periodicStart). Save(ctx) return r.translateConditionalWindowReset(ctx, client, id, n, err) } -func (r *userSubscriptionRepository) ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, newWindowStart time.Time) error { +func (r *userSubscriptionRepository) ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, dailyStart, periodicStart time.Time) error { client := clientFromContext(ctx, r.client) update := client.UserSubscription.UpdateOneID(id) if resetDaily { - update.SetDailyUsageUsd(0).SetDailyWindowStart(newWindowStart) + update.SetDailyUsageUsd(0).SetDailyWindowStart(dailyStart) } if resetWeekly { - update.SetWeeklyUsageUsd(0).SetWeeklyWindowStart(newWindowStart) + update.SetWeeklyUsageUsd(0).SetWeeklyWindowStart(periodicStart) } if resetMonthly { - update.SetMonthlyUsageUsd(0).SetMonthlyWindowStart(newWindowStart) + update.SetMonthlyUsageUsd(0).SetMonthlyWindowStart(periodicStart) } _, err := update.Save(ctx) return translatePersistenceError(err, service.ErrSubscriptionNotFound, nil) diff --git a/backend/internal/repository/user_subscription_repo_integration_test.go b/backend/internal/repository/user_subscription_repo_integration_test.go index 1d1964a20f..28f4c1b7d3 100644 --- a/backend/internal/repository/user_subscription_repo_integration_test.go +++ b/backend/internal/repository/user_subscription_repo_integration_test.go @@ -451,8 +451,9 @@ func (s *UserSubscriptionRepoSuite) TestActivateWindows() { group := s.mustCreateGroup("g-activate") sub := s.mustCreateSubscription(user.ID, group.ID, nil) + dailyStart := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) activateAt := time.Date(2025, 1, 1, 12, 0, 0, 0, time.UTC) - err := s.repo.ActivateWindows(s.ctx, sub.ID, activateAt) + err := s.repo.ActivateWindows(s.ctx, sub.ID, dailyStart, activateAt) s.Require().NoError(err, "ActivateWindows") got, err := s.repo.GetByID(s.ctx, sub.ID) @@ -460,7 +461,9 @@ func (s *UserSubscriptionRepoSuite) TestActivateWindows() { s.Require().NotNil(got.DailyWindowStart) s.Require().NotNil(got.WeeklyWindowStart) s.Require().NotNil(got.MonthlyWindowStart) - s.Require().WithinDuration(activateAt, *got.DailyWindowStart, time.Microsecond) + s.Require().WithinDuration(dailyStart, *got.DailyWindowStart, time.Microsecond) + s.Require().WithinDuration(activateAt, *got.WeeklyWindowStart, time.Microsecond) + s.Require().WithinDuration(activateAt, *got.MonthlyWindowStart, time.Microsecond) } func (s *UserSubscriptionRepoSuite) TestActivateWindows_StaleActivationPreservesExistingWindows() { @@ -469,15 +472,16 @@ func (s *UserSubscriptionRepoSuite) TestActivateWindows_StaleActivationPreserves sub := s.mustCreateSubscription(user.ID, group.ID, nil) activatedAt := time.Date(2025, 1, 1, 12, 0, 0, 0, time.UTC) manualResetAt := activatedAt.Add(2 * time.Hour) + manualDailyStart := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) - s.Require().NoError(s.repo.ActivateWindows(s.ctx, sub.ID, activatedAt)) - s.Require().NoError(s.repo.ResetUsageWindows(s.ctx, sub.ID, true, true, true, manualResetAt)) + s.Require().NoError(s.repo.ActivateWindows(s.ctx, sub.ID, activatedAt, activatedAt)) + s.Require().NoError(s.repo.ResetUsageWindows(s.ctx, sub.ID, true, true, true, manualDailyStart, manualResetAt)) // Simulate a concurrent request carrying the original unactivated snapshot. - s.Require().NoError(s.repo.ActivateWindows(s.ctx, sub.ID, activatedAt.Add(time.Hour))) + s.Require().NoError(s.repo.ActivateWindows(s.ctx, sub.ID, activatedAt.Add(time.Hour), activatedAt.Add(time.Hour))) got, err := s.repo.GetByID(s.ctx, sub.ID) s.Require().NoError(err) - s.Require().WithinDuration(manualResetAt, *got.DailyWindowStart, time.Microsecond) + s.Require().WithinDuration(manualDailyStart, *got.DailyWindowStart, time.Microsecond) s.Require().WithinDuration(manualResetAt, *got.WeeklyWindowStart, time.Microsecond) s.Require().WithinDuration(manualResetAt, *got.MonthlyWindowStart, time.Microsecond) } @@ -535,7 +539,7 @@ func (s *UserSubscriptionRepoSuite) TestResetUsageWindows_ClearsUsageAfterAutoma newWindowStart := oldWindowStart.Add(24 * time.Hour) s.Require().NoError(s.repo.ResetDailyUsage(s.ctx, sub.ID, &oldWindowStart, newWindowStart)) s.Require().NoError(s.repo.IncrementUsage(s.ctx, sub.ID, 3)) - s.Require().NoError(s.repo.ResetUsageWindows(s.ctx, sub.ID, true, false, false, newWindowStart)) + s.Require().NoError(s.repo.ResetUsageWindows(s.ctx, sub.ID, true, false, false, newWindowStart, newWindowStart)) got, err := s.repo.GetByID(s.ctx, sub.ID) s.Require().NoError(err) @@ -770,7 +774,7 @@ func (s *UserSubscriptionRepoSuite) TestActiveExpiredBoundaries_UsageAndReset_Ba s.Require().Equal(active.ID, got.ID, "expected active subscription") activateAt := time.Now().Add(-25 * time.Hour) - s.Require().NoError(s.repo.ActivateWindows(s.ctx, active.ID, activateAt), "ActivateWindows") + s.Require().NoError(s.repo.ActivateWindows(s.ctx, active.ID, activateAt, activateAt), "ActivateWindows") s.Require().NoError(s.repo.IncrementUsage(s.ctx, active.ID, 1.25), "IncrementUsage") after, err := s.repo.GetByID(s.ctx, active.ID) diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 04a549098f..196b081985 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -371,6 +371,10 @@ func TestAPIContracts(t *testing.T) { "video_price_720p": null, "video_price_1080p": null, "web_search_price_per_call": null, + "search_price_per_1k": null, + "audio_tts_price_per_million_chars": null, + "audio_stt_price_per_hour": null, + "audio_realtime_price_per_min": null, "allow_image_generation": false, "allow_batch_image_generation": false, "batch_image_discount_multiplier": 0, @@ -880,6 +884,9 @@ func TestAPIContracts(t *testing.T) { "invitation_code_enabled": false, "home_content": "", "hide_ccs_import_button": false, + "grok_default_text_model": "grok-4.5", + "grok_default_base_url_mode": "cli", + "grok_cross_client_model_map_enabled": true, "purchase_subscription_enabled": false, "purchase_subscription_url": "", "table_default_page_size": 20, @@ -971,6 +978,7 @@ func TestAPIContracts(t *testing.T) { "payment_alipay_mobile_precreate_deep_link": false, "balance_low_notify_enabled": false, "account_quota_notify_enabled": false, + "account_scheduling_thresholds": {"anthropic":100,"grok":100,"openai":100}, "subscription_expiry_notify_enabled": true, "balance_low_notify_threshold": 0, "balance_low_notify_recharge_url": "", @@ -1157,6 +1165,9 @@ func TestAPIContracts(t *testing.T) { "doc_url": "", "home_content": "", "hide_ccs_import_button": false, + "grok_default_text_model": "grok-4.5", + "grok_default_base_url_mode": "cli", + "grok_cross_client_model_map_enabled": true, "purchase_subscription_enabled": false, "purchase_subscription_url": "", "table_default_page_size": 20, @@ -1276,6 +1287,7 @@ func TestAPIContracts(t *testing.T) { "payment_alipay_mobile_precreate_deep_link": false, "balance_low_notify_enabled": false, "account_quota_notify_enabled": false, + "account_scheduling_thresholds": {"anthropic":100,"grok":100,"openai":100}, "subscription_expiry_notify_enabled": true, "balance_low_notify_threshold": 0, "balance_low_notify_recharge_url": "", @@ -2251,10 +2263,10 @@ func (stubUserSubscriptionRepo) UpdateStatus(ctx context.Context, subscriptionID func (stubUserSubscriptionRepo) UpdateNotes(ctx context.Context, subscriptionID int64, notes string) error { return errors.New("not implemented") } -func (stubUserSubscriptionRepo) ActivateWindows(ctx context.Context, id int64, start time.Time) error { +func (stubUserSubscriptionRepo) ActivateWindows(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error { return errors.New("not implemented") } -func (stubUserSubscriptionRepo) ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, newWindowStart time.Time) error { +func (stubUserSubscriptionRepo) ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, dailyStart, periodicStart time.Time) error { return errors.New("not implemented") } func (stubUserSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error { diff --git a/backend/internal/server/middleware/api_key_auth_google_test.go b/backend/internal/server/middleware/api_key_auth_google_test.go index 25adbb3121..6b5f10d48e 100644 --- a/backend/internal/server/middleware/api_key_auth_google_test.go +++ b/backend/internal/server/middleware/api_key_auth_google_test.go @@ -82,7 +82,7 @@ type fakeGoogleSubscriptionRepo struct { getByID func(ctx context.Context, id int64) (*service.UserSubscription, error) getActive func(ctx context.Context, userID, groupID int64) (*service.UserSubscription, error) updateStatus func(ctx context.Context, subscriptionID int64, status string) error - activateWindow func(ctx context.Context, id int64, start time.Time) error + activateWindow func(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error resetDaily func(ctx context.Context, id int64, start time.Time) error resetWeekly func(ctx context.Context, id int64, start time.Time) error resetMonthly func(ctx context.Context, id int64, start time.Time) error @@ -231,13 +231,13 @@ func (f fakeGoogleSubscriptionRepo) UpdateStatus(ctx context.Context, subscripti func (f fakeGoogleSubscriptionRepo) UpdateNotes(ctx context.Context, subscriptionID int64, notes string) error { return errors.New("not implemented") } -func (f fakeGoogleSubscriptionRepo) ActivateWindows(ctx context.Context, id int64, start time.Time) error { +func (f fakeGoogleSubscriptionRepo) ActivateWindows(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error { if f.activateWindow != nil { - return f.activateWindow(ctx, id, start) + return f.activateWindow(ctx, id, dailyStart, periodicStart) } return errors.New("not implemented") } -func (f fakeGoogleSubscriptionRepo) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { +func (f fakeGoogleSubscriptionRepo) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time, time.Time) error { return errors.New("not implemented") } func (f fakeGoogleSubscriptionRepo) ResetDailyUsage(ctx context.Context, id int64, _ *time.Time, start time.Time) error { @@ -886,7 +886,7 @@ func TestApiKeyAuthWithSubscriptionGoogle_SubscriptionLimitExceededReturns429(t return &clone, nil }, updateStatus: func(ctx context.Context, subscriptionID int64, status string) error { return nil }, - activateWindow: func(ctx context.Context, id int64, start time.Time) error { return nil }, + activateWindow: func(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error { return nil }, resetDaily: func(ctx context.Context, id int64, start time.Time) error { return nil }, resetWeekly: func(ctx context.Context, id int64, start time.Time) error { return nil }, resetMonthly: func(ctx context.Context, id int64, start time.Time) error { return nil }, diff --git a/backend/internal/server/middleware/api_key_auth_test.go b/backend/internal/server/middleware/api_key_auth_test.go index 1bfb3519b1..964fe8041b 100644 --- a/backend/internal/server/middleware/api_key_auth_test.go +++ b/backend/internal/server/middleware/api_key_auth_test.go @@ -119,7 +119,7 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { return &clone, nil }, updateStatus: func(ctx context.Context, subscriptionID int64, status string) error { return nil }, - activateWindow: func(ctx context.Context, id int64, start time.Time) error { return nil }, + activateWindow: func(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error { return nil }, resetDaily: func(ctx context.Context, id int64, start time.Time) error { sub.DailyWindowStart = &start sub.DailyUsageUSD = 0 @@ -252,7 +252,7 @@ func TestSimpleModeBypassesQuotaCheck(t *testing.T) { return &clone, nil }, updateStatus: func(ctx context.Context, subscriptionID int64, status string) error { return nil }, - activateWindow: func(ctx context.Context, id int64, start time.Time) error { return nil }, + activateWindow: func(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error { return nil }, resetDaily: func(ctx context.Context, id int64, start time.Time) error { return nil }, resetWeekly: func(ctx context.Context, id int64, start time.Time) error { return nil }, resetMonthly: func(ctx context.Context, id int64, start time.Time) error { return nil }, @@ -1631,7 +1631,7 @@ type stubUserSubscriptionRepo struct { getByID func(ctx context.Context, id int64) (*service.UserSubscription, error) getActive func(ctx context.Context, userID, groupID int64) (*service.UserSubscription, error) updateStatus func(ctx context.Context, subscriptionID int64, status string) error - activateWindow func(ctx context.Context, id int64, start time.Time) error + activateWindow func(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error resetDaily func(ctx context.Context, id int64, start time.Time) error resetWeekly func(ctx context.Context, id int64, start time.Time) error resetMonthly func(ctx context.Context, id int64, start time.Time) error @@ -1753,14 +1753,14 @@ func (r *stubUserSubscriptionRepo) UpdateNotes(ctx context.Context, subscription return errors.New("not implemented") } -func (r *stubUserSubscriptionRepo) ActivateWindows(ctx context.Context, id int64, start time.Time) error { +func (r *stubUserSubscriptionRepo) ActivateWindows(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error { if r.activateWindow != nil { - return r.activateWindow(ctx, id, start) + return r.activateWindow(ctx, id, dailyStart, periodicStart) } return errors.New("not implemented") } -func (r *stubUserSubscriptionRepo) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { +func (r *stubUserSubscriptionRepo) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time, time.Time) error { return errors.New("not implemented") } diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index 4afd4c3032..149d4f95bc 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -381,6 +381,7 @@ func registerAccountRoutes(admin *gin.RouterGroup, h *handler.Handlers, stepUpAu accounts.POST("/:id/revert-proxy-fallback", h.Admin.Account.RevertProxyFallback) accounts.GET("/:id/usage", h.Admin.Account.GetUsage) accounts.GET("/:id/today-stats", h.Admin.Account.GetTodayStats) + accounts.POST("/usage/batch", h.Admin.Account.GetBatchUsage) accounts.POST("/today-stats/batch", h.Admin.Account.GetBatchTodayStats) accounts.POST("/:id/clear-rate-limit", h.Admin.Account.ClearRateLimit) accounts.POST("/:id/reset-quota", h.Admin.Account.ResetQuota) @@ -465,9 +466,12 @@ func registerAntigravityOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) func registerGrokOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) { grok := admin.Group("/grok") { + grok.GET("/oauth/capabilities", h.Admin.GrokOAuth.GetCapabilities) grok.POST("/oauth/auth-url", h.Admin.GrokOAuth.GenerateAuthURL) grok.POST("/oauth/exchange-code", h.Admin.GrokOAuth.ExchangeCode) grok.POST("/oauth/refresh-token", h.Admin.GrokOAuth.RefreshToken) + grok.POST("/oauth/sso-token", h.Admin.GrokOAuth.ValidateSSOToken) + grok.POST("/oauth/password", h.Admin.GrokOAuth.AuthorizePassword) grok.POST("/oauth/create-from-oauth", h.Admin.GrokOAuth.CreateAccountFromOAuth) grok.POST("/sso-to-oauth", h.Admin.GrokOAuth.CreateAccountsFromSSO) grok.POST("/oauth/reconcile", h.Admin.GrokOAuth.ReconcileOAuthAccounts) diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index 1ce1fd9788..8d5fcd6a53 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -256,11 +256,65 @@ func RegisterGatewayRoutes( gateway.POST("/images/batches/:id/cancel", h.BatchImage.Cancel) gateway.DELETE("/images/batches/:id", h.BatchImage.DeleteRecord) gateway.DELETE("/images/batches/:id/outputs", h.BatchImage.DeleteOutputs) + // OpenAI-compatible clients may create through /videos; xAI receives the + // canonical /videos/generations route inside the Grok media forwarder. + gateway.POST("/videos", videoGenerationHandler) gateway.POST("/videos/generations", videoGenerationHandler) gateway.POST("/videos/edits", videoEditHandler) gateway.POST("/videos/extensions", videoExtensionHandler) + gateway.GET("/videos/generations/:request_id/content", videoContentHandler) + gateway.GET("/videos/edits/:request_id/content", videoContentHandler) + gateway.GET("/videos/extensions/:request_id/content", videoContentHandler) + gateway.GET("/videos/generations/:request_id", videoStatusHandler) + gateway.GET("/videos/edits/:request_id", videoStatusHandler) + gateway.GET("/videos/extensions/:request_id", videoStatusHandler) gateway.GET("/videos/:request_id", videoStatusHandler) gateway.GET("/videos/:request_id/content", videoContentHandler) + + // xAI Voice APIs (Grok platform only): HTTP TTS/STT + Realtime WS. + // Not part of the creation-center product surface — gateway relay only. + voiceHandler := func(endpoint string) gin.HandlerFunc { + return func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Voice API is not supported for this platform"}}) + return + } + h.OpenAIGateway.GrokVoice(c, endpoint) + } + } + gateway.POST("/tts", voiceHandler("tts")) + gateway.POST("/stt", voiceHandler("stt")) + gateway.POST("/custom-voices", voiceHandler("custom-voices")) + customVoicePathHandler := func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Voice API is not supported for this platform"}}) + return + } + h.OpenAIGateway.GrokVoice(c, grokCustomVoiceEndpoint(c)) + } + gateway.GET("/custom-voices", voiceHandler("custom-voices")) + gateway.GET("/custom-voices/:voice_id/audio", customVoicePathHandler) + gateway.GET("/custom-voices/:voice_id", customVoicePathHandler) + gateway.PATCH("/custom-voices/:voice_id", customVoicePathHandler) + gateway.DELETE("/custom-voices/:voice_id", customVoicePathHandler) + gateway.GET("/realtime", func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Realtime API is not supported for this platform"}}) + return + } + h.OpenAIGateway.GrokRealtime(c) + }) + gateway.POST("/web_search", func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Web Search API is not supported for this platform"}}) + return + } + h.Gateway.WebSearch(c) + }) } // Gemini 原生 API 兼容层(Gemini SDK/CLI 直连) @@ -334,12 +388,62 @@ func RegisterGatewayRoutes( r.POST("/images/generations/async", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, h.AsyncImage.Submit) r.POST("/images/edits/async", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, h.AsyncImage.Submit) r.GET("/images/tasks/:task_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, h.AsyncImage.Get) + r.POST("/videos", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoGenerationHandler) r.POST("/videos/generations", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoGenerationHandler) r.POST("/videos/edits", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoEditHandler) r.POST("/videos/extensions", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoExtensionHandler) + r.GET("/videos/generations/:request_id/content", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoContentHandler) + r.GET("/videos/edits/:request_id/content", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoContentHandler) + r.GET("/videos/extensions/:request_id/content", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoContentHandler) + r.GET("/videos/generations/:request_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoStatusHandler) + r.GET("/videos/edits/:request_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoStatusHandler) + r.GET("/videos/extensions/:request_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoStatusHandler) r.GET("/videos/:request_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoStatusHandler) r.GET("/videos/:request_id/content", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, videoContentHandler) + rootVoiceHandler := func(endpoint string) gin.HandlerFunc { + return func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Voice API is not supported for this platform"}}) + return + } + h.OpenAIGateway.GrokVoice(c, endpoint) + } + } + r.POST("/tts", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootVoiceHandler("tts")) + r.POST("/stt", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootVoiceHandler("stt")) + r.POST("/custom-voices", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootVoiceHandler("custom-voices")) + rootCustomVoicePathHandler := func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Voice API is not supported for this platform"}}) + return + } + h.OpenAIGateway.GrokVoice(c, grokCustomVoiceEndpoint(c)) + } + r.GET("/custom-voices", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootVoiceHandler("custom-voices")) + r.GET("/custom-voices/:voice_id/audio", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootCustomVoicePathHandler) + r.GET("/custom-voices/:voice_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootCustomVoicePathHandler) + r.PATCH("/custom-voices/:voice_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootCustomVoicePathHandler) + r.DELETE("/custom-voices/:voice_id", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, rootCustomVoicePathHandler) + r.GET("/realtime", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Realtime API is not supported for this platform"}}) + return + } + h.OpenAIGateway.GrokRealtime(c) + }) + r.POST("/web_search", bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic, func(c *gin.Context) { + if getGroupPlatform(c) != service.PlatformGrok { + service.MarkOpsClientBusinessLimited(c, service.OpsClientBusinessLimitedReasonLocalFeatureGate) + c.JSON(http.StatusNotFound, gin.H{"error": gin.H{"type": "not_found_error", "message": "Web Search API is not supported for this platform"}}) + return + } + h.Gateway.WebSearch(c) + }) + // Antigravity 模型列表 r.GET("/antigravity/models", gin.HandlerFunc(apiKeyAuth), requireGroupAnthropic, h.Gateway.AntigravityModels) @@ -503,6 +607,22 @@ func compositeGeminiTargetPlatformMiddleware(resolver *service.CompositeRouteRes } } +// grokCustomVoiceEndpoint derives the upstream Voice endpoint for the +// /custom-voices/:voice_id[/audio] routes. +// +// The /audio suffix must be decided from the matched route template, not from +// the raw URL path: a voice literally named "audio" makes GET +// /custom-voices/audio match /custom-voices/:voice_id, and a raw-path suffix +// check would rewrite it to custom-voices/audio/audio — turning a profile +// lookup into an audio download. +func grokCustomVoiceEndpoint(c *gin.Context) string { + endpoint := "custom-voices/" + c.Param("voice_id") + if strings.HasSuffix(c.FullPath(), "/:voice_id/audio") { + endpoint += "/audio" + } + return endpoint +} + func compositeGeminiModelFromParams(c *gin.Context) string { if c == nil { return "" diff --git a/backend/internal/server/routes/gateway_test.go b/backend/internal/server/routes/gateway_test.go index ada931e3a5..ded7bd2989 100644 --- a/backend/internal/server/routes/gateway_test.go +++ b/backend/internal/server/routes/gateway_test.go @@ -152,6 +152,8 @@ func TestGatewayRoutesGrokImagesAndVideosPathsAreRegistered(t *testing.T) { "/images/generations", "/images/edits", "/v1/videos/generations", + "/v1/videos", + "/videos", "/videos/generations", "/v1/videos/edits", "/videos/edits", @@ -170,8 +172,20 @@ func TestGatewayRoutesGrokImagesAndVideosPathsAreRegistered(t *testing.T) { for _, path := range []string{ "/v1/videos/request-123", "/videos/request-123", + "/v1/videos/generations/request-123", + "/videos/generations/request-123", + "/v1/videos/edits/request-123", + "/videos/edits/request-123", + "/v1/videos/extensions/request-123", + "/videos/extensions/request-123", "/v1/videos/request-123/content", "/videos/request-123/content", + "/v1/videos/generations/request-123/content", + "/videos/generations/request-123/content", + "/v1/videos/edits/request-123/content", + "/videos/edits/request-123/content", + "/v1/videos/extensions/request-123/content", + "/videos/extensions/request-123/content", } { req := httptest.NewRequest(http.MethodGet, path, nil) w := httptest.NewRecorder() @@ -182,6 +196,62 @@ func TestGatewayRoutesGrokImagesAndVideosPathsAreRegistered(t *testing.T) { } } +func TestGatewayRoutesGrokCustomVoiceCRUDPathsAreRegistered(t *testing.T) { + router := newGatewayRoutesTestRouter(service.PlatformGrok) + registered := make(map[string]bool) + for _, route := range router.Routes() { + registered[route.Method+" "+route.Path] = true + } + for _, route := range []string{ + "POST /v1/custom-voices", + "GET /v1/custom-voices", + "GET /v1/custom-voices/:voice_id", + "PATCH /v1/custom-voices/:voice_id", + "DELETE /v1/custom-voices/:voice_id", + "GET /v1/custom-voices/:voice_id/audio", + "POST /custom-voices", + "GET /custom-voices", + "GET /custom-voices/:voice_id", + "PATCH /custom-voices/:voice_id", + "DELETE /custom-voices/:voice_id", + "GET /custom-voices/:voice_id/audio", + } { + require.True(t, registered[route], "%s should be registered", route) + } +} + +func TestGrokCustomVoiceEndpointUsesRouteTemplateNotRawPath(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + var got string + capture := func(c *gin.Context) { + got = grokCustomVoiceEndpoint(c) + c.Status(http.StatusOK) + } + router.GET("/v1/custom-voices/:voice_id/audio", capture) + router.GET("/v1/custom-voices/:voice_id", capture) + + for _, tc := range []struct { + path string + want string + }{ + {path: "/v1/custom-voices/voice-123", want: "custom-voices/voice-123"}, + {path: "/v1/custom-voices/voice-123/audio", want: "custom-voices/voice-123/audio"}, + // A voice literally named "audio" matches /:voice_id, not /:voice_id/audio. + // A raw-path suffix check would turn this profile lookup into an audio download. + {path: "/v1/custom-voices/audio", want: "custom-voices/audio"}, + {path: "/v1/custom-voices/audio/audio", want: "custom-voices/audio/audio"}, + } { + got = "" + req := httptest.NewRequest(http.MethodGet, tc.path, nil) + w := httptest.NewRecorder() + + router.ServeHTTP(w, req) + require.Equal(t, http.StatusOK, w.Code, "path=%s", tc.path) + require.Equal(t, tc.want, got, "path=%s", tc.path) + } +} + func TestGatewayRoutesCompositeVideoLookupsUseGrokHandler(t *testing.T) { router := newGatewayRoutesTestRouter(service.PlatformComposite) @@ -241,6 +311,8 @@ func TestGatewayRoutesNonGrokVideosAreRejectedAtPlatformGate(t *testing.T) { body string }{ {http.MethodPost, "/v1/videos/generations", `{"model":"grok-imagine-video-1.5","prompt":"waves"}`}, + {http.MethodPost, "/v1/videos", `{"model":"grok-imagine-video-1.5","prompt":"waves"}`}, + {http.MethodPost, "/videos", `{"model":"grok-imagine-video-1.5","prompt":"waves"}`}, {http.MethodPost, "/videos/generations", `{"model":"grok-imagine-video-1.5","prompt":"waves"}`}, {http.MethodPost, "/v1/videos/edits", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`}, {http.MethodPost, "/videos/edits", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`}, @@ -248,8 +320,20 @@ func TestGatewayRoutesNonGrokVideosAreRejectedAtPlatformGate(t *testing.T) { {http.MethodPost, "/videos/extensions", `{"model":"grok-imagine-video","prompt":"waves","video":{"url":"https://example.com/in.mp4"}}`}, {http.MethodGet, "/v1/videos/request-123", ""}, {http.MethodGet, "/videos/request-123", ""}, + {http.MethodGet, "/v1/videos/generations/request-123", ""}, + {http.MethodGet, "/videos/generations/request-123", ""}, + {http.MethodGet, "/v1/videos/edits/request-123", ""}, + {http.MethodGet, "/videos/edits/request-123", ""}, + {http.MethodGet, "/v1/videos/extensions/request-123", ""}, + {http.MethodGet, "/videos/extensions/request-123", ""}, {http.MethodGet, "/v1/videos/request-123/content", ""}, {http.MethodGet, "/videos/request-123/content", ""}, + {http.MethodGet, "/v1/videos/generations/request-123/content", ""}, + {http.MethodGet, "/videos/generations/request-123/content", ""}, + {http.MethodGet, "/v1/videos/edits/request-123/content", ""}, + {http.MethodGet, "/videos/edits/request-123/content", ""}, + {http.MethodGet, "/v1/videos/extensions/request-123/content", ""}, + {http.MethodGet, "/videos/extensions/request-123/content", ""}, } { req := httptest.NewRequest(tc.method, tc.path, strings.NewReader(tc.body)) req.Header.Set("Content-Type", "application/json") diff --git a/backend/internal/server/routes/prompt_audit_route_coverage_test.go b/backend/internal/server/routes/prompt_audit_route_coverage_test.go index e4591ab4f5..f03cf5b314 100644 --- a/backend/internal/server/routes/prompt_audit_route_coverage_test.go +++ b/backend/internal/server/routes/prompt_audit_route_coverage_test.go @@ -41,14 +41,19 @@ func TestEveryGatewayPOSTRouteIsClassifiedForPromptAuditCoverage(t *testing.T) { "/images/generations/async": {"image_task_handler.go"}, "/images/edits/async": {"image_task_handler.go"}, "/images/batches": {"batch_image_handler.go"}, + "/videos": {"grok_media.go"}, "/videos/generations": {"grok_media.go"}, "/videos/edits": {"grok_media.go"}, "/videos/extensions": {"grok_media.go"}, "/models/*modelAction": {"gemini_v1beta_handler.go"}, + "/tts": {"grok_audio.go"}, + "/web_search": {"gateway_web_search.go"}, } excluded := map[string]string{ "/messages/count_tokens": "tokenization only; it does not execute a model request", "/images/batches/:id/cancel": "control-plane cancellation with no user prompt", + "/stt": "speech transcription is not a text-generation prompt", + "/custom-voices": "voice profile management has no model prompt", } unclassified := make([]string, 0) diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index ab54d3c805..cb74aaf50f 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -6,6 +6,7 @@ import ( "errors" "hash/fnv" "log/slog" + "net/url" "reflect" "sort" "strconv" @@ -71,6 +72,7 @@ type Account struct { modelMappingCacheRawPtr uintptr modelMappingCacheRawLen int modelMappingCacheRawSig uint64 + modelMappingCacheRuntimeVersion uint64 // header_overrides 热路径缓存(非持久化字段,同 model_mapping 缓存先例) headerOverrideCache map[string]string @@ -553,6 +555,7 @@ func stringMappingFromRaw(raw any) map[string]string { } func (a *Account) GetModelMapping() map[string]string { + runtimeVersion := xai.RuntimeModelMappingVersion() credentialsPtr := mapPtr(a.Credentials) rawMapping, _ := a.Credentials["model_mapping"].(map[string]any) rawPtr := mapPtr(rawMapping) @@ -563,7 +566,8 @@ func (a *Account) GetModelMapping() map[string]string { if a.modelMappingCacheReady && a.modelMappingCacheCredentialsPtr == credentialsPtr && a.modelMappingCacheRawPtr == rawPtr && - a.modelMappingCacheRawLen == rawLen { + a.modelMappingCacheRawLen == rawLen && + a.modelMappingCacheRuntimeVersion == runtimeVersion { rawSig = modelMappingSignature(rawMapping) rawSigReady = true if a.modelMappingCacheRawSig == rawSig { @@ -582,6 +586,7 @@ func (a *Account) GetModelMapping() map[string]string { a.modelMappingCacheRawPtr = rawPtr a.modelMappingCacheRawLen = rawLen a.modelMappingCacheRawSig = rawSig + a.modelMappingCacheRuntimeVersion = runtimeVersion return mapping } @@ -620,6 +625,11 @@ func (a *Account) resolveModelMapping(rawMapping map[string]any) map[string]stri "gemini-3-flash", "gemini-3.1-pro-high", "gemini-3.1-pro-low", + "gemini-3.6-flash", + "gemini-3.6-flash-high", + "gemini-3.6-flash-low", + "gemini-3.6-flash-medium", + "gemini-3.6-flash-tiered", }) applyAntigravityGemini31ProAliases(result) } @@ -1306,25 +1316,47 @@ func (a *Account) GetOpenAIRefreshToken() string { // traffic (OAuth authorization and token refresh) always uses the official // auth endpoints regardless of this value. func (a *Account) GetGrokBaseURL() string { - if !a.IsGrok() { + if a == nil || !a.IsGrok() { return "" } - baseURL := strings.TrimSpace(a.GetCredential("base_url")) if a.IsGrokOAuth() { - // Operators switch subscription traffic between the official CLI - // gateway, the official/regional API hosts and third-party relays - // (individual endpoints go down from time to time), so a stored - // value is always honored as-is. Only empty or unparseable values - // fall back to the default CLI gateway. - if baseURL == "" || !xai.IsParseableBaseURL(baseURL) { - return xai.DefaultCLIBaseURL + return a.GetGrokBaseURLOr(xai.DefaultCLIBaseURL) + } + return a.GetGrokBaseURLOr(xai.DefaultBaseURL) +} + +// GetGrokBaseURLOr resolves an explicit account endpoint, falling back to the +// supplied default. Official OAuth endpoints are normalized here; custom +// endpoints are retained for the request builder's operator URL policy. +func (a *Account) GetGrokBaseURLOr(defaultBaseURL string) string { + if a == nil || !a.IsGrok() { + return "" + } + defaultBaseURL = strings.TrimRight(strings.TrimSpace(defaultBaseURL), "/") + if defaultBaseURL == "" { + if a.IsGrokOAuth() { + defaultBaseURL = xai.DefaultCLIBaseURL + } else { + defaultBaseURL = xai.DefaultBaseURL } + } + baseURL := strings.TrimSpace(a.GetCredential("base_url")) + if baseURL == "" { + return defaultBaseURL + } + if !a.IsGrokOAuth() { return baseURL } - if baseURL != "" { - return baseURL + // Explicit regional/API or custom values remain pinned. Custom endpoints are checked by the + // operator URL policy at the request builder, which has access to config. + if validated, err := xai.ValidateTrustedBaseURL(baseURL); err == nil { + return validated } - return xai.DefaultBaseURL + if parsed, err := url.Parse(baseURL); err == nil && parsed.Scheme != "" && parsed.Host != "" && + parsed.User == nil && parsed.RawQuery == "" && parsed.Fragment == "" { + return strings.TrimRight(baseURL, "/") + } + return defaultBaseURL } // GetGrokMediaBaseURL selects the upstream used by Grok Imagine APIs. diff --git a/backend/internal/service/account_credentials_redact.go b/backend/internal/service/account_credentials_redact.go index 8eb70513ba..1e32483817 100644 --- a/backend/internal/service/account_credentials_redact.go +++ b/backend/internal/service/account_credentials_redact.go @@ -7,6 +7,8 @@ var SensitiveCredentialKeys = []string{ "access_token", "refresh_token", "id_token", "agent_private_key", // API Key 类 "api_key", "session_key", "cookie", + // Grok Web SSO / password (must never persist or echo after Build OAuth) + "password", "sso_token", "sso", "sso-rw", "clearTextPassword", // 云服务凭据 "aws_secret_access_key", "aws_session_token", "service_account_json", "service_account", "private_key", diff --git a/backend/internal/service/account_grok_media_eligibility.go b/backend/internal/service/account_grok_media_eligibility.go new file mode 100644 index 0000000000..5f05a92772 --- /dev/null +++ b/backend/internal/service/account_grok_media_eligibility.go @@ -0,0 +1,65 @@ +package service + +import infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + +// ValidateGrokMediaEligibilityExtra validates the optional per-account media +// routing override. A nil value removes the override and restores automatic +// provider-observation routing. +func ValidateGrokMediaEligibilityExtra(platform string, extra map[string]any) error { + if platform != PlatformGrok || extra == nil { + return nil + } + raw, exists := extra[GrokMediaEligibleExtraKey] + if !exists || raw == nil { + return nil + } + if _, ok := raw.(bool); !ok { + return infraerrors.BadRequest("GROK_MEDIA_ELIGIBILITY_INVALID", "grok_media_eligible must be a boolean or null") + } + return nil +} + +func normalizeGrokMediaEligibilityExtra(platform string, extra map[string]any) (map[string]any, error) { + if platform != PlatformGrok { + return extra, nil + } + if err := ValidateGrokMediaEligibilityExtra(platform, extra); err != nil { + return nil, err + } + if extra == nil { + return nil, nil + } + normalized := shallowCopyMap(extra) + if normalized[GrokMediaEligibleExtraKey] == nil { + delete(normalized, GrokMediaEligibleExtraKey) + } + return normalized, nil +} + +func normalizeGrokMediaEligibilityUpdateExtra(account *Account, input *UpdateAccountInput, normalized map[string]any) (map[string]any, error) { + if account == nil || account.Platform != PlatformGrok { + return normalized, nil + } + if input == nil { + return nil, infraerrors.BadRequest("INVALID_ACCOUNT_INPUT", "account update input is required") + } + if err := ValidateGrokMediaEligibilityExtra(account.Platform, input.Extra); err != nil { + return nil, err + } + if normalized == nil { + normalized = make(map[string]any) + } else { + normalized = shallowCopyMap(normalized) + } + raw, provided := input.Extra[GrokMediaEligibleExtraKey] + if provided { + if raw == nil { + delete(normalized, GrokMediaEligibleExtraKey) + } + return normalized, nil + } + if current, ok := account.Extra[GrokMediaEligibleExtraKey].(bool); ok { + normalized[GrokMediaEligibleExtraKey] = current + } + return normalized, nil +} diff --git a/backend/internal/service/account_grok_media_eligibility_test.go b/backend/internal/service/account_grok_media_eligibility_test.go index a79f9c02ec..47d1f2b94f 100644 --- a/backend/internal/service/account_grok_media_eligibility_test.go +++ b/backend/internal/service/account_grok_media_eligibility_test.go @@ -25,20 +25,6 @@ func TestGrokMediaGenerationEligibility(t *testing.T) { StatusCode: http.StatusOK, WeeklyStatusCode: http.StatusOK, } - freeBilling := &xai.BillingSummary{ - PeriodType: "monthly", - StatusCode: http.StatusOK, - WeeklyStatusCode: http.StatusOK, - MonthlyStatusCode: http.StatusOK, - MonthlyUpdatedAt: "2026-07-17T00:00:00Z", - } - inconclusiveBilling := &xai.BillingSummary{ - StatusCode: http.StatusOK, - WeeklyStatusCode: http.StatusOK, - MonthlyStatusCode: http.StatusBadGateway, - Partial: true, - FailedWindows: []string{"monthly"}, - } weeklyForbidden := &xai.BillingSummary{ StatusCode: http.StatusOK, WeeklyStatusCode: http.StatusForbidden, @@ -61,8 +47,6 @@ func TestGrokMediaGenerationEligibility(t *testing.T) { {name: "non oauth grok account stays eligible", account: &Account{Platform: PlatformGrok, Type: AccountTypeAPIKey}, want: true, wantReason: "non_oauth"}, {name: "unobserved oauth fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth}, want: false, wantReason: "billing_unobserved"}, {name: "weekly paid usage is eligible without inferring from period type", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyAllowance}}, want: true, wantReason: "eligible"}, - {name: "observed free account is rejected", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: freeBilling}}, want: false, wantReason: "billing_free_tier"}, - {name: "inconclusive billing fails closed", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: inconclusiveBilling}}, want: false, wantReason: "billing_inconclusive"}, {name: "billing forbidden is rejected", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: forbiddenBilling}}, want: false, wantReason: "billing_forbidden"}, {name: "weekly billing forbidden is rejected after partial success", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: weeklyForbidden}}, want: false, wantReason: "billing_forbidden"}, {name: "monthly billing forbidden is rejected after partial success", account: &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{grokBillingExtraKey: monthlyForbidden}}, want: false, wantReason: "billing_forbidden"}, diff --git a/backend/internal/service/account_scheduling_threshold_eval.go b/backend/internal/service/account_scheduling_threshold_eval.go new file mode 100644 index 0000000000..07ed0db57d --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_eval.go @@ -0,0 +1,468 @@ +package service + +import ( + "encoding/json" + "math" + "strconv" + "strings" + "time" +) + +// AccountSchedulingThresholdDecision captures the pure pause decision for one account. +type AccountSchedulingThresholdDecision struct { + ShouldPause bool + Platform string + Window string + Scope string + ThresholdPercent int + UsedPercent float64 + Until *time.Time +} + +type accountSchedulingThresholdCandidate struct { + window string + scope string + usedPercent float64 + until *time.Time +} + +const accountSchedulingThresholdCredentialKey = "account_scheduling_threshold" + +// EvaluateAccountSchedulingThreshold evaluates whether an account should be paused +// based on the current per-platform scheduling threshold snapshot. +func EvaluateAccountSchedulingThreshold(account *Account, thresholds map[string]int, now time.Time) AccountSchedulingThresholdDecision { + decision := AccountSchedulingThresholdDecision{} + if account == nil { + return decision + } + + decision.Platform = strings.ToLower(strings.TrimSpace(account.Platform)) + if decision.Platform == "" { + return decision + } + if !isAllowedSchedulingThresholdPlatform(decision.Platform) { + return decision + } + + threshold, ok := resolveEffectiveAccountSchedulingThreshold(account, thresholds, decision.Platform) + decision.ThresholdPercent = threshold + if !ok || threshold >= 100 { + return decision + } + + var winner *accountSchedulingThresholdCandidate + switch decision.Platform { + case PlatformOpenAI: + winner = pickLatestResetSchedulingCandidate(openAIThresholdCandidates(account), threshold, now) + case PlatformAnthropic: + winner = pickLatestResetSchedulingCandidate(anthropicThresholdCandidates(account), threshold, now) + case PlatformGrok: + winner = pickLatestResetSchedulingCandidate(grokThresholdCandidates(account), threshold, now) + default: + return decision + } + + if winner == nil { + return decision + } + + decision.ShouldPause = true + decision.Window = winner.window + decision.Scope = winner.scope + decision.UsedPercent = winner.usedPercent + decision.Until = winner.until + return decision +} + +func isAllowedSchedulingThresholdPlatform(platform string) bool { + for _, allowed := range AllowedSchedulingThresholdPlatforms { + if platform == allowed { + return true + } + } + return false +} + +func resolveEffectiveAccountSchedulingThreshold(account *Account, thresholds map[string]int, platform string) (int, bool) { + if account != nil { + if threshold, ok := accountSchedulingThresholdOverride(account); ok { + return threshold, true + } + } + return lookupAccountSchedulingThreshold(thresholds, platform) +} + +func accountSchedulingThresholdOverride(account *Account) (int, bool) { + if account == nil || len(account.Credentials) == 0 { + return 0, false + } + raw, ok := account.Credentials[accountSchedulingThresholdCredentialKey] + if !ok { + return 0, false + } + return parseAccountSchedulingThresholdValue(raw) +} + +func parseAccountSchedulingThresholdValue(raw any) (int, bool) { + var value int + switch v := raw.(type) { + case int: + value = v + case int64: + value = int(v) + case float64: + value = int(math.Round(v)) + case float32: + value = int(math.Round(float64(v))) + case json.Number: + parsed, err := v.Float64() + if err != nil { + return 0, false + } + value = int(math.Round(parsed)) + case string: + raw := strings.TrimSpace(v) + parsed, err := strconv.Atoi(raw) + if err == nil { + value = parsed + break + } + parsedFloat, floatErr := strconv.ParseFloat(raw, 64) + if floatErr != nil { + return 0, false + } + value = int(math.Round(parsedFloat)) + default: + return 0, false + } + if value < 1 || value > 100 { + return 0, false + } + return value, true +} + +func lookupAccountSchedulingThreshold(thresholds map[string]int, platform string) (int, bool) { + if len(thresholds) == 0 { + return 0, false + } + value, ok := thresholds[platform] + return value, ok +} + +func openAIThresholdCandidates(account *Account) []*accountSchedulingThresholdCandidate { + if account == nil { + return nil + } + if !openAICodexSnapshotIdentityTrusted(account) { + return nil + } + return []*accountSchedulingThresholdCandidate{ + openAIThresholdCandidate(account.Extra, "5h"), + openAIThresholdCandidate(account.Extra, "7d"), + } +} + +func openAICodexSnapshotIdentityTrusted(account *Account) bool { + if account == nil || !account.IsOpenAIOAuth() || len(account.Extra) == 0 { + return true + } + + if identityValuesConflict( + firstStringValue(account.Credentials, "email"), + firstStringValue(account.Extra, "email", "email_address"), + ) { + return false + } + if identityValuesConflict( + firstStringValue(account.Credentials, "chatgpt_account_id"), + firstStringValue(account.Extra, "chatgpt_account_id", "account_id"), + ) { + return false + } + if identityValuesConflict( + firstStringValue(account.Credentials, "workspace_id", "chatgpt_workspace_id", "organization_id", "org_id"), + firstStringValue(account.Extra, "workspace_id", "chatgpt_workspace_id", "organization_id", "org_id"), + ) { + return false + } + return true +} + +func identityValuesConflict(left, right string) bool { + left = strings.TrimSpace(left) + right = strings.TrimSpace(right) + return left != "" && right != "" && !strings.EqualFold(left, right) +} + +// firstStringValue returns the first non-empty string among the given map keys. +// Used by OpenAI codex snapshot identity matching for scheduling thresholds. +func firstStringValue(values map[string]any, keys ...string) string { + if len(values) == 0 { + return "" + } + for _, key := range keys { + raw, ok := values[key] + if !ok || raw == nil { + continue + } + switch typed := raw.(type) { + case string: + if v := strings.TrimSpace(typed); v != "" { + return v + } + default: + if v := strings.TrimSpace(stringValue(raw)); v != "" { + return v + } + } + } + return "" +} + +func openAIThresholdCandidate(extra map[string]any, window string) *accountSchedulingThresholdCandidate { + if len(extra) == 0 { + return nil + } + + var ( + usedPercentKey string + resetAtKey string + ) + switch window { + case "5h": + usedPercentKey = "codex_5h_used_percent" + resetAtKey = "codex_5h_reset_at" + case "7d": + usedPercentKey = "codex_7d_used_percent" + resetAtKey = "codex_7d_reset_at" + default: + return nil + } + + usedPercent, ok := extra[usedPercentKey] + if !ok { + return nil + } + return &accountSchedulingThresholdCandidate{ + window: window, + usedPercent: utilizationAsPercent(usedPercent), + until: parseSchedulingResetAt(extra[resetAtKey]), + } +} + +func anthropicThresholdCandidates(account *Account) []*accountSchedulingThresholdCandidate { + if account == nil { + return nil + } + + var candidates []*accountSchedulingThresholdCandidate + if usedPercent := utilizationAsPercent(account.Extra["session_window_utilization"]); usedPercent > 0 { + candidates = append(candidates, &accountSchedulingThresholdCandidate{ + window: "5h", + usedPercent: usedPercent, + until: cloneTimePtr(account.SessionWindowEnd), + }) + } + if usedPercent := utilizationAsPercent(account.Extra["passive_usage_7d_utilization"]); usedPercent > 0 { + candidates = append(candidates, &accountSchedulingThresholdCandidate{ + window: "7d", + usedPercent: usedPercent, + until: parseSchedulingResetAt(account.Extra["passive_usage_7d_reset"]), + }) + } + return candidates +} + +// NOTE: Gemini / Kiro / Antigravity are intentionally NOT threshold-pausing +// platforms (see AllowedSchedulingThresholdPlatforms and the evaluator switch, +// asserted by TestEvaluateAccountSchedulingThreshold_UnsupportedPlatformsDoNotPause). +// Their former per-platform candidate readers were dead code — never reachable +// from EvaluateAccountSchedulingThreshold — and have been removed to avoid the +// false impression that configuring a threshold for them has any effect. The +// kiro_sched_* / antigravity_sched_* extras are still written purely as +// observability snapshots. + +// grokThresholdCandidates uses only header-projected +// grok_sched_utilization / grok_sched_reset_at (rolling quota window, reset +// capped at ~25h when written). Official billing 7d/30d windows are not used +// for auto-pause here. +func grokThresholdCandidates(account *Account) []*accountSchedulingThresholdCandidate { + if account == nil { + return nil + } + return []*accountSchedulingThresholdCandidate{ + { + window: "quota", + scope: "grok", + usedPercent: schedulingPercentValue(account.Extra["grok_sched_utilization"]), + until: parseSchedulingResetAt(account.Extra["grok_sched_reset_at"]), + }, + } +} + +func pickLatestResetSchedulingCandidate(candidates []*accountSchedulingThresholdCandidate, threshold int, now time.Time) *accountSchedulingThresholdCandidate { + var winner *accountSchedulingThresholdCandidate + for _, candidate := range candidates { + if !candidateMatchesThreshold(candidate, threshold, now) { + continue + } + if winner == nil || candidate.until.After(*winner.until) { + winner = candidate + continue + } + if winner.until.Equal(*candidate.until) && candidate.usedPercent > winner.usedPercent { + winner = candidate + } + } + return winner +} + +func candidateMatchesThreshold(candidate *accountSchedulingThresholdCandidate, threshold int, now time.Time) bool { + if candidate == nil || candidate.until == nil || !candidate.until.After(now) { + return false + } + return candidate.usedPercent >= float64(threshold) +} + +func utilizationAsPercent(raw any) float64 { + switch v := raw.(type) { + case float64: + if v >= 0 && v <= 1 { + return v * 100 + } + return v + case float32: + value := float64(v) + if value >= 0 && value <= 1 { + return value * 100 + } + return value + case int: + return float64(v) + case int64: + return float64(v) + case json.Number: + value, err := v.Float64() + if err != nil { + return 0 + } + if strings.Contains(v.String(), ".") && value >= 0 && value <= 1 { + return value * 100 + } + return value + case string: + trimmed := strings.TrimSpace(v) + value, err := strconv.ParseFloat(trimmed, 64) + if err != nil { + return 0 + } + if strings.Contains(trimmed, ".") && value >= 0 && value <= 1 { + return value * 100 + } + return value + default: + return 0 + } +} + +func schedulingPercentValue(raw any) float64 { + switch v := raw.(type) { + case float64: + return v + case float32: + return float64(v) + case int: + return float64(v) + case int64: + return float64(v) + case json.Number: + value, err := v.Float64() + if err != nil { + return 0 + } + return value + case string: + value, err := strconv.ParseFloat(strings.TrimSpace(v), 64) + if err != nil { + return 0 + } + return value + default: + return 0 + } +} + +func parseSchedulingResetAt(raw any) *time.Time { + switch v := raw.(type) { + case nil: + return nil + case time.Time: + ts := v + return &ts + case *time.Time: + return cloneTimePtr(v) + case string: + trimmed := strings.TrimSpace(v) + if trimmed == "" { + return nil + } + ts, err := parseSchedulingTime(trimmed) + if err != nil { + return nil + } + return &ts + case json.Number: + if value, err := v.Int64(); err == nil && value > 0 { + ts := time.Unix(value, 0) + return &ts + } + if value, err := v.Float64(); err == nil && value > 0 { + ts := time.Unix(int64(value), 0) + return &ts + } + case float64: + if v > 0 { + ts := time.Unix(int64(v), 0) + return &ts + } + case float32: + if v > 0 { + ts := time.Unix(int64(v), 0) + return &ts + } + case int: + if v > 0 { + ts := time.Unix(int64(v), 0) + return &ts + } + case int64: + if v > 0 { + ts := time.Unix(v, 0) + return &ts + } + } + return nil +} + +func parseSchedulingTime(raw string) (time.Time, error) { + formats := []string{ + time.RFC3339, + time.RFC3339Nano, + "2006-01-02T15:04:05Z", + "2006-01-02T15:04:05.000Z", + } + for _, format := range formats { + if ts, err := time.Parse(format, raw); err == nil { + return ts, nil + } + } + return time.Time{}, strconv.ErrSyntax +} + +func cloneTimePtr(src *time.Time) *time.Time { + if src == nil { + return nil + } + value := *src + return &value +} diff --git a/backend/internal/service/account_scheduling_threshold_eval_test.go b/backend/internal/service/account_scheduling_threshold_eval_test.go new file mode 100644 index 0000000000..22a1efdf0c --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_eval_test.go @@ -0,0 +1,339 @@ +//go:build unit + +package service + +import ( + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +func TestEvaluateAccountSchedulingThreshold_OpenAIChoosesLatestResetWindow(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + wantUntil := now.Add(72 * time.Hour) + account := &Account{ + Platform: PlatformOpenAI, + Extra: map[string]any{ + "codex_5h_used_percent": 90.0, + "codex_5h_reset_at": now.Add(2 * time.Hour).Format(time.RFC3339), + "codex_7d_used_percent": 85.0, + "codex_7d_reset_at": wantUntil.Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 80, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, PlatformOpenAI, decision.Platform) + require.Equal(t, "7d", decision.Window) + require.Empty(t, decision.Scope) + require.Equal(t, 85.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} + +func TestEvaluateAccountSchedulingThreshold_OpenAIIgnoresMismatchedCodexSnapshotIdentity(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 13, 8, 50, 0, 0, time.UTC) + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "email": "CageLeen9208@outlook.com", + "chatgpt_account_id": "1f945aa7-d9a9-4369-9542-0c702ff4adb0", + "workspace_id": "org-nU4goUxMmureroyswT5oYPv4", + "chatgpt_workspace_id": "org-nU4goUxMmureroyswT5oYPv4", + }, + Extra: map[string]any{ + "email": "MasonDobies01@outlook.com", + "name": "Paul Clark", + "workspace_id": "org-avRk1G4qdXg7qph3cRIraNKf", + "codex_7d_used_percent": 100.0, + "codex_7d_reset_at": now.Add(7 * 24 * time.Hour).Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 99, + }, now) + + require.False(t, decision.ShouldPause) +} + +func TestEvaluateAccountSchedulingThreshold_AnthropicIgnoresExpiredFiveHourWindow(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + expiredEnd := now.Add(-30 * time.Minute) + wantUntil := now.Add(5 * 24 * time.Hour) + account := &Account{ + Platform: PlatformAnthropic, + SessionWindowEnd: &expiredEnd, + Extra: map[string]any{ + "session_window_utilization": 0.99, + "passive_usage_7d_utilization": 0.82, + "passive_usage_7d_reset": float64(wantUntil.Unix()), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformAnthropic: 80, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, PlatformAnthropic, decision.Platform) + require.Equal(t, "7d", decision.Window) + require.Empty(t, decision.Scope) + require.Equal(t, 82.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} + +func TestEvaluateAccountSchedulingThreshold_FractionalPlatformsKeepFractionSemantics(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + openAIUntil := now.Add(24 * time.Hour) + openAIAccount := &Account{ + Platform: PlatformOpenAI, + Extra: map[string]any{ + "codex_5h_used_percent": 0.91, + "codex_5h_reset_at": openAIUntil.Format(time.RFC3339), + }, + } + + openAIDecision := EvaluateAccountSchedulingThreshold(openAIAccount, map[string]int{ + PlatformOpenAI: 90, + }, now) + + require.True(t, openAIDecision.ShouldPause) + require.Equal(t, 91.0, openAIDecision.UsedPercent) + + anthropicUntil := now.Add(5 * time.Hour) + anthropicAccount := &Account{ + Platform: PlatformAnthropic, + SessionWindowEnd: &anthropicUntil, + Extra: map[string]any{ + "session_window_utilization": 0.92, + }, + } + + anthropicDecision := EvaluateAccountSchedulingThreshold(anthropicAccount, map[string]int{ + PlatformAnthropic: 90, + }, now) + + require.True(t, anthropicDecision.ShouldPause) + require.Equal(t, 92.0, anthropicDecision.UsedPercent) +} + +func TestEvaluateAccountSchedulingThreshold_AccountOverrideCanLowerOpenAIThreshold(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + wantUntil := now.Add(12 * time.Hour) + account := &Account{ + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "account_scheduling_threshold": 80, + }, + Extra: map[string]any{ + "codex_7d_used_percent": 85.0, + "codex_7d_reset_at": wantUntil.Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 90, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, PlatformOpenAI, decision.Platform) + require.Equal(t, 80, decision.ThresholdPercent) + require.Equal(t, "7d", decision.Window) + require.Empty(t, decision.Scope) + require.Equal(t, 85.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} + +func TestEvaluateAccountSchedulingThreshold_AccountOverrideHundredDisablesOpenAI(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + account := &Account{ + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "account_scheduling_threshold": 100, + }, + Extra: map[string]any{ + "codex_7d_used_percent": 99.0, + "codex_7d_reset_at": now.Add(24 * time.Hour).Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 80, + }, now) + + require.False(t, decision.ShouldPause) + require.Equal(t, 100, decision.ThresholdPercent) +} + +func TestEvaluateAccountSchedulingThreshold_AccountOverrideRoundsDecimalThreshold(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + wantUntil := now.Add(12 * time.Hour) + account := &Account{ + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "account_scheduling_threshold": 75.5, + }, + Extra: map[string]any{ + "codex_7d_used_percent": 80.0, + "codex_7d_reset_at": wantUntil.Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformOpenAI: 90, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, 76, decision.ThresholdPercent) + require.Equal(t, 80.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} + +func TestEvaluateAccountSchedulingThreshold_UnsupportedPlatformsDoNotPause(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + cases := []struct { + name string + platform string + threshold int + extra map[string]any + }{ + { + name: "gemini", + platform: PlatformGemini, + threshold: 80, + extra: map[string]any{ + "gemini_usage_raw": map[string]any{ + "buckets": []any{ + map[string]any{ + "modelId": "gemini-2.5-pro", + "remainingFraction": 0.05, + "resetTime": now.Add(2 * time.Hour).Format(time.RFC3339), + }, + }, + }, + }, + }, + { + name: "kiro", + platform: PlatformKiro, + threshold: 90, + extra: map[string]any{ + "kiro_sched_utilization": 99.0, + "kiro_sched_reset_at": now.Add(24 * time.Hour).Format(time.RFC3339), + }, + }, + { + name: "antigravity", + platform: PlatformAntigravity, + threshold: 90, + extra: map[string]any{ + "antigravity_sched_utilization": 92.0, + "antigravity_sched_reset_at": now.Add(48 * time.Hour).Format(time.RFC3339), + "antigravity_sched_scope": "gemini", + }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + account := &Account{ + Platform: tc.platform, + Credentials: map[string]any{ + "account_scheduling_threshold": 1, + }, + Extra: tc.extra, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + tc.platform: tc.threshold, + }, now) + + require.False(t, decision.ShouldPause) + require.Zero(t, decision.ThresholdPercent) + }) + } +} + +func TestEvaluateAccountSchedulingThreshold_GrokUsesConfiguredThresholds(t *testing.T) { + t.Parallel() + + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + wantUntil := now.Add(2 * time.Hour) + account := &Account{ + Platform: PlatformGrok, + Extra: map[string]any{ + "grok_sched_utilization": 92.0, + "grok_sched_reset_at": wantUntil.Format(time.RFC3339), + }, + } + + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{ + PlatformGrok: 90, + }, now) + + require.True(t, decision.ShouldPause) + require.Equal(t, PlatformGrok, decision.Platform) + require.Equal(t, 90, decision.ThresholdPercent) + require.Equal(t, "grok", decision.Scope) + require.Equal(t, 92.0, decision.UsedPercent) + require.NotNil(t, decision.Until) + require.True(t, wantUntil.Equal(*decision.Until)) +} + +func TestEvaluateAccountSchedulingThreshold_GrokUsesOnlyHeaderQuotaWindow(t *testing.T) { + t.Parallel() + // Billing seven_day/thirty_day must not drive pause; only grok_sched_* may. + now := time.Date(2026, 6, 3, 12, 0, 0, 0, time.UTC) + weeklyEnd := now.Add(3 * time.Hour) + weeklyPct := 99.0 + headerUntil := now.Add(2 * time.Hour) + account := &Account{ + Platform: PlatformGrok, + Extra: map[string]any{ + "grok_sched_utilization": 50.0, // below threshold + "grok_sched_reset_at": headerUntil.Format(time.RFC3339), + grokBillingExtraKey: &xai.BillingSummary{ + UsagePercent: &weeklyPct, + PeriodEnd: weeklyEnd.Format(time.RFC3339), + }, + }, + } + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{PlatformGrok: 90}, now) + require.False(t, decision.ShouldPause, "high billing % alone must not pause under scheduling windows") + + account.Extra["grok_sched_utilization"] = 95.0 + decision = EvaluateAccountSchedulingThreshold(account, map[string]int{PlatformGrok: 90}, now) + require.True(t, decision.ShouldPause) + require.Equal(t, "grok", decision.Scope) + require.Equal(t, "quota", decision.Window) + require.NotNil(t, decision.Until) + require.True(t, headerUntil.Equal(*decision.Until)) +} diff --git a/backend/internal/service/account_scheduling_threshold_integration_test.go b/backend/internal/service/account_scheduling_threshold_integration_test.go new file mode 100644 index 0000000000..d2b80dbbcf --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_integration_test.go @@ -0,0 +1,136 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +type thresholdSelectionAccountRepoStub struct { + rateLimitAccountRepoStub + accounts []Account +} + +func (r *thresholdSelectionAccountRepoStub) ListSchedulableByPlatform(_ context.Context, platform string) ([]Account, error) { + filtered := make([]Account, 0, len(r.accounts)) + for _, account := range r.accounts { + if account.Platform == platform { + filtered = append(filtered, account) + } + } + return filtered, nil +} + +func (r *thresholdSelectionAccountRepoStub) ListSchedulableByGroupIDAndPlatform(ctx context.Context, _ int64, platform string) ([]Account, error) { + return r.ListSchedulableByPlatform(ctx, platform) +} + +func (r *thresholdSelectionAccountRepoStub) ListSchedulableUngroupedByPlatform(ctx context.Context, platform string) ([]Account, error) { + return r.ListSchedulableByPlatform(ctx, platform) +} + +func TestGatewayService_ListSchedulableAccounts_DoesNotFilterUnsupportedThresholdPlatforms(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":90}` + + accountRepo := &thresholdSelectionAccountRepoStub{ + accounts: []Account{ + { + ID: 3101, + Platform: PlatformKiro, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{ + "account_scheduling_threshold": 1, + }, + Extra: map[string]any{ + "kiro_sched_utilization": 95.0, + "kiro_sched_reset_at": time.Now().UTC().Add(2 * time.Hour).Format(time.RFC3339), + }, + }, + { + ID: 3102, + Platform: PlatformKiro, + Status: StatusActive, + Schedulable: true, + Extra: map[string]any{ + "kiro_sched_utilization": 42.0, + "kiro_sched_reset_at": time.Now().UTC().Add(2 * time.Hour).Format(time.RFC3339), + }, + }, + }, + } + + rateLimitService := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rateLimitService.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + svc := &GatewayService{ + accountRepo: accountRepo, + cfg: &config.Config{}, + rateLimitService: rateLimitService, + } + + accounts, useMixed, err := svc.listSchedulableAccounts(context.Background(), nil, PlatformKiro, false) + + require.NoError(t, err) + require.False(t, useMixed) + require.Len(t, accounts, 2) + require.Equal(t, int64(3101), accounts[0].ID) + require.Equal(t, int64(3102), accounts[1].ID) + require.Equal(t, 0, accountRepo.tempCalls) +} + +func TestOpenAIGatewayService_ListSchedulableAccounts_FiltersThresholdBlockedAccounts(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":85}` + + accountRepo := &thresholdSelectionAccountRepoStub{ + accounts: []Account{ + { + ID: 4101, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + Extra: map[string]any{ + "codex_7d_used_percent": 91.0, + "codex_7d_reset_at": time.Now().UTC().Add(12 * time.Hour).Format(time.RFC3339), + }, + }, + { + ID: 4102, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + Extra: map[string]any{ + "codex_7d_used_percent": 40.0, + "codex_7d_reset_at": time.Now().UTC().Add(12 * time.Hour).Format(time.RFC3339), + }, + }, + }, + } + + rateLimitService := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rateLimitService.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + svc := &OpenAIGatewayService{ + accountRepo: accountRepo, + cfg: &config.Config{}, + rateLimitService: rateLimitService, + } + + accounts, err := svc.listSchedulableAccounts(context.Background(), nil, PlatformOpenAI) + + require.NoError(t, err) + require.Len(t, accounts, 1) + require.Equal(t, int64(4102), accounts[0].ID) + require.Equal(t, 1, accountRepo.tempCalls) +} diff --git a/backend/internal/service/account_scheduling_threshold_reason.go b/backend/internal/service/account_scheduling_threshold_reason.go new file mode 100644 index 0000000000..b43c7e22a1 --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_reason.go @@ -0,0 +1,188 @@ +package service + +import ( + "encoding/json" + "fmt" + "strings" + "time" +) + +const AccountSchedulingThresholdReasonSource = "account_scheduling_threshold" + +const ( + defaultTempUnschedReasonErrorMessage = "temporary scheduling block reason unavailable" + defaultAccountSchedulingThresholdErrorMessage = "account scheduling threshold reached" +) + +type tempUnschedReasonPayload struct { + Source string `json:"source,omitempty"` + Platform string `json:"platform,omitempty"` + Window string `json:"window,omitempty"` + Scope string `json:"scope,omitempty"` + ThresholdPercent int `json:"threshold_percent,omitempty"` + UsedPercent float64 `json:"used_percent,omitempty"` + UntilUnix int64 `json:"until_unix,omitempty"` + TriggeredAtUnix int64 `json:"triggered_at_unix,omitempty"` + ErrorMessage string `json:"error_message"` +} + +type AccountSchedulingThresholdReasonInput struct { + Platform string + Window string + Scope string + ThresholdPercent int + UsedPercent float64 + Until time.Time + Now time.Time +} + +func BuildTempUnschedReasonPayload(source string, errorMessage string) string { + payload := tempUnschedReasonPayload{ + Source: strings.TrimSpace(source), + ErrorMessage: normalizeTempUnschedReasonErrorMessage(errorMessage, defaultTempUnschedReasonErrorMessage), + } + + raw, err := json.Marshal(payload) + if err != nil { + return payload.ErrorMessage + } + return string(raw) +} + +func BuildAccountSchedulingThresholdReason(errorMessage string) string { + return BuildTempUnschedReasonPayload( + AccountSchedulingThresholdReasonSource, + normalizeTempUnschedReasonErrorMessage(errorMessage, defaultAccountSchedulingThresholdErrorMessage), + ) +} + +func BuildDetailedAccountSchedulingThresholdReason(input AccountSchedulingThresholdReasonInput) string { + triggeredAt := input.Now + if triggeredAt.IsZero() { + triggeredAt = time.Now().UTC() + } + payload := tempUnschedReasonPayload{ + Source: AccountSchedulingThresholdReasonSource, + Platform: strings.TrimSpace(input.Platform), + Window: strings.TrimSpace(input.Window), + Scope: strings.TrimSpace(input.Scope), + ThresholdPercent: input.ThresholdPercent, + UsedPercent: input.UsedPercent, + TriggeredAtUnix: triggeredAt.Unix(), + ErrorMessage: buildAccountSchedulingThresholdErrorMessage(input), + } + if !input.Until.IsZero() { + payload.UntilUnix = input.Until.UTC().Unix() + } + + raw, err := json.Marshal(payload) + if err != nil { + return payload.ErrorMessage + } + return string(raw) +} + +func IsAccountSchedulingThresholdReason(rawReason string) bool { + payload, ok := parseTempUnschedReasonPayload(rawReason) + if !ok { + return false + } + return payload.Source == AccountSchedulingThresholdReasonSource +} + +func parseTempUnschedReasonPayload(rawReason string) (tempUnschedReasonPayload, bool) { + rawReason = strings.TrimSpace(rawReason) + if rawReason == "" { + return tempUnschedReasonPayload{}, false + } + + var payload tempUnschedReasonPayload + if err := json.Unmarshal([]byte(rawReason), &payload); err != nil { + return tempUnschedReasonPayload{}, false + } + payload.Source = strings.TrimSpace(payload.Source) + payload.ErrorMessage = strings.TrimSpace(payload.ErrorMessage) + return payload, true +} + +func normalizeTempUnschedReasonErrorMessage(errorMessage string, fallback string) string { + errorMessage = strings.TrimSpace(errorMessage) + if errorMessage != "" { + return errorMessage + } + + fallback = strings.TrimSpace(fallback) + if fallback != "" { + return fallback + } + return defaultTempUnschedReasonErrorMessage +} + +func buildAccountSchedulingThresholdErrorMessage(input AccountSchedulingThresholdReasonInput) string { + platform := strings.TrimSpace(input.Platform) + if platform == "" { + platform = "account" + } + + target := strings.TrimSpace(input.Window) + if scope := strings.TrimSpace(input.Scope); scope != "" { + if target == "" { + target = scope + } else { + target = target + "/" + scope + } + } + if target == "" { + target = "usage window" + } + + threshold := input.ThresholdPercent + if threshold <= 0 { + threshold = 100 + } + + untilText := "the window reset" + if !input.Until.IsZero() { + untilText = input.Until.UTC().Format(time.RFC3339) + } + + return fmt.Sprintf( + "%s scheduling threshold reached for %s: %.1f%% used >= %d%%; paused until %s", + platform, + target, + input.UsedPercent, + threshold, + untilText, + ) +} + +func tempUnschedStateFromStoredReason(rawReason string, fallbackUntilUnix int64) *TempUnschedState { + state := &TempUnschedState{ + UntilUnix: fallbackUntilUnix, + RuleIndex: -1, + } + + rawReason = strings.TrimSpace(rawReason) + if rawReason == "" { + state.ErrorMessage = defaultTempUnschedReasonErrorMessage + return state + } + + parsed := TempUnschedState{RuleIndex: -1} + if err := json.Unmarshal([]byte(rawReason), &parsed); err == nil { + if fallbackUntilUnix > parsed.UntilUnix { + parsed.UntilUnix = fallbackUntilUnix + } + if strings.TrimSpace(parsed.ErrorMessage) == "" { + if IsAccountSchedulingThresholdReason(rawReason) { + parsed.ErrorMessage = defaultAccountSchedulingThresholdErrorMessage + } else { + parsed.ErrorMessage = rawReason + } + } + return &parsed + } + + state.ErrorMessage = rawReason + return state +} diff --git a/backend/internal/service/account_scheduling_threshold_reason_test.go b/backend/internal/service/account_scheduling_threshold_reason_test.go new file mode 100644 index 0000000000..34eea91cd4 --- /dev/null +++ b/backend/internal/service/account_scheduling_threshold_reason_test.go @@ -0,0 +1,76 @@ +//go:build unit + +package service + +import ( + "encoding/json" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestBuildAccountSchedulingThresholdReason_UsesSourceAndFallbackMessage(t *testing.T) { + raw := BuildAccountSchedulingThresholdReason(" \t ") + + var payload map[string]string + require.NoError(t, json.Unmarshal([]byte(raw), &payload)) + require.Equal(t, AccountSchedulingThresholdReasonSource, payload["source"]) + require.Equal(t, defaultAccountSchedulingThresholdErrorMessage, payload["error_message"]) + require.True(t, IsAccountSchedulingThresholdReason(raw)) +} + +func TestIsAccountSchedulingThresholdReason(t *testing.T) { + require.True(t, IsAccountSchedulingThresholdReason(BuildAccountSchedulingThresholdReason("threshold reached"))) + require.False(t, IsAccountSchedulingThresholdReason(BuildTempUnschedReasonPayload("", "temporary block"))) + require.False(t, IsAccountSchedulingThresholdReason("plain text reason")) +} + +func TestTempUnschedStateFromStoredReason_EmptyReasonUsesFallbackErrorMessage(t *testing.T) { + state := tempUnschedStateFromStoredReason(" \n ", 1735689600) + + require.NotNil(t, state) + require.Equal(t, int64(1735689600), state.UntilUnix) + require.Equal(t, defaultTempUnschedReasonErrorMessage, state.ErrorMessage) +} + +func TestTempUnschedStateFromStoredReason_MissingRuleIndexIsSystemRule(t *testing.T) { + state := tempUnschedStateFromStoredReason(`{"error_message":"system cooldown"}`, 123) + require.Equal(t, -1, state.RuleIndex) +} + +func TestTempUnschedStateFromStoredReason_SchedulingThresholdJSONWithoutMessageUsesThresholdFallback(t *testing.T) { + raw := `{"source":"` + AccountSchedulingThresholdReasonSource + `"}` + + state := tempUnschedStateFromStoredReason(raw, 1735689600) + + require.NotNil(t, state) + require.Equal(t, int64(1735689600), state.UntilUnix) + require.Equal(t, defaultAccountSchedulingThresholdErrorMessage, state.ErrorMessage) +} + +func TestBuildDetailedAccountSchedulingThresholdReason_IncludesReadableFields(t *testing.T) { + now := time.Unix(1735689600, 0).UTC() + until := now.Add(5 * time.Hour) + + raw := BuildDetailedAccountSchedulingThresholdReason(AccountSchedulingThresholdReasonInput{ + Platform: PlatformOpenAI, + Window: "7d", + ThresholdPercent: 90, + UsedPercent: 92.5, + Until: until, + Now: now, + }) + + var payload map[string]any + require.NoError(t, json.Unmarshal([]byte(raw), &payload)) + require.Equal(t, AccountSchedulingThresholdReasonSource, payload["source"]) + require.Equal(t, PlatformOpenAI, payload["platform"]) + require.Equal(t, "7d", payload["window"]) + require.Equal(t, float64(90), payload["threshold_percent"]) + require.Equal(t, float64(92.5), payload["used_percent"]) + require.Equal(t, float64(until.Unix()), payload["until_unix"]) + require.Equal(t, float64(now.Unix()), payload["triggered_at_unix"]) + require.Contains(t, payload["error_message"], "openai scheduling threshold reached") + require.Contains(t, payload["error_message"], "92.5% used >= 90%") +} diff --git a/backend/internal/service/account_service.go b/backend/internal/service/account_service.go index d6905d7329..65d4d3f268 100644 --- a/backend/internal/service/account_service.go +++ b/backend/internal/service/account_service.go @@ -232,7 +232,7 @@ func (s *AccountService) Create(ctx context.Context, req CreateAccountRequest) ( Notes: normalizeAccountNotes(req.Notes), Platform: req.Platform, Type: req.Type, - Credentials: req.Credentials, + Credentials: SanitizeStoredCredentials(req.Platform, req.Credentials), Extra: req.Extra, ProxyID: req.ProxyID, Concurrency: req.Concurrency, @@ -325,7 +325,7 @@ func (s *AccountService) Update(ctx context.Context, id int64, req UpdateAccount } if req.Credentials != nil { - account.Credentials = *req.Credentials + account.Credentials = SanitizeStoredCredentials(account.Platform, *req.Credentials) } if req.Extra != nil { diff --git a/backend/internal/service/account_test_service.go b/backend/internal/service/account_test_service.go index 63d404d45c..5a62598db2 100644 --- a/backend/internal/service/account_test_service.go +++ b/backend/internal/service/account_test_service.go @@ -5,14 +5,22 @@ import ( "bytes" "context" "crypto/rand" + "encoding/base64" + "encoding/binary" "encoding/hex" "encoding/json" "errors" "fmt" + "image" + _ "image/gif" + _ "image/jpeg" + _ "image/png" "io" "log" + "mime/multipart" "net/http" "net/http/httptest" + "net/url" "regexp" "strings" "sync" @@ -27,6 +35,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/util/urlvalidator" "github.com/gin-gonic/gin" "github.com/google/uuid" + "github.com/tidwall/gjson" ) // sseDataPrefix matches SSE data lines with optional whitespace after colon. @@ -46,16 +55,53 @@ type TestEvent struct { Status string `json:"status,omitempty"` Code string `json:"code,omitempty"` ImageURL string `json:"image_url,omitempty"` + // AudioURL / VideoURL are data: or https URLs for in-browser media players. + AudioURL string `json:"audio_url,omitempty"` + VideoURL string `json:"video_url,omitempty"` MimeType string `json:"mime_type,omitempty"` Data any `json:"data,omitempty"` Success bool `json:"success,omitempty"` Error string `json:"error,omitempty"` } +// AccountTestOptions carries optional media for admin connectivity tests. +// ImageDataURL / AudioDataURL are full data URLs (data:;base64,...). +type AccountTestOptions struct { + ImageDataURL string + AudioDataURL string +} + +func firstAccountTestOptions(opts []AccountTestOptions) AccountTestOptions { + if len(opts) == 0 { + return AccountTestOptions{} + } + return opts[0] +} + +// maxAccountTestMediaBytes caps inbound data-URL payloads for admin tests (~8 MiB). +const maxAccountTestMediaBytes = 8 << 20 + const ( defaultGeminiTextTestPrompt = "hi" defaultGeminiImageTestPrompt = "Generate a cute orange cat astronaut sticker on a clean pastel background." defaultOpenAIImageTestPrompt = "Generate a cute orange cat astronaut sticker on a clean pastel background." + defaultGrokImageTestPrompt = "Generate a cute orange cat astronaut sticker on a clean pastel background." + defaultGrokVideoTestPrompt = "A red ball bouncing once on a white floor, short simple motion." + defaultGrokSearchTestQuery = "xAI Grok" + defaultGrokTTSTestText = "Hello from Sub2API account connectivity test." + + // Grok account-test modes (admin UI). Empty / default / text = Responses probe. + // image/video may also be inferred from model_id when mode is default. + AccountTestModeGrokText = "text" + AccountTestModeGrokImage = "image" + AccountTestModeGrokVideo = "video" + AccountTestModeGrokSearch = "search" + AccountTestModeGrokTTS = "tts" + AccountTestModeGrokSTT = "stt" + AccountTestModeGrokRealtime = "realtime" + + defaultGrokRealtimeTestModel = "grok-voice-latest" + grokRealtimeProbeTimeout = 12 * time.Second ) // isOpenAIImageModel checks if the model is an OpenAI image generation model (e.g. gpt-image-2). @@ -63,6 +109,32 @@ func isOpenAIImageModel(model string) bool { return strings.HasPrefix(strings.ToLower(model), "gpt-image-") } +func isGrokVideoGenerationModel(model string) bool { + return isGrokVideoBillingModel(model) || + strings.HasPrefix(strings.ToLower(strings.TrimSpace(model)), "grok-video") +} + +func normalizeGrokAccountTestMode(mode string) string { + switch strings.ToLower(strings.TrimSpace(mode)) { + case AccountTestModeGrokText: + return AccountTestModeGrokText + case AccountTestModeGrokImage: + return AccountTestModeGrokImage + case AccountTestModeGrokVideo: + return AccountTestModeGrokVideo + case AccountTestModeGrokSearch: + return AccountTestModeGrokSearch + case AccountTestModeGrokTTS: + return AccountTestModeGrokTTS + case AccountTestModeGrokSTT: + return AccountTestModeGrokSTT + case AccountTestModeGrokRealtime: + return AccountTestModeGrokRealtime + default: + return AccountTestModeDefault + } +} + // AccountTestService handles account testing operations type AccountTestService struct { accountRepo AccountRepository @@ -72,9 +144,19 @@ type AccountTestService struct { antigravityGatewayService *AntigravityGatewayService httpUpstream HTTPUpstream cfg *config.Config + settingService *SettingService tlsFPProfileService *TLSFingerprintProfileService agentIdentityTaskMu sync.Mutex agentIdentityWS agentIdentityWSConnectionInvalidator + // grokWSDialer is optional; realtime account tests use the default OpenAI-style + // WS dialer when nil (supports proxy + coder/websocket handshake). + grokWSDialer openAIWSClientDialer +} + +func (s *AccountTestService) SetSettingService(settingService *SettingService) { + if s != nil { + s.settingService = settingService + } } // NewAccountTestService creates a new AccountTestService @@ -177,8 +259,10 @@ func createTestPayload(modelID string) (map[string]any, error) { // All account types use full Claude Code client characteristics, only auth header differs // modelID is optional - if empty, defaults to claude.DefaultTestModel // mode is optional - "compact" routes OpenAI accounts to the /responses/compact probe path -func (s *AccountTestService) TestAccountConnection(c *gin.Context, accountID int64, modelID string, prompt string, mode string) error { +// opts is optional media (image/audio data URLs for real generation / STT). +func (s *AccountTestService) TestAccountConnection(c *gin.Context, accountID int64, modelID string, prompt string, mode string, opts ...AccountTestOptions) error { ctx := c.Request.Context() + testOpts := firstAccountTestOptions(opts) // Get account account, err := s.accountRepo.GetByID(ctx, accountID) @@ -210,7 +294,7 @@ func (s *AccountTestService) TestAccountConnection(c *gin.Context, accountID int } if account.Platform == PlatformGrok { - return s.testGrokAccountConnection(c, account, modelID) + return s.testGrokAccountConnection(c, account, modelID, prompt, mode, testOpts) } if account.Platform == PlatformAntigravity { @@ -706,14 +790,61 @@ func (s *AccountTestService) testOpenAIAccountConnection(c *gin.Context, account return s.processOpenAIStream(c, resp.Body) } -// testGrokAccountConnection tests a Grok OAuth or API-key account through xAI's Responses API. -func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account *Account, modelID string) error { +// testGrokAccountConnection routes Grok admin connectivity tests by explicit mode first, +// then by selected model family for media. Standalone modes (search/tts/stt) never share +// the text Responses path; image/video never hit Responses either. +// +// Modes: +// - default/text → Responses (optional model) +// - image → /v1/images/generations (model optional; defaults to grok-imagine-image) +// - video → /v1/videos/generations (model optional; defaults to grok-imagine-video) +// - search → standalone web-search probe (gateway /v1/web_search semantics) +// - tts → HTTP /v1/tts +// - stt → HTTP /v1/stt (synthetic tiny wav probe) +// - realtime → WS /v1/realtime dial + optional first server event +// +// When mode is default, image/video can still be inferred from model_id for backward compat. +func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account *Account, modelID, prompt, mode string, opts AccountTestOptions) error { ctx := c.Request.Context() - if s.httpUpstream == nil { + // Realtime is WebSocket-only and does not need HTTP upstream. + mode = normalizeGrokAccountTestMode(mode) + if mode != AccountTestModeGrokRealtime && s.httpUpstream == nil { return s.sendErrorAndEnd(c, "HTTP upstream not configured") } + authToken, err := s.grokTestAccessToken(ctx, account) + if err != nil { + return s.sendErrorAndEnd(c, err.Error()) + } + + // Explicit standalone / media modes always win over model id. + switch mode { + case AccountTestModeGrokSearch: + return s.testGrokWebSearch(c, ctx, account, authToken, prompt) + case AccountTestModeGrokTTS: + return s.testGrokTTS(c, ctx, account, authToken, prompt) + case AccountTestModeGrokSTT: + return s.testGrokSTT(c, ctx, account, authToken, opts.AudioDataURL) + case AccountTestModeGrokRealtime: + return s.testGrokRealtime(c, ctx, account, authToken, modelID) + case AccountTestModeGrokImage: + return s.testGrokImageGeneration(c, ctx, account, authToken, resolveGrokImageTestModel(account, modelID), resolveGrokImagePrompt(prompt), opts.ImageDataURL) + case AccountTestModeGrokVideo: + return s.testGrokVideoGeneration(c, ctx, account, authToken, resolveGrokVideoTestModel(account, modelID), resolveGrokVideoPrompt(prompt), opts) + case AccountTestModeGrokText: + // Force text Responses even if model_id looks like media. + testModelID := strings.TrimSpace(modelID) + if testModelID == "" { + testModelID = grokDefaultResponsesModel + } + if mapped := strings.TrimSpace(account.GetMappedModel(testModelID)); mapped != "" { + testModelID = mapped + } + return s.testGrokResponsesConnection(c, ctx, account, authToken, testModelID) + } + + // mode == default: infer from model family (legacy UI / API clients). testModelID := strings.TrimSpace(modelID) if testModelID == "" { testModelID = grokDefaultResponsesModel @@ -722,73 +853,121 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * testModelID = mapped } - var authToken string + switch { + case isGrokImageGenerationModel(testModelID): + return s.testGrokImageGeneration(c, ctx, account, authToken, testModelID, resolveGrokImagePrompt(prompt), opts.ImageDataURL) + case isGrokVideoGenerationModel(testModelID): + return s.testGrokVideoGeneration(c, ctx, account, authToken, testModelID, resolveGrokVideoPrompt(prompt), opts) + default: + return s.testGrokResponsesConnection(c, ctx, account, authToken, testModelID) + } +} + +func resolveGrokImagePrompt(prompt string) string { + if strings.TrimSpace(prompt) == "" { + return defaultGrokImageTestPrompt + } + return strings.TrimSpace(prompt) +} + +func resolveGrokVideoPrompt(prompt string) string { + if strings.TrimSpace(prompt) == "" { + return defaultGrokVideoTestPrompt + } + return strings.TrimSpace(prompt) +} + +func resolveGrokImageTestModel(account *Account, modelID string) string { + testModelID := strings.TrimSpace(modelID) + if testModelID == "" { + testModelID = "grok-imagine-image" + } + if mapped := strings.TrimSpace(account.GetMappedModel(testModelID)); mapped != "" { + return mapped + } + return testModelID +} + +func resolveGrokVideoTestModel(account *Account, modelID string) string { + testModelID := strings.TrimSpace(modelID) + if testModelID == "" { + testModelID = "grok-imagine-video" + } + if mapped := strings.TrimSpace(account.GetMappedModel(testModelID)); mapped != "" { + return mapped + } + return testModelID +} + +func (s *AccountTestService) grokTestAccessToken(ctx context.Context, account *Account) (string, error) { switch account.Type { case AccountTypeOAuth: if s.grokTokenProvider == nil { - return s.sendErrorAndEnd(c, "Grok token provider not configured") + return "", fmt.Errorf("grok token provider not configured") } - var err error - // 手动测试不走生产调度资格门:关闭调度、限流/过载/临时冷却中的账号 - // 也应能被管理员探测(#4598),与 Codex/OpenAI 测试行为一致。 - authToken, err = s.grokTokenProvider.GetAccessTokenForManualTest(ctx, account) + // Manual tests skip production scheduling eligibility so paused/rate-limited + // accounts can still be probed by admins (same as Codex/OpenAI tests). + token, err := s.grokTokenProvider.GetAccessTokenForManualTest(ctx, account) if err != nil { - return s.sendErrorAndEnd(c, fmt.Sprintf("Failed to get Grok access token: %s", err.Error())) + return "", fmt.Errorf("failed to get grok access token: %s", err.Error()) } + return token, nil case AccountTypeAPIKey: - authToken = strings.TrimSpace(account.GetCredential("api_key")) + authToken := strings.TrimSpace(account.GetCredential("api_key")) if authToken == "" { - return s.sendErrorAndEnd(c, "Grok API key is missing") + return "", fmt.Errorf("grok api key is missing") } + return authToken, nil default: - return s.sendErrorAndEnd(c, fmt.Sprintf("Unsupported Grok account type: %s", account.Type)) + return "", fmt.Errorf("unsupported grok account type: %s", account.Type) } +} - apiURL, err := buildGrokResponsesURL(account, s.cfg) - if err != nil { - return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok base URL: %s", err.Error())) +func (s *AccountTestService) grokTestProxyURL(account *Account) string { + if account.ProxyID != nil && account.Proxy != nil { + return account.Proxy.URL() } + return "" +} +func (s *AccountTestService) prepareGrokTestSSE(c *gin.Context) { c.Writer.Header().Set("Content-Type", "text/event-stream") c.Writer.Header().Set("Cache-Control", "no-cache") c.Writer.Header().Set("Connection", "keep-alive") c.Writer.Header().Set("X-Accel-Buffering", "no") c.Writer.Flush() +} - payloadBytes, err := buildGrokQuotaProbeBody(testModelID) - if err != nil { - return s.sendErrorAndEnd(c, "Failed to create Grok test payload") - } - - if !agentIdentityTaskRecoveryWasTried(ctx) { - s.sendEvent(c, TestEvent{Type: "test_start", Model: testModelID}) - } - - req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(payloadBytes)) - if err != nil { - return s.sendErrorAndEnd(c, "Failed to create Grok request") - } +func (s *AccountTestService) applyGrokTestRequestHeaders(req *http.Request, account *Account, authToken string, accept string) { req.Header.Set("Content-Type", "application/json") - req.Header.Set("Accept", "application/json, text/event-stream") + if accept != "" { + req.Header.Set("Accept", accept) + } req.Header.Set("Authorization", "Bearer "+authToken) - if account.IsGrokOAuth() { + // Match gateway media/voice: CLI identity headers only on the CLI chat proxy. + // api.x.ai media (images/videos) rejects or mistreats OAuth when CLI headers + // are stamped on the official API host (e.g. ZDR upload_url false positives). + if account.IsGrokOAuth() && req.URL != nil && isGrokCLIProxyTarget(req.URL.String()) { applyGrokCLIHeaders(req.Header) } - // 连通性测试与真实转发保持同一套账号级请求头覆写。 account.ApplyHeaderOverrides(req.Header) +} - proxyURL := "" - if account.ProxyID != nil && account.Proxy != nil { - proxyURL = account.Proxy.URL() +func (s *AccountTestService) observeGrokTestResponse(ctx context.Context, account *Account, resp *http.Response) { + if resp == nil { + return } - - resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) - if err != nil { - return s.sendErrorAndEnd(c, fmt.Sprintf("Grok Responses API request failed: %s", err.Error())) - } - defer func() { _ = resp.Body.Close() }() - now := time.Now() + // Error bodies carry Grok's free-usage, billing, and content-policy + // classifications when quota headers are absent. Read only non-success + // responses here, then restore the body because the caller still needs it + // for the user-facing test result. + var responseBody []byte + if resp.StatusCode >= http.StatusBadRequest && resp.Body != nil { + responseBody, _ = io.ReadAll(resp.Body) + _ = resp.Body.Close() + resp.Body = io.NopCloser(bytes.NewReader(responseBody)) + } snapshot := parseGrokQuotaSnapshot(resp.Header, resp.StatusCode, now) if snapshot != nil && s.accountRepo != nil { resetAt, limited := grokRateLimitResetAtForAccount(account, snapshot, now) @@ -806,25 +985,933 @@ func (s *AccountTestService) testGrokAccountConnection(c *gin.Context, account * } else if s.accountRepo != nil && isSuccessfulGrokRateLimitRecovery(account, &xai.QuotaSnapshot{StatusCode: resp.StatusCode}) { clearGrokRateLimitAfterRecovery(ctx, s.accountRepo, account) } - - if resp.StatusCode != http.StatusOK { - body, _ := io.ReadAll(resp.Body) + if s.accountRepo == nil || len(responseBody) == 0 { if resp.StatusCode == http.StatusPaymentRequired && s.accountRepo != nil { stateCtx, cancel := openAIAccountStateContext(ctx) defer cancel() - _ = s.accountRepo.SetTempUnschedulable( - stateCtx, - account.ID, - time.Now().Add(30*time.Minute), - "grok payment required", - ) + _ = s.accountRepo.SetTempUnschedulable(stateCtx, account.ID, now.Add(30*time.Minute), "grok payment required") } + return + } + if isGrokContentPolicyRejection(resp.StatusCode, responseBody) { + return + } + decision := classifyGrokUpstreamFailure(resp.StatusCode, responseBody, "") + if decision.Class == GrokFailureFreeUsage { + if resetAt, limited := grokRateLimitResetAtForAccount(account, snapshot, now); limited && resetAt.After(now) { + persistGrokRateLimit(ctx, s.accountRepo, account, resetAt) + } else { + stateCtx, cancel := openAIAccountStateContext(ctx) + _ = s.accountRepo.SetTempUnschedulable(stateCtx, account.ID, now.Add(grokFreeUsageProbeCooldown), "grok free usage exhausted") + cancel() + } + return + } + if decision.Class == GrokFailureBilling && (isGrokSpendingLimitError(responseBody) || strings.Contains(strings.ToLower(decision.Reason), "credit")) { + persistGrokRateLimit(ctx, s.accountRepo, account, grokSpendingLimitResetAt(account, now)) + return + } + cooldown := time.Duration(0) + reason := "" + switch resp.StatusCode { + case http.StatusUnauthorized: + cooldown, reason = 10*time.Minute, "grok oauth token unauthorized" + case http.StatusPaymentRequired: + cooldown, reason = 30*time.Minute, "grok payment required" + case http.StatusForbidden: + cooldown, reason = 30*time.Minute, "grok entitlement or subscription tier denied" + default: + if resp.StatusCode >= 500 { + cooldown, reason = 2*time.Minute, "grok upstream temporary error" + } + } + if decision.Class == GrokFailureBilling && cooldown == 0 { + cooldown, reason = 30*time.Minute, "grok payment required" + } + if cooldown > 0 { + stateCtx, cancel := openAIAccountStateContext(ctx) + defer cancel() + until := now.Add(cooldown) + if account.TempUnschedulableUntil != nil && account.TempUnschedulableUntil.After(until) { + until = *account.TempUnschedulableUntil + } + _ = s.accountRepo.SetTempUnschedulable( + stateCtx, + account.ID, + until, + reason, + ) + } +} + +func (s *AccountTestService) testGrokResponsesConnection(c *gin.Context, ctx context.Context, account *Account, authToken, testModelID string) error { + apiURL, err := buildGrokResponsesURL(account, s.cfg, s.settingService) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok base URL: %s", err.Error())) + } + + s.prepareGrokTestSSE(c) + + payloadBytes, err := buildGrokQuotaProbeBody(testModelID) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create Grok test payload") + } + + if !agentIdentityTaskRecoveryWasTried(ctx) { + s.sendEvent(c, TestEvent{Type: "test_start", Model: testModelID}) + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(payloadBytes)) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create Grok request") + } + s.applyGrokTestRequestHeaders(req, account, authToken, "application/json, text/event-stream") + + resp, err := s.httpUpstream.Do(req, s.grokTestProxyURL(account), account.ID, account.Concurrency) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok Responses API request failed: %s", err.Error())) + } + defer func() { _ = resp.Body.Close() }() + + s.observeGrokTestResponse(ctx, account, resp) + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) return s.sendErrorAndEnd(c, fmt.Sprintf("Grok Responses API returned %d: %s", resp.StatusCode, string(body))) } return s.processOpenAIStream(c, resp.Body) } +func (s *AccountTestService) testGrokImageGeneration(c *gin.Context, ctx context.Context, account *Account, authToken, modelID, prompt, imageDataURL string) error { + // With a source image, prefer /images/edits; otherwise /images/generations. + endpoint := GrokMediaEndpointImagesGenerations + imageDataURL = strings.TrimSpace(imageDataURL) + hasSourceImage := imageDataURL != "" + if hasSourceImage { + endpoint = GrokMediaEndpointImagesEdits + } + + // Align model aliases with gateway (e.g. grok-imagine → grok-imagine-image-quality). + modelID = NormalizeGrokMediaModelForEndpoint(endpoint, modelID, hasSourceImage) + if modelID == "" { + modelID = "grok-imagine-image-quality" + } + + apiURL, err := buildGrokMediaURL(account, s.cfg, endpoint, "") + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok media base URL: %s", err.Error())) + } + + s.prepareGrokTestSSE(c) + s.sendEvent(c, TestEvent{Type: "test_start", Model: modelID}) + if endpoint == GrokMediaEndpointImagesEdits { + s.sendEvent(c, TestEvent{Type: "status", Text: "Calling Grok /v1/images/edits with uploaded source image..."}) + } else { + s.sendEvent(c, TestEvent{Type: "status", Text: "Calling Grok /v1/images/generations..."}) + } + + // Zero-data-retention teams reject URL format; always request base64 for admin tests. + payload := map[string]any{ + "model": modelID, + "prompt": prompt, + "n": 1, + "response_format": "b64_json", + } + if hasSourceImage { + normalized, err := normalizeAccountTestImageDataURL(imageDataURL) + if err != nil { + return s.sendErrorAndEnd(c, err.Error()) + } + // Match gateway prepareGrokMediaForwardBody shape: {url, type:image_url}. + payload["image"] = grokMediaImageObject(normalized) + s.sendEvent(c, TestEvent{Type: "content", Text: fmt.Sprintf("source image ready (%d chars data URL)\n", len(normalized))}) + } + payloadBytes, err := json.Marshal(payload) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to marshal Grok image request") + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(payloadBytes)) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create Grok image request") + } + s.applyGrokTestRequestHeaders(req, account, authToken, "application/json") + req.ContentLength = int64(len(payloadBytes)) + req.GetBody = func() (io.ReadCloser, error) { + return io.NopCloser(bytes.NewReader(payloadBytes)), nil + } + + // One retry on transport EOF (proxies occasionally drop large edit payloads). + var resp *http.Response + var doErr error + for attempt := 0; attempt < 2; attempt++ { + if attempt > 0 { + s.sendEvent(c, TestEvent{Type: "status", Text: "Retrying Grok image request after transport error..."}) + req, err = http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(payloadBytes)) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create Grok image retry request") + } + s.applyGrokTestRequestHeaders(req, account, authToken, "application/json") + req.ContentLength = int64(len(payloadBytes)) + } + resp, doErr = s.httpUpstream.Do(req, s.grokTestProxyURL(account), account.ID, account.Concurrency) + if doErr == nil { + break + } + if !isTransientGrokTransportError(doErr) || attempt == 1 { + return s.sendErrorAndEnd(c, formatGrokImageTransportError(doErr, hasSourceImage, len(payloadBytes))) + } + } + defer func() { _ = resp.Body.Close() }() + s.observeGrokTestResponse(ctx, account, resp) + + body, err := io.ReadAll(resp.Body) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Failed to read Grok image response: %s", err.Error())) + } + if resp.StatusCode != http.StatusOK { + return s.sendErrorAndEnd(c, formatGrokImagesAPIError(resp.StatusCode, body, hasSourceImage)) + } + + var result struct { + Data []struct { + URL string `json:"url"` + B64JSON string `json:"b64_json"` + RevisedPrompt string `json:"revised_prompt"` + MimeType string `json:"mime_type"` + } `json:"data"` + } + if err := json.Unmarshal(body, &result); err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Failed to parse Grok image response: %s", err.Error())) + } + if len(result.Data) == 0 { + return s.sendErrorAndEnd(c, "No images returned from Grok API") + } + + for _, item := range result.Data { + if item.RevisedPrompt != "" { + s.sendEvent(c, TestEvent{Type: "content", Text: item.RevisedPrompt}) + } + mimeType := strings.TrimSpace(item.MimeType) + if mimeType == "" { + mimeType = "image/jpeg" + } + switch { + case strings.TrimSpace(item.B64JSON) != "": + s.sendEvent(c, TestEvent{ + Type: "image", + ImageURL: "data:" + mimeType + ";base64," + item.B64JSON, + MimeType: mimeType, + }) + case strings.TrimSpace(item.URL) != "": + s.sendEvent(c, TestEvent{Type: "image", ImageURL: item.URL, MimeType: mimeType}) + } + } + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil +} + +func (s *AccountTestService) testGrokVideoGeneration(c *gin.Context, ctx context.Context, account *Account, authToken, modelID, prompt string, opts AccountTestOptions) error { + apiURL, err := buildGrokMediaURL(account, s.cfg, GrokMediaEndpointVideosGenerations, "") + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok media base URL: %s", err.Error())) + } + + s.prepareGrokTestSSE(c) + s.sendEvent(c, TestEvent{Type: "test_start", Model: modelID}) + s.sendEvent(c, TestEvent{Type: "status", Text: "Calling Grok /v1/videos/generations..."}) + + payload := map[string]any{ + "model": modelID, + "prompt": prompt, + "duration": 6, + "aspect_ratio": "16:9", + "resolution": "480p", + } + if img := strings.TrimSpace(opts.ImageDataURL); img != "" { + normalized, err := normalizeAccountTestImageDataURL(img) + if err != nil { + return s.sendErrorAndEnd(c, err.Error()) + } + // First-frame / image-to-video input (xAI image field). + payload["image"] = grokMediaImageObject(normalized) + s.sendEvent(c, TestEvent{Type: "content", Text: "using uploaded first-frame / reference image\n"}) + } + payloadBytes, _ := json.Marshal(payload) + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(payloadBytes)) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create Grok video request") + } + s.applyGrokTestRequestHeaders(req, account, authToken, "application/json") + + resp, err := s.httpUpstream.Do(req, s.grokTestProxyURL(account), account.ID, account.Concurrency) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok video request failed: %s", err.Error())) + } + defer func() { _ = resp.Body.Close() }() + s.observeGrokTestResponse(ctx, account, resp) + + body, err := io.ReadAll(resp.Body) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Failed to read Grok video response: %s", err.Error())) + } + if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusAccepted && resp.StatusCode != http.StatusCreated { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok videos API returned %d: %s", resp.StatusCode, string(body))) + } + + requestID := strings.TrimSpace(gjson.GetBytes(body, "request_id").String()) + if requestID == "" { + requestID = strings.TrimSpace(gjson.GetBytes(body, "id").String()) + } + if requestID == "" { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok video create response missing request_id: %s", string(body))) + } + + s.sendEvent(c, TestEvent{Type: "content", Text: fmt.Sprintf("video request accepted: %s\n", requestID)}) + s.sendEvent(c, TestEvent{Type: "status", Text: "Polling video status until done (max ~60s)..."}) + + statusURL, err := buildGrokMediaURL(account, s.cfg, GrokMediaEndpointVideoStatus, requestID) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok video status URL: %s", err.Error())) + } + + deadline := time.Now().Add(60 * time.Second) + for time.Now().Before(deadline) { + if ctx.Err() != nil { + return s.sendErrorAndEnd(c, "Grok video poll canceled") + } + statusReq, err := http.NewRequestWithContext(ctx, http.MethodGet, statusURL, nil) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create Grok video status request") + } + s.applyGrokTestRequestHeaders(statusReq, account, authToken, "application/json") + statusResp, err := s.httpUpstream.Do(statusReq, s.grokTestProxyURL(account), account.ID, account.Concurrency) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok video status failed: %s", err.Error())) + } + statusBody, _ := io.ReadAll(statusResp.Body) + _ = statusResp.Body.Close() + if statusResp.StatusCode != http.StatusOK && statusResp.StatusCode != http.StatusAccepted { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok video status returned %d: %s", statusResp.StatusCode, string(statusBody))) + } + st := strings.ToLower(strings.TrimSpace(gjson.GetBytes(statusBody, "status").String())) + progress := gjson.GetBytes(statusBody, "progress") + if progress.Exists() { + s.sendEvent(c, TestEvent{Type: "status", Text: fmt.Sprintf("status=%s progress=%v", st, progress.Value())}) + } else { + s.sendEvent(c, TestEvent{Type: "status", Text: "status=" + st}) + } + switch st { + case "done", "completed", "succeeded", "success": + return s.emitGrokVideoResult(c, ctx, account, authToken, requestID, statusBody) + case "failed", "error", "canceled", "cancelled": + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok video failed: %s", string(statusBody))) + } + select { + case <-ctx.Done(): + return s.sendErrorAndEnd(c, "Grok video poll canceled") + case <-time.After(3 * time.Second): + } + } + return s.sendErrorAndEnd(c, "Grok video still processing after 60s (request_id="+requestID+")") +} + +// emitGrokVideoResult surfaces a playable video URL or downloads /content as data URL. +func (s *AccountTestService) emitGrokVideoResult(c *gin.Context, ctx context.Context, account *Account, authToken, requestID string, statusBody []byte) error { + videoURL := firstNonEmpty( + strings.TrimSpace(gjson.GetBytes(statusBody, "video.url").String()), + strings.TrimSpace(gjson.GetBytes(statusBody, "url").String()), + strings.TrimSpace(gjson.GetBytes(statusBody, "video_url").String()), + strings.TrimSpace(gjson.GetBytes(statusBody, "download_url").String()), + ) + if videoURL != "" && (strings.HasPrefix(videoURL, "http://") || strings.HasPrefix(videoURL, "https://") || strings.HasPrefix(videoURL, "data:")) { + s.sendEvent(c, TestEvent{Type: "content", Text: "video ready: " + videoURL + "\n"}) + s.sendEvent(c, TestEvent{Type: "video", VideoURL: videoURL, MimeType: "video/mp4"}) + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil + } + + // Fetch binary content via official /videos/{id}/content (Bearer-authenticated). + contentURL, err := buildGrokMediaURL(account, s.cfg, GrokMediaEndpointVideoContent, requestID) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok video content URL: %s", err.Error())) + } + s.sendEvent(c, TestEvent{Type: "status", Text: "Downloading video content for preview..."}) + req, err := http.NewRequestWithContext(ctx, http.MethodGet, contentURL, nil) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create Grok video content request") + } + s.applyGrokTestRequestHeaders(req, account, authToken, "video/*, application/octet-stream, */*") + resp, err := s.httpUpstream.Do(req, s.grokTestProxyURL(account), account.ID, account.Concurrency) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok video content download failed: %s", err.Error())) + } + defer func() { _ = resp.Body.Close() }() + body, _ := io.ReadAll(io.LimitReader(resp.Body, 64<<20)) // 64 MiB cap for admin preview + if resp.StatusCode != http.StatusOK { + // Fall back to status URL when binary content is unavailable. + if videoURL != "" { + s.sendEvent(c, TestEvent{Type: "content", Text: "video completed; content download unavailable, reported url=" + videoURL + "\n"}) + s.sendEvent(c, TestEvent{Type: "video", VideoURL: videoURL, MimeType: "video/mp4"}) + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil + } + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok video content returned %d: %s", resp.StatusCode, truncateString(string(body), 300))) + } + ct := resp.Header.Get("Content-Type") + if ct == "" || strings.HasPrefix(ct, "application/octet-stream") { + ct = "video/mp4" + } + // Keep only type/subtype for data URL. + if i := strings.Index(ct, ";"); i >= 0 { + ct = strings.TrimSpace(ct[:i]) + } + dataURL := "data:" + ct + ";base64," + base64.StdEncoding.EncodeToString(body) + s.sendEvent(c, TestEvent{Type: "content", Text: fmt.Sprintf("video content downloaded: content-type=%s bytes=%d\n", ct, len(body))}) + s.sendEvent(c, TestEvent{Type: "video", VideoURL: dataURL, MimeType: ct}) + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil +} + +func (s *AccountTestService) testGrokWebSearch(c *gin.Context, ctx context.Context, account *Account, authToken, query string) error { + query = strings.TrimSpace(query) + if query == "" { + query = defaultGrokSearchTestQuery + } + + // Account-test "web_search" mode mirrors the standalone gateway endpoint + // POST /v1/web_search (not a free-form chat with tools). Implementation still + // uses the same DoGrokNativeResponsesJSON helper as the gateway handler so + // results match production search. + s.prepareGrokTestSSE(c) + s.sendEvent(c, TestEvent{Type: "test_start", Model: "grok-web-search"}) + s.sendEvent(c, TestEvent{Type: "status", Text: "Calling standalone web_search probe (same as gateway /v1/web_search)..."}) + + // Keep parity with handler.buildGrokWebSearchPrompt / include sources. + const maxResults = 5 + prompt := fmt.Sprintf( + `Search the web for the user query below. Return ONLY valid JSON with this exact shape: {"results":[{"url":"https://...","title":"page title","snippet":"concise factual summary"}]}. Return at most %d unique results. Every URL must be an actual web_search source. Populate a non-empty title and snippet for every result. Do not wrap the JSON in markdown. + +User query: +%s`, maxResults, query) + payload := map[string]any{ + "model": grokDefaultResponsesModel, + "input": prompt, + "tools": []map[string]any{{"type": "web_search"}}, + "include": []string{"web_search_call.action.sources"}, + "store": false, + "stream": false, + } + payloadBytes, _ := json.Marshal(payload) + + apiURL, err := buildGrokResponsesURL(account, s.cfg, s.settingService) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok base URL: %s", err.Error())) + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(payloadBytes)) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create standalone web_search probe request") + } + s.applyGrokTestRequestHeaders(req, account, authToken, "application/json") + + resp, err := s.httpUpstream.Do(req, s.grokTestProxyURL(account), account.ID, account.Concurrency) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("standalone web_search probe failed: %s", err.Error())) + } + defer func() { _ = resp.Body.Close() }() + s.observeGrokTestResponse(ctx, account, resp) + + body, _ := io.ReadAll(resp.Body) + if resp.StatusCode != http.StatusOK { + return s.sendErrorAndEnd(c, fmt.Sprintf("standalone web_search probe returned %d: %s", resp.StatusCode, string(body))) + } + + // Normalize like gateway extractGrokWebSearchSources (URL-only sources are enough for connectivity). + sourceCount := 0 + gjson.GetBytes(body, "output").ForEach(func(_, item gjson.Result) bool { + if item.Get("type").String() != "web_search_call" { + return true + } + sources := item.Get("action.sources") + if sources.IsArray() { + sourceCount += len(sources.Array()) + } + return true + }) + searchCount := countGrokNativeSearchCallsFromJSONBytes(body) + s.sendEvent(c, TestEvent{Type: "content", Text: fmt.Sprintf("web_search ok: query=%q tool_calls=%d sources=%d\n", query, searchCount, sourceCount)}) + // Optional: first structured result title if model returned JSON text. + gjson.GetBytes(body, "output").ForEach(func(_, item gjson.Result) bool { + if item.Get("type").String() != "message" { + return true + } + for _, part := range item.Get("content").Array() { + text := strings.TrimSpace(part.Get("text").String()) + if text == "" { + continue + } + if len(text) > 300 { + text = text[:300] + "..." + } + s.sendEvent(c, TestEvent{Type: "content", Text: text + "\n"}) + return false + } + return true + }) + if searchCount == 0 && sourceCount == 0 { + return s.sendErrorAndEnd(c, "standalone web_search probe completed but no search sources/tool calls were observed") + } + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil +} + +func (s *AccountTestService) testGrokTTS(c *gin.Context, ctx context.Context, account *Account, authToken, text string) error { + text = strings.TrimSpace(text) + if text == "" { + text = defaultGrokTTSTestText + } + apiURL, err := buildGrokVoiceURL(account, s.cfg, "tts") + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok TTS URL: %s", err.Error())) + } + + s.prepareGrokTestSSE(c) + s.sendEvent(c, TestEvent{Type: "test_start", Model: "grok-voice-tts"}) + s.sendEvent(c, TestEvent{Type: "status", Text: "Calling standalone /v1/tts..."}) + + // xAI requires `language`; optional voice_id. Prefer the shape that matches + // live gateway probes (text + language [+ voice_id]). + payloads := []map[string]any{ + {"text": text, "language": "en", "voice_id": "Ara"}, + {"text": text, "language": "en"}, + {"text": text, "language": "English", "voice_id": "Ara"}, + } + var lastBody string + var lastCode int + for _, payload := range payloads { + payloadBytes, _ := json.Marshal(payload) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewReader(payloadBytes)) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create Grok TTS request") + } + s.applyGrokTestRequestHeaders(req, account, authToken, "audio/*, application/json, */*") + resp, err := s.httpUpstream.Do(req, s.grokTestProxyURL(account), account.ID, account.Concurrency) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok TTS failed: %s", err.Error())) + } + body, _ := io.ReadAll(resp.Body) + _ = resp.Body.Close() + s.observeGrokTestResponse(ctx, account, resp) + lastCode = resp.StatusCode + lastBody = string(body) + if resp.StatusCode == http.StatusOK { + ct := resp.Header.Get("Content-Type") + if ct == "" { + ct = "audio/mpeg" + } + if i := strings.Index(ct, ";"); i >= 0 { + ct = strings.TrimSpace(ct[:i]) + } + // Cap preview size so SSE stays manageable (~4 MiB audio). + if len(body) > 4<<20 { + body = body[:4<<20] + } + audioURL := "data:" + ct + ";base64," + base64.StdEncoding.EncodeToString(body) + s.sendEvent(c, TestEvent{Type: "content", Text: fmt.Sprintf("tts ok: content-type=%s bytes=%d\n", ct, len(body))}) + s.sendEvent(c, TestEvent{Type: "audio", AudioURL: audioURL, MimeType: ct}) + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil + } + if resp.StatusCode < 400 || resp.StatusCode >= 500 { + break + } + } + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok TTS returned %d: %s", lastCode, lastBody)) +} + +// testGrokSTT posts audio to /v1/stt. When audioDataURL is set, uses the +// uploaded file; otherwise a tiny synthetic silent WAV for connectivity only. +func (s *AccountTestService) testGrokSTT(c *gin.Context, ctx context.Context, account *Account, authToken, audioDataURL string) error { + apiURL, err := buildGrokVoiceURL(account, s.cfg, "stt") + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok STT URL: %s", err.Error())) + } + + s.prepareGrokTestSSE(c) + s.sendEvent(c, TestEvent{Type: "test_start", Model: "grok-voice-stt"}) + + var audioBytes []byte + filename := "probe.wav" + if audioDataURL = strings.TrimSpace(audioDataURL); audioDataURL != "" { + if err := validateAccountTestDataURL(audioDataURL, "audio/"); err != nil { + return s.sendErrorAndEnd(c, err.Error()) + } + raw, mime, err := decodeAccountTestDataURL(audioDataURL) + if err != nil { + return s.sendErrorAndEnd(c, "Invalid audio data URL: "+err.Error()) + } + audioBytes = raw + filename = sttFilenameForMIME(mime) + s.sendEvent(c, TestEvent{Type: "status", Text: "Calling standalone /v1/stt with uploaded audio..."}) + s.sendEvent(c, TestEvent{Type: "content", Text: fmt.Sprintf("uploaded audio: mime=%s bytes=%d\n", mime, len(audioBytes))}) + } else { + audioBytes = minimalSilentWAV() + s.sendEvent(c, TestEvent{Type: "status", Text: "Calling standalone /v1/stt with a synthetic silent WAV..."}) + } + + var bodyBuf bytes.Buffer + w := multipart.NewWriter(&bodyBuf) + part, err := w.CreateFormFile("file", filename) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to build STT multipart body") + } + if _, err := part.Write(audioBytes); err != nil { + return s.sendErrorAndEnd(c, "Failed to write STT audio part") + } + _ = w.WriteField("model", "grok-stt") + _ = w.WriteField("language", "en") + if err := w.Close(); err != nil { + return s.sendErrorAndEnd(c, "Failed to finalize STT multipart body") + } + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, &bodyBuf) + if err != nil { + return s.sendErrorAndEnd(c, "Failed to create Grok STT request") + } + req.Header.Set("Content-Type", w.FormDataContentType()) + req.Header.Set("Accept", "application/json") + req.Header.Set("Authorization", "Bearer "+authToken) + if account.IsGrokOAuth() { + applyGrokCLIHeaders(req.Header) + } + account.ApplyHeaderOverrides(req.Header) + + resp, err := s.httpUpstream.Do(req, s.grokTestProxyURL(account), account.ID, account.Concurrency) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok STT failed: %s", err.Error())) + } + defer func() { _ = resp.Body.Close() }() + s.observeGrokTestResponse(ctx, account, resp) + respBody, _ := io.ReadAll(resp.Body) + if resp.StatusCode != http.StatusOK { + // 4xx on synthetic audio still proves the STT endpoint is wired; report clearly. + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok STT returned %d: %s", resp.StatusCode, string(respBody))) + } + text := strings.TrimSpace(gjson.GetBytes(respBody, "text").String()) + if text == "" { + text = strings.TrimSpace(string(respBody)) + if len(text) > 200 { + text = text[:200] + "..." + } + } + s.sendEvent(c, TestEvent{Type: "content", Text: fmt.Sprintf("stt ok: %s\n", text)}) + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil +} + +// testGrokRealtime dials the standalone xAI Voice Realtime WebSocket +// (wss://api.x.ai/v1/realtime?model=...) to verify auth + endpoint reachability. +// It does not run a full audio session — success is WS handshake, optionally +// enriched with the first server event type when one arrives quickly. +func (s *AccountTestService) testGrokRealtime(c *gin.Context, ctx context.Context, account *Account, authToken, modelID string) error { + model := strings.TrimSpace(modelID) + if model == "" { + model = defaultGrokRealtimeTestModel + } + if mapped := strings.TrimSpace(account.GetMappedModel(model)); mapped != "" { + model = mapped + } + + base, err := buildGrokVoiceURL(account, s.cfg, "realtime") + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok Realtime URL: %s", err.Error())) + } + u, err := url.Parse(base) + if err != nil { + return s.sendErrorAndEnd(c, fmt.Sprintf("Invalid Grok Realtime URL: %s", err.Error())) + } + switch strings.ToLower(u.Scheme) { + case "https": + u.Scheme = "wss" + case "http": + u.Scheme = "ws" + case "wss", "ws": + // already websocket + default: + return s.sendErrorAndEnd(c, "Invalid Grok Realtime URL scheme") + } + q := u.Query() + if q.Get("model") == "" { + q.Set("model", model) + } + u.RawQuery = q.Encode() + wsURL := u.String() + + s.prepareGrokTestSSE(c) + s.sendEvent(c, TestEvent{Type: "test_start", Model: model}) + s.sendEvent(c, TestEvent{Type: "status", Text: "Dialing standalone wss /v1/realtime (connectivity probe)..."}) + s.sendEvent(c, TestEvent{Type: "content", Text: fmt.Sprintf("realtime target: %s\n", redactGrokRealtimeURLForLog(wsURL))}) + + headers := http.Header{} + headers.Set("Authorization", "Bearer "+authToken) + if account.IsGrokOAuth() { + applyGrokCLIHeaders(headers) + } + account.ApplyHeaderOverrides(headers) + + dialer := s.grokWSDialer + if dialer == nil { + dialer = newDefaultOpenAIWSClientDialer() + } + + dialCtx, cancel := context.WithTimeout(ctx, grokRealtimeProbeTimeout) + defer cancel() + + conn, status, _, dialErr := dialer.Dial(dialCtx, wsURL, headers, s.grokTestProxyURL(account)) + if dialErr != nil { + detail := dialErr.Error() + var hs *openAIWSHandshakeError + if errors.As(dialErr, &hs) && len(hs.Body) > 0 { + body := strings.TrimSpace(string(hs.Body)) + if len(body) > 300 { + body = body[:300] + "..." + } + detail = fmt.Sprintf("%s body=%s", detail, body) + } + if status > 0 { + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok Realtime WS handshake failed (HTTP %d): %s", status, detail)) + } + return s.sendErrorAndEnd(c, fmt.Sprintf("Grok Realtime WS dial failed: %s", detail)) + } + defer func() { _ = conn.Close() }() + + s.sendEvent(c, TestEvent{Type: "content", Text: "realtime ws handshake ok\n"}) + + // Best-effort: read one server event if it arrives quickly (session.created etc.). + // Handshake alone is enough for connectivity; missing first event is not a failure. + readCtx, readCancel := context.WithTimeout(ctx, 3*time.Second) + defer readCancel() + if msg, readErr := conn.ReadMessage(readCtx); readErr == nil && len(msg) > 0 { + eventType := strings.TrimSpace(gjson.GetBytes(msg, "type").String()) + if eventType == "" { + eventType = "unknown" + } + preview := strings.TrimSpace(string(msg)) + if len(preview) > 240 { + preview = preview[:240] + "..." + } + s.sendEvent(c, TestEvent{ + Type: "content", + Text: fmt.Sprintf("realtime first event: type=%s payload=%s\n", eventType, preview), + }) + } else { + s.sendEvent(c, TestEvent{ + Type: "content", + Text: "realtime handshake succeeded (no server event within 3s; still connectivity OK)\n", + }) + } + + s.sendEvent(c, TestEvent{Type: "test_complete", Success: true}) + return nil +} + +// validateAccountTestDataURL ensures data URLs are well-formed and size-bounded. +func validateAccountTestDataURL(raw, requiredPrefix string) error { + raw = strings.TrimSpace(raw) + if raw == "" { + return fmt.Errorf("media data URL is empty") + } + if !strings.HasPrefix(raw, "data:") { + return fmt.Errorf("media must be a data: URL (data:;base64,...)") + } + // Rough size check before decode (base64 expands ~4/3). + if len(raw) > maxAccountTestMediaBytes*2 { + return fmt.Errorf("media data URL exceeds size limit") + } + _, mime, err := decodeAccountTestDataURL(raw) + if err != nil { + return err + } + if requiredPrefix != "" && !strings.HasPrefix(strings.ToLower(mime), strings.ToLower(requiredPrefix)) { + return fmt.Errorf("expected media type prefix %q, got %q", requiredPrefix, mime) + } + return nil +} + +// normalizeAccountTestImageDataURL validates an image data URL, enforces xAI +// minimum dimensions (8x8), and rewrites to a clean data:image/;base64,... form. +func normalizeAccountTestImageDataURL(raw string) (string, error) { + if err := validateAccountTestDataURL(raw, "image/"); err != nil { + return "", err + } + data, mime, err := decodeAccountTestDataURL(raw) + if err != nil { + return "", err + } + // Soft cap decoded bytes (~4 MiB) for edit payloads to avoid upstream/proxy EOF. + const maxDecodedImage = 4 << 20 + if len(data) > maxDecodedImage { + return "", fmt.Errorf( + "source image is too large (%d bytes decoded). Please use a smaller image (under ~4 MB) for admin edit tests", + len(data), + ) + } + cfg, _, err := image.DecodeConfig(bytes.NewReader(data)) + if err != nil { + // Keep raw data URL if decoder does not understand the codec (e.g. webp + // without golang.org/x/image/webp); still send upstream and let xAI validate. + return "data:" + mime + ";base64," + base64.StdEncoding.EncodeToString(data), nil + } + if cfg.Width < 8 || cfg.Height < 8 { + return "", fmt.Errorf( + "source image is too small (%dx%d). xAI requires both width and height to be at least 8 pixels", + cfg.Width, cfg.Height, + ) + } + // Prefer a stable mime from config when known. + if mime == "" || mime == "application/octet-stream" { + mime = "image/png" + } + return "data:" + mime + ";base64," + base64.StdEncoding.EncodeToString(data), nil +} + +func isTransientGrokTransportError(err error) bool { + if err == nil { + return false + } + msg := strings.ToLower(err.Error()) + return strings.Contains(msg, "unexpected eof") || + strings.Contains(msg, "connection reset") || + strings.Contains(msg, "broken pipe") || + strings.Contains(msg, "i/o timeout") || + strings.Contains(msg, "timeout awaiting response") +} + +func formatGrokImageTransportError(err error, hasSourceImage bool, payloadBytes int) string { + base := fmt.Sprintf("Grok image request failed: %s", err.Error()) + if !hasSourceImage { + return base + } + return base + fmt.Sprintf( + " (edit payload ~%d bytes). Tips: use a smaller source image (<4 MB / lower resolution), ensure the account proxy is stable, and retry. xAI /images/edits expects image as {\"url\":\"data:image/...;base64,...\",\"type\":\"image_url\"}.", + payloadBytes, + ) +} + +func formatGrokImagesAPIError(status int, body []byte, hasSourceImage bool) string { + msg := strings.TrimSpace(string(body)) + if len(msg) > 800 { + msg = msg[:800] + "..." + } + prefix := fmt.Sprintf("Grok images API returned %d: %s", status, msg) + lower := strings.ToLower(msg) + if hasSourceImage && (strings.Contains(lower, "too small") || strings.Contains(lower, "at least 8")) { + return prefix + " — upload a source image with both width and height ≥ 8 px." + } + return prefix +} + +func decodeAccountTestDataURL(raw string) (data []byte, mime string, err error) { + raw = strings.TrimSpace(raw) + if !strings.HasPrefix(raw, "data:") { + return nil, "", fmt.Errorf("not a data URL") + } + rest := strings.TrimPrefix(raw, "data:") + comma := strings.Index(rest, ",") + if comma < 0 { + return nil, "", fmt.Errorf("invalid data URL (missing comma)") + } + meta := rest[:comma] + payload := rest[comma+1:] + mime = "application/octet-stream" + if semi := strings.Index(meta, ";"); semi >= 0 { + if t := strings.TrimSpace(meta[:semi]); t != "" { + mime = t + } + } else if t := strings.TrimSpace(meta); t != "" { + mime = t + } + if !strings.Contains(strings.ToLower(meta), ";base64") { + return nil, "", fmt.Errorf("only base64 data URLs are supported") + } + decoded, err := base64.StdEncoding.DecodeString(payload) + if err != nil { + // Some browsers emit URL-safe base64 without padding. + decoded, err = base64.RawStdEncoding.DecodeString(strings.TrimRight(payload, "=")) + if err != nil { + return nil, "", fmt.Errorf("base64 decode failed: %w", err) + } + } + if len(decoded) == 0 { + return nil, "", fmt.Errorf("decoded media is empty") + } + if len(decoded) > maxAccountTestMediaBytes { + return nil, "", fmt.Errorf("media exceeds %d byte limit", maxAccountTestMediaBytes) + } + return decoded, mime, nil +} + +func sttFilenameForMIME(mime string) string { + switch strings.ToLower(strings.TrimSpace(mime)) { + case "audio/mpeg", "audio/mp3": + return "upload.mp3" + case "audio/wav", "audio/x-wav", "audio/wave": + return "upload.wav" + case "audio/webm": + return "upload.webm" + case "audio/ogg", "audio/opus": + return "upload.ogg" + case "audio/mp4", "audio/m4a", "audio/x-m4a": + return "upload.m4a" + default: + return "upload.bin" + } +} + +// redactGrokRealtimeURLForLog strips query secrets while keeping model for diagnostics. +func redactGrokRealtimeURLForLog(raw string) string { + u, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || u == nil { + return raw + } + // Keep model query only. + model := u.Query().Get("model") + u.RawQuery = "" + if model != "" { + u.RawQuery = "model=" + url.QueryEscape(model) + } + // Never log bearer in fragment/userinfo. + u.User = nil + u.Fragment = "" + return u.String() +} + +// minimalSilentWAV returns a valid tiny mono 8kHz 16-bit PCM WAV (~0.05s silence). +func minimalSilentWAV() []byte { + // 400 samples * 2 bytes = 800 data bytes + const sampleRate = 8000 + const numSamples = 400 + dataSize := numSamples * 2 + buf := make([]byte, 44+dataSize) + copy(buf[0:], []byte("RIFF")) + binary.LittleEndian.PutUint32(buf[4:], uint32(36+dataSize)) + copy(buf[8:], []byte("WAVE")) + copy(buf[12:], []byte("fmt ")) + binary.LittleEndian.PutUint32(buf[16:], 16) // PCM chunk size + binary.LittleEndian.PutUint16(buf[20:], 1) // PCM + binary.LittleEndian.PutUint16(buf[22:], 1) // mono + binary.LittleEndian.PutUint32(buf[24:], sampleRate) + binary.LittleEndian.PutUint32(buf[28:], sampleRate*2) // byte rate + binary.LittleEndian.PutUint16(buf[32:], 2) // block align + binary.LittleEndian.PutUint16(buf[34:], 16) // bits + copy(buf[36:], []byte("data")) + binary.LittleEndian.PutUint32(buf[40:], uint32(dataSize)) + // samples already zero (silence) + return buf +} + // testOpenAIChatCompletionsConnection tests an OpenAI-compatible APIKey account // through the raw /v1/chat/completions endpoint. func (s *AccountTestService) testOpenAIChatCompletionsConnection( diff --git a/backend/internal/service/account_test_service_grok_test.go b/backend/internal/service/account_test_service_grok_test.go index fcfe1257b6..781e29e47d 100644 --- a/backend/internal/service/account_test_service_grok_test.go +++ b/backend/internal/service/account_test_service_grok_test.go @@ -3,7 +3,13 @@ package service import ( + "bytes" "context" + "encoding/base64" + "errors" + "image" + "image/color" + "image/png" "io" "net/http" "net/http/httptest" @@ -22,6 +28,60 @@ type grokAccountTestRateLimitRepo struct { resetAt time.Time } +func TestObserveGrokTestResponseClassifiesBodyOnlyQuotaErrors(t *testing.T) { + account := &Account{ID: 1901, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + svc := &AccountTestService{accountRepo: repo} + + resp := &http.Response{ + StatusCode: http.StatusBadRequest, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"error":{"code":"subscription:free-usage-exhausted","message":"included free usage exhausted"}}`)), + } + svc.observeGrokTestResponse(context.Background(), account, resp) + require.Equal(t, 1, repo.tempUnschedCalls) + require.Equal(t, "grok free usage exhausted", repo.lastTempUnschedReason) + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Contains(t, string(body), "free-usage-exhausted") +} + +func TestObserveGrokTestResponseDoesNotQuarantineContentPolicy(t *testing.T) { + account := &Account{ID: 1902, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + svc := &AccountTestService{accountRepo: repo} + resp := &http.Response{ + StatusCode: http.StatusForbidden, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"error":{"code":"new_sensitive","message":"text is sensitive"}}`)), + } + svc.observeGrokTestResponse(context.Background(), account, resp) + require.Zero(t, repo.tempUnschedCalls) + require.Zero(t, repo.rateLimitedCalls) +} + +func TestObserveGrokTestResponseKeepsEntitlement403Cooldown(t *testing.T) { + account := &Account{ID: 1903, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + svc := &AccountTestService{accountRepo: repo} + resp := &http.Response{ + StatusCode: http.StatusForbidden, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"error":{"message":"subscription required"}}`)), + } + before := time.Now() + svc.observeGrokTestResponse(context.Background(), account, resp) + require.Equal(t, 1, repo.tempUnschedCalls) + require.Equal(t, "grok entitlement or subscription tier denied", repo.lastTempUnschedReason) + require.Greater(t, repo.lastTempUnschedUntil, before.Add(29*time.Minute)) +} + func (r *grokAccountTestRateLimitRepo) SetRateLimited(_ context.Context, _ int64, resetAt time.Time) error { r.rateLimitedCalls++ r.resetAt = resetAt @@ -199,3 +259,386 @@ func TestAccountTestService_Grok429WithoutQuotaHeadersUsesFallback(t *testing.T) require.Equal(t, 1, repo.rateLimitedCalls) require.WithinDuration(t, before.Add(grokRateLimitFallbackCooldown), repo.resetAt, time.Second) } + +func TestAccountTestService_GrokImageModelUsesImagesGenerations(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 17, Name: "grok-oauth-image", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + "base_url": "https://cli-chat-proxy.grok.com/v1", + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"data":[{"b64_json":"QUJD","mime_type":"image/jpeg"}]}`, + )), + }} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: upstream, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/17/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "grok-imagine-image", "a red apple", AccountTestModeDefault) + + require.NoError(t, err) + require.Equal(t, "https://api.x.ai/v1/images/generations", upstream.lastReq.URL.String()) + require.Equal(t, "grok-imagine-image", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, "a red apple", gjson.GetBytes(upstream.lastBody, "prompt").String()) + require.Equal(t, "b64_json", gjson.GetBytes(upstream.lastBody, "response_format").String()) + require.Contains(t, rec.Body.String(), `"type":"image"`) + require.Contains(t, rec.Body.String(), "data:image/jpeg;base64,QUJD") + require.Contains(t, rec.Body.String(), `"type":"test_complete"`) +} + +func TestAccountTestService_GrokWebSearchModeUsesResponsesWebSearchTool(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 18, Name: "grok-oauth-search", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader( + `{"id":"r1","output":[{"type":"web_search_call","id":"ws1","status":"completed"},{"type":"message","role":"assistant","content":[{"type":"output_text","text":"Grok is built by xAI."}]}]}`, + )), + }} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: upstream, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/18/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "grok-4.5", "xAI Grok", AccountTestModeGrokSearch) + + require.NoError(t, err) + require.Equal(t, "https://cli-chat-proxy.grok.com/v1/responses", upstream.lastReq.URL.String()) + // Standalone web_search wraps the query in the gateway-style prompt. + require.Contains(t, gjson.GetBytes(upstream.lastBody, "input").String(), "xAI Grok") + require.Equal(t, "web_search", gjson.GetBytes(upstream.lastBody, "tools.0.type").String()) + require.Equal(t, "web_search_call.action.sources", gjson.GetBytes(upstream.lastBody, "include.0").String()) + require.Contains(t, rec.Body.String(), "web_search ok") + require.Contains(t, rec.Body.String(), `"type":"test_complete"`) +} + +func TestAccountTestService_GrokTTSIncludesLanguage(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 19, Name: "grok-oauth-tts", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"audio/mpeg"}}, + Body: io.NopCloser(strings.NewReader("ID3fakeaudio")), + }} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: upstream, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/19/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "", "hello voice", AccountTestModeGrokTTS) + + require.NoError(t, err) + require.Equal(t, "https://api.x.ai/v1/tts", upstream.lastReq.URL.String()) + require.Equal(t, "hello voice", gjson.GetBytes(upstream.lastBody, "text").String()) + require.Equal(t, "en", gjson.GetBytes(upstream.lastBody, "language").String()) + require.Contains(t, rec.Body.String(), "tts ok") + require.Contains(t, rec.Body.String(), `"type":"audio"`) + require.Contains(t, rec.Body.String(), "data:audio/mpeg;base64,") + require.Contains(t, rec.Body.String(), `"type":"test_complete"`) +} + +func TestAccountTestService_GrokImageEditUsesUploadedImage(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 24, Name: "grok-oauth-image-edit", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"data":[{"b64_json":"QUJD","mime_type":"image/png"}]}`, + )), + }} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: upstream, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/24/test", nil) + + // 8x8 solid PNG (xAI min dimension is 8px). + src := minimalAccountTestPNGDataURL(8, 8) + err := svc.TestAccountConnection(c, account.ID, "grok-imagine-image", "edit me", AccountTestModeGrokImage, AccountTestOptions{ + ImageDataURL: src, + }) + + require.NoError(t, err) + require.Equal(t, "https://api.x.ai/v1/images/edits", upstream.lastReq.URL.String()) + require.True(t, strings.HasPrefix(gjson.GetBytes(upstream.lastBody, "image.url").String(), "data:image/png;base64,")) + require.Equal(t, "image_url", gjson.GetBytes(upstream.lastBody, "image.type").String()) + require.Equal(t, "b64_json", gjson.GetBytes(upstream.lastBody, "response_format").String()) + // concrete image model ids pass through; only bare "grok-imagine" is aliased. + require.Equal(t, "grok-imagine-image", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Contains(t, rec.Body.String(), `"type":"image"`) +} + +func TestAccountTestService_GrokImageEditRejectsTinySource(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 25, Name: "grok-oauth-image-tiny", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: &httpUpstreamRecorder{}, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/25/test", nil) + + // 1x1 PNG data URL + tiny := "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg==" + err := svc.TestAccountConnection(c, account.ID, "grok-imagine-image-quality", "edit", AccountTestModeGrokImage, AccountTestOptions{ + ImageDataURL: tiny, + }) + require.Error(t, err) + require.Contains(t, rec.Body.String(), "too small") +} + +// minimalAccountTestPNGDataURL builds a solid RGBA PNG as a data URL for tests. +func minimalAccountTestPNGDataURL(w, h int) string { + img := image.NewRGBA(image.Rect(0, 0, w, h)) + for y := 0; y < h; y++ { + for x := 0; x < w; x++ { + img.Set(x, y, color.RGBA{R: 200, G: 40, B: 40, A: 255}) + } + } + var buf bytes.Buffer + _ = png.Encode(&buf, img) + return "data:image/png;base64," + base64.StdEncoding.EncodeToString(buf.Bytes()) +} + +func TestAccountTestService_GrokExplicitImageModeDefaultsModel(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 20, Name: "grok-oauth-image-mode", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"data":[{"b64_json":"QUJD","mime_type":"image/jpeg"}]}`, + )), + }} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: upstream, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/20/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "", "", AccountTestModeGrokImage) + + require.NoError(t, err) + require.Equal(t, "https://api.x.ai/v1/images/generations", upstream.lastReq.URL.String()) + require.Equal(t, "grok-imagine-image", gjson.GetBytes(upstream.lastBody, "model").String()) + require.Equal(t, "b64_json", gjson.GetBytes(upstream.lastBody, "response_format").String()) + require.Contains(t, rec.Body.String(), `"type":"test_complete"`) +} + +func TestAccountTestService_GrokVideoUpstreamErrorIsNotMaskedAsSuccess(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 21, Name: "grok-oauth-video-err", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusBadRequest, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"code":"invalid-argument","error":"bad video request"}`, + )), + }} + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + httpUpstream: upstream, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/21/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "grok-imagine-video", "bounce ball", AccountTestModeGrokVideo) + + require.Error(t, err) + require.Equal(t, "https://api.x.ai/v1/videos/generations", upstream.lastReq.URL.String()) + require.Contains(t, rec.Body.String(), `"type":"error"`) + require.Contains(t, rec.Body.String(), "Grok videos API returned 400") + require.NotContains(t, rec.Body.String(), `"success":true`) +} + +type grokRealtimeTestConn struct { + msg []byte +} + +func (c *grokRealtimeTestConn) WriteJSON(context.Context, any) error { return nil } +func (c *grokRealtimeTestConn) ReadMessage(context.Context) ([]byte, error) { + if c == nil || len(c.msg) == 0 { + return nil, context.DeadlineExceeded + } + return c.msg, nil +} +func (c *grokRealtimeTestConn) Ping(context.Context) error { return nil } +func (c *grokRealtimeTestConn) Close() error { return nil } + +type grokRealtimeTestDialer struct { + lastURL string + lastAuth string + lastProxy string + conn openAIWSClientConn + err error + status int +} + +func (d *grokRealtimeTestDialer) Dial(_ context.Context, wsURL string, headers http.Header, proxyURL string) (openAIWSClientConn, int, http.Header, error) { + d.lastURL = wsURL + d.lastAuth = headers.Get("Authorization") + d.lastProxy = proxyURL + if d.err != nil { + return nil, d.status, nil, d.err + } + if d.conn == nil { + d.conn = &grokRealtimeTestConn{} + } + return d.conn, 0, nil, nil +} + +func TestAccountTestService_GrokRealtimeModeDialsWS(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 22, Name: "grok-oauth-realtime", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + dialer := &grokRealtimeTestDialer{ + conn: &grokRealtimeTestConn{msg: []byte(`{"type":"session.created","session":{"id":"sess_1"}}`)}, + } + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + grokWSDialer: dialer, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/22/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "", "", AccountTestModeGrokRealtime) + + require.NoError(t, err) + require.Contains(t, dialer.lastURL, "wss://api.x.ai/v1/realtime") + require.Contains(t, dialer.lastURL, "model=grok-voice-latest") + require.Equal(t, "Bearer grok-access-token", dialer.lastAuth) + require.Contains(t, rec.Body.String(), "realtime ws handshake ok") + require.Contains(t, rec.Body.String(), "session.created") + require.Contains(t, rec.Body.String(), `"type":"test_complete"`) + require.Contains(t, rec.Body.String(), `"success":true`) +} + +func TestAccountTestService_GrokRealtimeModeDialFailure(t *testing.T) { + gin.SetMode(gin.TestMode) + account := &Account{ + ID: 23, Name: "grok-oauth-realtime-fail", Platform: PlatformGrok, + Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, + Credentials: map[string]any{ + "access_token": "grok-access-token", + "refresh_token": "grok-refresh-token", + "expires_at": time.Now().Add(2 * time.Hour).UTC().Format(time.RFC3339), + }, + } + repo := &mockAccountRepoForGemini{accountsByID: map[int64]*Account{account.ID: account}} + dialer := &grokRealtimeTestDialer{ + status: 401, + err: &openAIWSHandshakeError{Body: []byte(`{"error":"unauthorized"}`), Err: errors.New("websocket handshake failed")}, + } + svc := &AccountTestService{ + accountRepo: repo, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + grokWSDialer: dialer, + } + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/23/test", nil) + + err := svc.TestAccountConnection(c, account.ID, "", "", AccountTestModeGrokRealtime) + + require.Error(t, err) + require.Contains(t, rec.Body.String(), `"type":"error"`) + require.Contains(t, rec.Body.String(), "Realtime") +} diff --git a/backend/internal/service/account_usage_service.go b/backend/internal/service/account_usage_service.go index 7033ad26c1..084d153ff2 100644 --- a/backend/internal/service/account_usage_service.go +++ b/backend/internal/service/account_usage_service.go @@ -199,20 +199,22 @@ type UsageInfo struct { AntigravityQuota map[string]*AntigravityModelQuota `json:"antigravity_quota,omitempty"` // Grok / xAI 被动额度快照 - GrokRequestQuota *xai.QuotaWindow `json:"grok_request_quota,omitempty"` - GrokTokenQuota *xai.QuotaWindow `json:"grok_token_quota,omitempty"` - GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"` - GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"` - GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"` - GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"` - GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"` - GrokLastStatusCode int `json:"grok_last_status_code,omitempty"` - GrokFreeTokenLimit int64 `json:"grok_free_token_limit,omitempty"` - GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"` - GrokLocalUsage24h *WindowStats `json:"grok_local_usage_24h,omitempty"` - GrokLocalUsage7d *WindowStats `json:"grok_local_usage_7d,omitempty"` - GrokLocalUsageMonthly *WindowStats `json:"grok_local_usage_monthly,omitempty"` - GrokBilling *xai.BillingSummary `json:"grok_billing,omitempty"` + GrokRequestQuota *xai.QuotaWindow `json:"grok_request_quota,omitempty"` + GrokTokenQuota *xai.QuotaWindow `json:"grok_token_quota,omitempty"` + GrokRetryAfterSeconds *int `json:"grok_retry_after_seconds,omitempty"` + GrokEntitlementStatus string `json:"grok_entitlement_status,omitempty"` + GrokQuotaSnapshotState string `json:"grok_quota_snapshot_state,omitempty"` + GrokLastQuotaProbeAt string `json:"grok_last_quota_probe_at,omitempty"` + GrokLastHeadersSeenAt string `json:"grok_last_headers_seen_at,omitempty"` + GrokLastStatusCode int `json:"grok_last_status_code,omitempty"` + GrokFreeTokenLimit int64 `json:"grok_free_token_limit,omitempty"` + GrokLocalUsage *WindowStats `json:"grok_local_usage,omitempty"` + GrokLocalUsage24h *WindowStats `json:"grok_local_usage_24h,omitempty"` + GrokLocalUsage7d *WindowStats `json:"grok_local_usage_7d,omitempty"` + GrokLocalUsageMonthly *WindowStats `json:"grok_local_usage_monthly,omitempty"` + // ThirtyDay is the official monthly billing window (used/monthlyLimit %). + ThirtyDay *UsageProgress `json:"thirty_day,omitempty"` + GrokBilling *xai.BillingSummary `json:"grok_billing,omitempty"` // Antigravity 账号级信息 SubscriptionTier string `json:"subscription_tier,omitempty"` // 归一化订阅等级: FREE/PRO/ULTRA/UNKNOWN @@ -333,23 +335,28 @@ func NewAccountUsageService( } } -// GetUsage 获取账号使用量 -// OAuth账号: 调用Anthropic API获取真实数据(需要profile scope),API响应缓存10分钟,窗口统计缓存1分钟 -// Setup Token账号: 根据session_window推算5h窗口,7d数据不可用(没有profile scope) -// API Key账号: 不支持usage查询 -func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, force ...bool) (*UsageInfo, error) { - forceProbe := len(force) > 0 && force[0] +func supportsAnthropicPassiveUsage(account *Account) bool { + return account != nil && account.IsAnthropicOAuthOrSetupToken() +} - account, err := s.accountRepo.GetByID(ctx, accountID) - if err != nil { - return nil, fmt.Errorf("get account failed: %w", err) +func batchUsageErrorMessage(err error) string { + if err == nil { + return "" } + return err.Error() +} + +func (s *AccountUsageService) getUsageForAccount(ctx context.Context, account *Account, forceProbe bool) (*UsageInfo, error) { + if account == nil { + return nil, fmt.Errorf("account is required") + } + accountID := account.ID // Dedicated UI load-test accounts must remain fully interactive without ever // contacting Anthropic with synthetic credentials. Reuse the same persisted // passive snapshot that the account table loads on mount. if account.IsSyntheticUITest() && account.IsAnthropicOAuthOrSetupToken() { - return s.GetPassiveUsage(ctx, accountID) + return s.getPassiveUsageForAccount(ctx, account) } if account.Platform == PlatformOpenAI && account.Type == AccountTypeOAuth { @@ -482,6 +489,96 @@ func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, for return nil, fmt.Errorf("account type %s does not support usage query", account.Type) } +// GetUsage 获取账号使用量 +// OAuth账号: 调用Anthropic API获取真实数据(需要profile scope),API响应缓存10分钟,窗口统计缓存1分钟 +// Setup Token账号: 根据session_window推算5h窗口,7d数据不可用(没有profile scope) +// API Key账号: 不支持usage查询 +func (s *AccountUsageService) GetUsage(ctx context.Context, accountID int64, force ...bool) (*UsageInfo, error) { + forceProbe := len(force) > 0 && force[0] + + account, err := s.accountRepo.GetByID(ctx, accountID) + if err != nil { + return nil, fmt.Errorf("get account failed: %w", err) + } + + return s.getUsageForAccount(ctx, account, forceProbe) +} + +// GetUsageBatch 批量获取账号使用量。 +// Anthropic OAuth/SetupToken 统一走 passive 链路,其他账号复用现有主动查询逻辑。 +// 单个账号失败不会中断整批请求,错误会按账号返回。 +func (s *AccountUsageService) GetUsageBatch(ctx context.Context, accountIDs []int64, force bool) (map[int64]*UsageInfo, map[int64]string, error) { + uniqueIDs := make([]int64, 0, len(accountIDs)) + seen := make(map[int64]struct{}, len(accountIDs)) + for _, accountID := range accountIDs { + if accountID <= 0 { + continue + } + if _, ok := seen[accountID]; ok { + continue + } + seen[accountID] = struct{}{} + uniqueIDs = append(uniqueIDs, accountID) + } + + usageByAccount := make(map[int64]*UsageInfo, len(uniqueIDs)) + errorsByAccount := make(map[int64]string) + if len(uniqueIDs) == 0 { + return usageByAccount, errorsByAccount, nil + } + + accounts, err := s.accountRepo.GetByIDs(ctx, uniqueIDs) + if err != nil { + return nil, nil, fmt.Errorf("get accounts failed: %w", err) + } + + accountsByID := make(map[int64]*Account, len(accounts)) + for _, account := range accounts { + if account == nil { + continue + } + accountsByID[account.ID] = account + } + + var mu sync.Mutex + g, gctx := errgroup.WithContext(ctx) + g.SetLimit(6) + + for _, accountID := range uniqueIDs { + id := accountID + account := accountsByID[id] + if account == nil { + errorsByAccount[id] = ErrAccountNotFound.Error() + continue + } + + g.Go(func() error { + var usage *UsageInfo + var usageErr error + if supportsAnthropicPassiveUsage(account) { + usage, usageErr = s.getPassiveUsageForAccount(gctx, account) + } else { + usage, usageErr = s.getUsageForAccount(gctx, account, force) + } + + mu.Lock() + defer mu.Unlock() + if usageErr != nil { + errorsByAccount[id] = batchUsageErrorMessage(usageErr) + return nil + } + usageByAccount[id] = usage + return nil + }) + } + + if err := g.Wait(); err != nil { + return nil, nil, err + } + + return usageByAccount, errorsByAccount, nil +} + // GetPassiveUsage 从 Account.Extra 中的被动采样数据构建 UsageInfo,不调用外部 API。 // 仅适用于 Anthropic OAuth / SetupToken 账号。 func (s *AccountUsageService) GetPassiveUsage(ctx context.Context, accountID int64) (*UsageInfo, error) { @@ -490,7 +587,11 @@ func (s *AccountUsageService) GetPassiveUsage(ctx context.Context, accountID int return nil, fmt.Errorf("get account failed: %w", err) } - if !account.IsAnthropicOAuthOrSetupToken() { + return s.getPassiveUsageForAccount(ctx, account) +} + +func (s *AccountUsageService) getPassiveUsageForAccount(ctx context.Context, account *Account) (*UsageInfo, error) { + if !supportsAnthropicPassiveUsage(account) { return nil, fmt.Errorf("passive usage only supported for Anthropic OAuth/SetupToken accounts") } @@ -1025,6 +1126,13 @@ func (s *AccountUsageService) getGrokUsage(ctx context.Context, account *Account ctx, s.usageLogRepo, account.ID, usage.GrokBilling, time.Now().UTC(), ) } + // Attach local window stats to official 7d/30d progress bars. + if usage.SevenDay != nil && usage.GrokLocalUsage7d != nil { + usage.SevenDay.WindowStats = usage.GrokLocalUsage7d + } + if usage.ThirtyDay != nil && usage.GrokLocalUsageMonthly != nil { + usage.ThirtyDay.WindowStats = usage.GrokLocalUsageMonthly + } } enrichUsageWithAccountError(usage, account) diff --git a/backend/internal/service/account_usage_service_batch_test.go b/backend/internal/service/account_usage_service_batch_test.go new file mode 100644 index 0000000000..067bf0ddd1 --- /dev/null +++ b/backend/internal/service/account_usage_service_batch_test.go @@ -0,0 +1,191 @@ +package service + +import ( + "context" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" +) + +// Minimal UsageLogRepository stub for batch usage tests (HEAD lacks geminiUsageLogRepoStub). +type usageBatchLogRepoStub struct{} + +var _ UsageLogRepository = (*usageBatchLogRepoStub)(nil) + +func (r *usageBatchLogRepoStub) Create(context.Context, *UsageLog) (bool, error) { + return false, nil +} +func (r *usageBatchLogRepoStub) GetByID(context.Context, int64) (*UsageLog, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) Delete(context.Context, int64) error { return nil } +func (r *usageBatchLogRepoStub) ListByUser(context.Context, int64, pagination.PaginationParams) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByAPIKey(context.Context, int64, pagination.PaginationParams) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByAccount(context.Context, int64, pagination.PaginationParams) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByUserAndTimeRange(context.Context, int64, time.Time, time.Time) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByAPIKeyAndTimeRange(context.Context, int64, time.Time, time.Time) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByAccountAndTimeRange(context.Context, int64, time.Time, time.Time) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) ListByModelAndTimeRange(context.Context, string, time.Time, time.Time) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) GetAccountWindowStats(context.Context, int64, time.Time) (*usagestats.AccountStats, error) { + return &usagestats.AccountStats{}, nil +} +func (r *usageBatchLogRepoStub) GetAccountTodayStats(context.Context, int64) (*usagestats.AccountStats, error) { + return &usagestats.AccountStats{}, nil +} +func (r *usageBatchLogRepoStub) GetDashboardStats(context.Context) (*usagestats.DashboardStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUsageTrendWithFilters(context.Context, time.Time, time.Time, string, int64, int64, int64, int64, string, *int16, *bool, *int8) ([]usagestats.TrendDataPoint, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetModelStatsWithFilters(context.Context, time.Time, time.Time, int64, int64, int64, int64, *int16, *bool, *int8) ([]usagestats.ModelStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetEndpointStatsWithFilters(context.Context, time.Time, time.Time, int64, int64, int64, int64, string, *int16, *bool, *int8) ([]usagestats.EndpointStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUpstreamEndpointStatsWithFilters(context.Context, time.Time, time.Time, int64, int64, int64, int64, string, *int16, *bool, *int8) ([]usagestats.EndpointStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetGroupStatsWithFilters(context.Context, time.Time, time.Time, int64, int64, int64, int64, *int16, *bool, *int8) ([]usagestats.GroupStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserBreakdownStats(context.Context, time.Time, time.Time, usagestats.UserBreakdownDimension, int) ([]usagestats.UserBreakdownItem, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAllGroupUsageSummary(context.Context, time.Time) ([]usagestats.GroupUsageSummary, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAPIKeyUsageTrend(context.Context, time.Time, time.Time, string, int) ([]usagestats.APIKeyUsageTrendPoint, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserUsageTrend(context.Context, time.Time, time.Time, string, int) ([]usagestats.UserUsageTrendPoint, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserSpendingRanking(context.Context, time.Time, time.Time, int) (*usagestats.UserSpendingRankingResponse, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetBatchUserUsageStats(context.Context, []int64, time.Time, time.Time) (map[int64]*usagestats.BatchUserUsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetBatchAPIKeyUsageStats(context.Context, []int64, time.Time, time.Time) (map[int64]*usagestats.BatchAPIKeyUsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserDashboardStats(context.Context, int64) (*usagestats.UserDashboardStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAPIKeyDashboardStats(context.Context, int64) (*usagestats.UserDashboardStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserUsageTrendByUserID(context.Context, int64, time.Time, time.Time, string) ([]usagestats.TrendDataPoint, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserModelStats(context.Context, int64, time.Time, time.Time) ([]usagestats.ModelStat, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) ListWithFilters(context.Context, pagination.PaginationParams, usagestats.UsageLogFilters) ([]UsageLog, *pagination.PaginationResult, error) { + return nil, nil, nil +} +func (r *usageBatchLogRepoStub) GetGlobalStats(context.Context, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetStatsWithFilters(context.Context, usagestats.UsageLogFilters) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAccountUsageStats(context.Context, int64, time.Time, time.Time) (*usagestats.AccountUsageStatsResponse, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetUserStatsAggregated(context.Context, int64, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAPIKeyStatsAggregated(context.Context, int64, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetAccountStatsAggregated(context.Context, int64, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetModelStatsAggregated(context.Context, string, time.Time, time.Time) (*usagestats.UsageStats, error) { + return nil, nil +} +func (r *usageBatchLogRepoStub) GetDailyStatsAggregated(context.Context, int64, time.Time, time.Time) ([]map[string]any, error) { + return nil, nil +} + +func TestAccountUsageService_GetUsageBatch_BestEffortByAccount(t *testing.T) { + t.Parallel() + + resetAt := time.Now().Add(2 * time.Hour).UTC().Truncate(time.Second) + + repo := &stubOpenAIAccountRepo{ + accounts: []Account{ + { + ID: 7001, + Platform: PlatformAnthropic, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "passive_usage_7d_utilization": 0.62, + }, + }, + { + ID: 7002, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Extra: map[string]any{ + "codex_usage_updated_at": time.Now().UTC().Format(time.RFC3339), + "codex_5h_used_percent": 18.0, + "codex_5h_reset_at": resetAt.Format(time.RFC3339), + "codex_7d_used_percent": 34.0, + "codex_7d_reset_at": resetAt.Add(24 * time.Hour).Format(time.RFC3339), + "workspace_id": "org-test", + "chatgpt_account_id": "acct-test", + "openai_snapshot_version": "test", + }, + }, + { + ID: 7003, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + }, + }, + } + + svc := &AccountUsageService{ + accountRepo: repo, + usageLogRepo: &usageBatchLogRepoStub{}, + cache: NewUsageCache(), + } + + usageByAccount, errorsByAccount, err := svc.GetUsageBatch(context.Background(), []int64{7001, 7002, 7003, 7002}, false) + if err != nil { + t.Fatalf("GetUsageBatch() error = %v", err) + } + + if usageByAccount[7001] == nil || usageByAccount[7001].Source != "passive" { + t.Fatalf("expected anthropic passive usage, got %#v", usageByAccount[7001]) + } + + if usageByAccount[7002] == nil || usageByAccount[7002].FiveHour == nil || usageByAccount[7002].FiveHour.Utilization != 18.0 { + t.Fatalf("expected openai snapshot usage, got %#v", usageByAccount[7002]) + } + + if !strings.Contains(strings.ToLower(errorsByAccount[7003]), "does not support usage query") { + t.Fatalf("expected API key account error to be preserved, got %q", errorsByAccount[7003]) + } +} diff --git a/backend/internal/service/account_wildcard_test.go b/backend/internal/service/account_wildcard_test.go index 6ce804bb16..15c5a66fa1 100644 --- a/backend/internal/service/account_wildcard_test.go +++ b/backend/internal/service/account_wildcard_test.go @@ -6,8 +6,31 @@ import ( "testing" "github.com/Wei-Shaw/sub2api/internal/domain" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) +func TestGrokAccountModelMappingCacheInvalidatesWithRuntimeSettings(t *testing.T) { + original := xai.RuntimeModelMappingOptions() + t.Cleanup(func() { xai.SetRuntimeModelMappingOptions(original) }) + account := &Account{Platform: PlatformGrok, Credentials: map[string]any{}} + + xai.SetRuntimeModelMappingOptions(xai.ModelMappingOptions{}) + requireMappedModel(t, account, "claude-sonnet-4-5", "claude-sonnet-4-5") + + xai.SetRuntimeModelMappingOptions(xai.ModelMappingOptions{ + DefaultText: "grok-build-0.1", + EnableCrossClientMap: true, + }) + requireMappedModel(t, account, "claude-sonnet-4-5", "grok-build-0.1") +} + +func requireMappedModel(t *testing.T, account *Account, requested, expected string) { + t.Helper() + if actual := account.GetMappedModel(requested); actual != expected { + t.Fatalf("GetMappedModel(%q) = %q, want %q", requested, actual, expected) + } +} + func TestMatchWildcard(t *testing.T) { tests := []struct { name string diff --git a/backend/internal/service/admin_account.go b/backend/internal/service/admin_account.go index 044a7ae87b..9958846658 100644 --- a/backend/internal/service/admin_account.go +++ b/backend/internal/service/admin_account.go @@ -394,63 +394,7 @@ func normalizeOpenAILongContextBillingUpdateExtra(account *Account, input *Updat return normalized, nil } -// ValidateGrokMediaEligibilityExtra validates the optional media-routing -// override. null removes the override and returns the account to automatic -// provider-observation based routing. -func ValidateGrokMediaEligibilityExtra(platform string, extra map[string]any) error { - if platform != PlatformGrok || extra == nil { - return nil - } - raw, exists := extra[GrokMediaEligibleExtraKey] - if !exists || raw == nil { - return nil - } - if _, ok := raw.(bool); !ok { - return infraerrors.BadRequest( - "GROK_MEDIA_ELIGIBILITY_INVALID", - "grok_media_eligible must be a boolean or null", - ) - } - return nil -} - -func normalizeGrokMediaEligibilityExtra(platform string, extra map[string]any) (map[string]any, error) { - if platform != PlatformGrok { - return extra, nil - } - if err := ValidateGrokMediaEligibilityExtra(platform, extra); err != nil { - return nil, err - } - normalized := maps.Clone(extra) - if normalized != nil && normalized[GrokMediaEligibleExtraKey] == nil { - delete(normalized, GrokMediaEligibleExtraKey) - } - return normalized, nil -} - -func normalizeGrokMediaEligibilityUpdateExtra(account *Account, input *UpdateAccountInput, normalized map[string]any) (map[string]any, error) { - if account == nil || account.Platform != PlatformGrok { - return normalized, nil - } - if err := ValidateGrokMediaEligibilityExtra(account.Platform, input.Extra); err != nil { - return nil, err - } - normalized = maps.Clone(normalized) - if normalized == nil { - normalized = make(map[string]any) - } - raw, provided := input.Extra[GrokMediaEligibleExtraKey] - if provided { - if raw == nil { - delete(normalized, GrokMediaEligibleExtraKey) - } - return normalized, nil - } - if current, ok := account.Extra[GrokMediaEligibleExtraKey].(bool); ok { - normalized[GrokMediaEligibleExtraKey] = current - } - return normalized, nil -} +// Grok media eligibility helpers live in account_grok_media_eligibility.go. func buildAccountForCreate(input *CreateAccountInput, accountExtra map[string]any) (*Account, error) { // Probe/session state is system-managed. New accounts always start with automatic refresh disabled. @@ -551,6 +495,8 @@ func (s *adminServiceImpl) CreateAccount(ctx context.Context, input *CreateAccou if err := NormalizeHeaderOverrideCredentials(input.Credentials); err != nil { return nil, err } + // Never persist ephemeral SSO/password secrets after OAuth conversion. + input.Credentials = SanitizeStoredCredentials(input.Platform, input.Credentials) account, err := buildAccountForCreate(input, accountExtra) if err != nil { @@ -661,6 +607,8 @@ func (s *adminServiceImpl) UpdateAccount(ctx context.Context, id int64, input *U if err := NormalizeHeaderOverrideCredentials(account.Credentials); err != nil { return nil, err } + // Strip SSO/password residue that must never sit next to OAuth tokens. + account.Credentials = SanitizeStoredCredentials(account.Platform, account.Credentials) } // Extra 使用 map:需要区分“未提供(nil)”与“显式清空({})”。 // 关闭配额限制时前端会删除 quota_* 键并提交 extra:{},此时也必须落库。 @@ -1071,6 +1019,11 @@ func (s *adminServiceImpl) BulkUpdateAccounts(ctx context.Context, input *BulkUp if err := NormalizeHeaderOverrideCredentials(input.Credentials); err != nil { return nil, err } + // Bulk may mix platforms; always drop ephemeral SSO/password keys (cookie + // only when platform is known Grok — empty platform still strips password/*). + if input.Credentials != nil { + input.Credentials = SanitizeStoredCredentials("", input.Credentials) + } // Prepare bulk updates for columns and JSONB fields. repoUpdates := AccountBulkUpdate{ diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index 910f668438..5744f40cfe 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -328,6 +328,10 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn videoPrice720P := normalizePrice(input.VideoPrice720P) videoPrice1080P := normalizePrice(input.VideoPrice1080P) webSearchPricePerCall := normalizePrice(input.WebSearchPricePerCall) + searchPricePer1k := normalizePrice(input.SearchPricePer1k) + audioRealtimePricePerMin := normalizePrice(input.AudioRealtimePricePerMin) + audioTTSPricePerMillionChars := normalizePrice(input.AudioTTSPricePerMillionChars) + audioSTTPricePerHour := normalizePrice(input.AudioSTTPricePerHour) imageRateMultiplier := 1.0 if input.ImageRateMultiplier != nil { if *input.ImageRateMultiplier < 0 { @@ -476,7 +480,12 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn VideoPrice480P: videoPrice480P, VideoPrice720P: videoPrice720P, VideoPrice1080P: videoPrice1080P, + VideoModelPrices: NormalizeVideoModelPrices(input.VideoModelPrices), WebSearchPricePerCall: webSearchPricePerCall, + SearchPricePer1k: searchPricePer1k, + AudioRealtimePricePerMin: audioRealtimePricePerMin, + AudioTTSPricePerMillionChars: audioTTSPricePerMillionChars, + AudioSTTPricePerHour: audioSTTPricePerHour, ClaudeCodeOnly: input.ClaudeCodeOnly, FallbackGroupID: input.FallbackGroupID, FallbackGroupIDOnInvalidRequest: fallbackOnInvalidRequest, @@ -755,9 +764,25 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if input.VideoPrice1080P != nil { group.VideoPrice1080P = normalizePrice(input.VideoPrice1080P) } + // nil = leave unchanged; empty map = clear per-model prices. + if input.VideoModelPrices != nil { + group.VideoModelPrices = NormalizeVideoModelPrices(input.VideoModelPrices) + } if input.WebSearchPricePerCall != nil { group.WebSearchPricePerCall = normalizePrice(input.WebSearchPricePerCall) } + if input.SearchPricePer1k != nil { + group.SearchPricePer1k = normalizePrice(input.SearchPricePer1k) + } + if input.AudioRealtimePricePerMin != nil { + group.AudioRealtimePricePerMin = normalizePrice(input.AudioRealtimePricePerMin) + } + if input.AudioTTSPricePerMillionChars != nil { + group.AudioTTSPricePerMillionChars = normalizePrice(input.AudioTTSPricePerMillionChars) + } + if input.AudioSTTPricePerHour != nil { + group.AudioSTTPricePerHour = normalizePrice(input.AudioSTTPricePerHour) + } // Claude Code 客户端限制 if input.ClaudeCodeOnly != nil { diff --git a/backend/internal/service/admin_group_duplicate.go b/backend/internal/service/admin_group_duplicate.go index 841af5e5eb..deda6d3822 100644 --- a/backend/internal/service/admin_group_duplicate.go +++ b/backend/internal/service/admin_group_duplicate.go @@ -67,6 +67,21 @@ func cloneGroupModelRouting(value map[string][]int64) map[string][]int64 { return cloned } +func cloneGroupVideoModelPrices(value map[string]map[string]float64) map[string]map[string]float64 { + if value == nil { + return nil + } + cloned := make(map[string]map[string]float64, len(value)) + for model, prices := range value { + clonedPrices := make(map[string]float64, len(prices)) + for resolution, price := range prices { + clonedPrices[resolution] = price + } + cloned[model] = clonedPrices + } + return cloned +} + func cloneGroupMessagesDispatchModelConfig(value OpenAIMessagesDispatchModelConfig) OpenAIMessagesDispatchModelConfig { cloned := value if value.ExactModelMappings != nil { @@ -113,7 +128,12 @@ func cloneGroupForDuplicate(source *Group, operationID string) *Group { VideoPrice480P: cloneGroupValuePointer(source.VideoPrice480P), VideoPrice720P: cloneGroupValuePointer(source.VideoPrice720P), VideoPrice1080P: cloneGroupValuePointer(source.VideoPrice1080P), + VideoModelPrices: cloneGroupVideoModelPrices(source.VideoModelPrices), WebSearchPricePerCall: cloneGroupValuePointer(source.WebSearchPricePerCall), + SearchPricePer1k: cloneGroupValuePointer(source.SearchPricePer1k), + AudioRealtimePricePerMin: cloneGroupValuePointer(source.AudioRealtimePricePerMin), + AudioTTSPricePerMillionChars: cloneGroupValuePointer(source.AudioTTSPricePerMillionChars), + AudioSTTPricePerHour: cloneGroupValuePointer(source.AudioSTTPricePerHour), ClaudeCodeOnly: source.ClaudeCodeOnly, FallbackGroupID: cloneGroupValuePointer(source.FallbackGroupID), FallbackGroupIDOnInvalidRequest: cloneGroupValuePointer(source.FallbackGroupIDOnInvalidRequest), diff --git a/backend/internal/service/admin_group_duplicate_test.go b/backend/internal/service/admin_group_duplicate_test.go index 44db1afdc7..f4a7bef2c5 100644 --- a/backend/internal/service/admin_group_duplicate_test.go +++ b/backend/internal/service/admin_group_duplicate_test.go @@ -121,37 +121,40 @@ func groupDuplicateTestPointer[T any](value T) *T { return &value } func TestDuplicateGroupCopiesConfigurationDeeplyAndResetsRuntimeState(t *testing.T) { createdAt := time.Date(2026, time.July, 1, 2, 3, 4, 0, time.UTC) source := &Group{ - ID: 41, - Name: "高级订阅", - Description: "configuration", - Platform: PlatformOpenAI, - RateMultiplier: 1.75, - PeakRateEnabled: true, - PeakStart: "09:00", - PeakEnd: "18:00", - PeakRateMultiplier: 1.2, - IsExclusive: true, - Status: StatusActive, - Hydrated: true, - SubscriptionType: SubscriptionTypeSubscription, - DailyLimitUSD: groupDuplicateTestPointer(11.0), - WeeklyLimitUSD: groupDuplicateTestPointer(22.0), - MonthlyLimitUSD: groupDuplicateTestPointer(33.0), - DefaultValidityDays: 91, - AllowImageGeneration: true, - AllowBatchImageGeneration: true, - ImageRateIndependent: true, - ImageRateMultiplier: 1.4, - ImagePrice1K: groupDuplicateTestPointer(0.01), - ImagePrice2K: groupDuplicateTestPointer(0.02), - ImagePrice4K: groupDuplicateTestPointer(0.04), - BatchImageDiscountMultiplier: 0.4, - BatchImageHoldMultiplier: 0.7, - VideoRateIndependent: true, - VideoRateMultiplier: 2.1, - VideoPrice480P: groupDuplicateTestPointer(0.1), - VideoPrice720P: groupDuplicateTestPointer(0.2), - VideoPrice1080P: groupDuplicateTestPointer(0.3), + ID: 41, + Name: "高级订阅", + Description: "configuration", + Platform: PlatformOpenAI, + RateMultiplier: 1.75, + PeakRateEnabled: true, + PeakStart: "09:00", + PeakEnd: "18:00", + PeakRateMultiplier: 1.2, + IsExclusive: true, + Status: StatusActive, + Hydrated: true, + SubscriptionType: SubscriptionTypeSubscription, + DailyLimitUSD: groupDuplicateTestPointer(11.0), + WeeklyLimitUSD: groupDuplicateTestPointer(22.0), + MonthlyLimitUSD: groupDuplicateTestPointer(33.0), + DefaultValidityDays: 91, + AllowImageGeneration: true, + AllowBatchImageGeneration: true, + ImageRateIndependent: true, + ImageRateMultiplier: 1.4, + ImagePrice1K: groupDuplicateTestPointer(0.01), + ImagePrice2K: groupDuplicateTestPointer(0.02), + ImagePrice4K: groupDuplicateTestPointer(0.04), + BatchImageDiscountMultiplier: 0.4, + BatchImageHoldMultiplier: 0.7, + VideoRateIndependent: true, + VideoRateMultiplier: 2.1, + VideoPrice480P: groupDuplicateTestPointer(0.1), + VideoPrice720P: groupDuplicateTestPointer(0.2), + VideoPrice1080P: groupDuplicateTestPointer(0.3), + VideoModelPrices: map[string]map[string]float64{ + VideoPriceFamilyGrokImagineVideo15: {VideoBillingResolution720P: 0.14}, + }, WebSearchPricePerCall: groupDuplicateTestPointer(0.005), ClaudeCodeOnly: true, FallbackGroupID: groupDuplicateTestPointer(int64(7)), @@ -204,6 +207,7 @@ func TestDuplicateGroupCopiesConfigurationDeeplyAndResetsRuntimeState(t *testing require.Equal(t, source.PeakRateMultiplier, duplicate.PeakRateMultiplier) require.Equal(t, source.DefaultValidityDays, duplicate.DefaultValidityDays) require.Equal(t, source.ImagePrice4K, duplicate.ImagePrice4K) + require.Equal(t, source.VideoModelPrices, duplicate.VideoModelPrices) require.Equal(t, source.WebSearchPricePerCall, duplicate.WebSearchPricePerCall) require.Equal(t, source.FallbackGroupID, duplicate.FallbackGroupID) require.Equal(t, source.ModelRouting, duplicate.ModelRouting) @@ -222,12 +226,14 @@ func TestDuplicateGroupCopiesConfigurationDeeplyAndResetsRuntimeState(t *testing }, repo.createdBindings[duplicate.ID]) duplicate.ModelRouting["gpt-*"][0] = 999 + duplicate.VideoModelPrices[VideoPriceFamilyGrokImagineVideo15][VideoBillingResolution720P] = 999 duplicate.SupportedModelScopes[0] = "changed" duplicate.MessagesDispatchModelConfig.ExactModelMappings["claude-special"] = "changed" duplicate.ModelsListConfig.Models[0] = "changed" duplicate.ReasoningEffortMappings[0].To = "changed" *duplicate.DailyLimitUSD = 999 require.Equal(t, int64(13), source.ModelRouting["gpt-*"][0]) + require.Equal(t, 0.14, source.VideoModelPrices[VideoPriceFamilyGrokImagineVideo15][VideoBillingResolution720P]) require.Equal(t, "claude", source.SupportedModelScopes[0]) require.Equal(t, "gpt-special", source.MessagesDispatchModelConfig.ExactModelMappings["claude-special"]) require.Equal(t, "gpt-5.4", source.ModelsListConfig.Models[0]) diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index a620f86f14..d9c4060c9e 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -238,10 +238,18 @@ type CreateGroupInput struct { VideoPrice480P *float64 VideoPrice720P *float64 VideoPrice1080P *float64 + // VideoModelPrices 可选按模型族×分辨率覆盖视频每秒单价。 + VideoModelPrices map[string]map[string]float64 // Codex alpha/search 网页搜索单次价格(USD/次,仅 openai 平台使用);nil/负数按默认价 0.01 处理 WebSearchPricePerCall *float64 - ClaudeCodeOnly bool // 仅允许 Claude Code 客户端 - FallbackGroupID *int64 // 降级分组 ID + // 搜索工具单价 per 1k + SearchPricePer1k *float64 + // Grok Voice 显式定价(分组级) + AudioRealtimePricePerMin *float64 + AudioTTSPricePerMillionChars *float64 + AudioSTTPricePerHour *float64 + ClaudeCodeOnly bool // 仅允许 Claude Code 客户端 + FallbackGroupID *int64 // 降级分组 ID // 无效请求兜底分组 ID(仅 anthropic 平台使用) FallbackGroupIDOnInvalidRequest *int64 // 模型路由配置(仅 anthropic 平台使用) @@ -303,10 +311,18 @@ type UpdateGroupInput struct { VideoPrice480P *float64 VideoPrice720P *float64 VideoPrice1080P *float64 + // VideoModelPrices 可选按模型族×分辨率覆盖;nil 表示不修改,空 map 表示清除。 + VideoModelPrices map[string]map[string]float64 // Codex alpha/search 网页搜索单次价格(USD/次);nil 表示不修改,负数表示清除回默认价 0.01 WebSearchPricePerCall *float64 - ClaudeCodeOnly *bool // 仅允许 Claude Code 客户端 - FallbackGroupID *int64 // 降级分组 ID + // 搜索工具单价;nil 不修改,负数清除 + SearchPricePer1k *float64 + // Grok Voice 显式定价;nil 表示不修改,负数表示清除 + AudioRealtimePricePerMin *float64 + AudioTTSPricePerMillionChars *float64 + AudioSTTPricePerHour *float64 + ClaudeCodeOnly *bool // 仅允许 Claude Code 客户端 + FallbackGroupID *int64 // 降级分组 ID // 无效请求兜底分组 ID(仅 anthropic 平台使用) FallbackGroupIDOnInvalidRequest *int64 // 模型路由配置(仅 anthropic 平台使用) diff --git a/backend/internal/service/antigravity_gateway_claude.go b/backend/internal/service/antigravity_gateway_claude.go index 2011625984..3511f60288 100644 --- a/backend/internal/service/antigravity_gateway_claude.go +++ b/backend/internal/service/antigravity_gateway_claude.go @@ -28,6 +28,7 @@ import ( // ├─ 成功 → 正常返回 // └─ 失败 → 设置模型限流 + 清除粘性绑定 → 切换账号 func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context, account *Account, body []byte, isStickySession bool) (*ForwardResult, error) { + beginUpstreamResponseModelObservation(c) // 上游透传账号直接转发,不走 OAuth token 刷新 if account.Type == AccountTypeUpstream { return s.ForwardUpstream(ctx, c, account, body) @@ -448,14 +449,16 @@ func (s *AntigravityGatewayService) Forward(ctx context.Context, c *gin.Context, } return &ForwardResult{ - RequestID: requestID, - Usage: *usage, - Model: originalModel, - UpstreamModel: billingModel, - Stream: claudeReq.Stream, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, - ClientDisconnect: clientDisconnect, + RequestID: requestID, + Usage: *usage, + Model: originalModel, + UpstreamModel: billingModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: claudeReq.Stream, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnect, }, nil } diff --git a/backend/internal/service/antigravity_gateway_compat.go b/backend/internal/service/antigravity_gateway_compat.go index af1b8ebbe4..ab8f0d17a2 100644 --- a/backend/internal/service/antigravity_gateway_compat.go +++ b/backend/internal/service/antigravity_gateway_compat.go @@ -168,6 +168,7 @@ func (s *AntigravityGatewayService) forwardAntigravityCompat( account *Account, request antigravityCompatRequest, ) (*ForwardResult, error) { + beginUpstreamResponseModelObservation(c) call, err := s.prepareAntigravityCompatCall(ctx, c, account, request) if err != nil { return nil, err @@ -324,15 +325,17 @@ func (s *AntigravityGatewayService) consumeAntigravityCompatResponse( } return &ForwardResult{ - RequestID: requestID, - Usage: *streamResult.usage, - Model: call.request.originalModel, - UpstreamModel: call.billingModel, - Stream: call.request.clientStream, - Duration: time.Since(call.request.startTime), - FirstTokenMs: streamResult.firstTokenMs, - ReasoningEffort: call.request.reasoningEffort, - ClientDisconnect: streamResult.clientDisconnect, + RequestID: requestID, + Usage: *streamResult.usage, + Model: call.request.originalModel, + UpstreamModel: call.billingModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: call.request.clientStream, + Duration: time.Since(call.request.startTime), + FirstTokenMs: streamResult.firstTokenMs, + ReasoningEffort: call.request.reasoningEffort, + ClientDisconnect: streamResult.clientDisconnect, }, nil } @@ -477,7 +480,7 @@ func (s *AntigravityGatewayService) handleChatCompletionsNonStreamingFromAntigra startTime time.Time, originalModel string, ) (*antigravityStreamResult, error) { - claudeResponse, result, err := s.collectClaudeStreamResponse(resp, startTime, originalModel) + claudeResponse, result, err := s.collectClaudeStreamResponse(c, resp, startTime, originalModel) if err != nil { return nil, s.mapAntigravityCompatCollectionError(c, err) } @@ -496,7 +499,7 @@ func (s *AntigravityGatewayService) handleResponsesNonStreamingFromAntigravity( startTime time.Time, originalModel string, ) (*antigravityStreamResult, error) { - claudeResponse, result, err := s.collectClaudeStreamResponse(resp, startTime, originalModel) + claudeResponse, result, err := s.collectClaudeStreamResponse(c, resp, startTime, originalModel) if err != nil { return nil, s.mapAntigravityCompatCollectionError(c, err) } diff --git a/backend/internal/service/antigravity_gateway_compat_stream.go b/backend/internal/service/antigravity_gateway_compat_stream.go index 3243193ceb..26d30ddb46 100644 --- a/backend/internal/service/antigravity_gateway_compat_stream.go +++ b/backend/internal/service/antigravity_gateway_compat_stream.go @@ -301,6 +301,7 @@ func (s *AntigravityGatewayService) handleAntigravityCompatStream( return s.handleAntigravityCompatReadError(c, session, event.err, maxLineSize, prefix) } resetAntigravityCompatTimer(timeoutTimer, timeout) + s.observeAntigravityGeminiSSELine(c, event.line) session.consume(event.line) case <-timeoutCh: diff --git a/backend/internal/service/antigravity_gateway_gemini.go b/backend/internal/service/antigravity_gateway_gemini.go index f0d20244a8..2b622afd37 100644 --- a/backend/internal/service/antigravity_gateway_gemini.go +++ b/backend/internal/service/antigravity_gateway_gemini.go @@ -42,6 +42,7 @@ func WithForwardGeminiSession(groupID int64, sessionHash string) ForwardGeminiOp } func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Context, account *Account, originalModel string, action string, stream bool, body []byte, isStickySession bool, options ...ForwardGeminiOption) (*ForwardResult, error) { + beginUpstreamResponseModelObservation(c) startTime := time.Now() forwardOpts := forwardGeminiOptions{} for _, apply := range options { @@ -203,6 +204,7 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co if err == nil && fallbackResp.StatusCode < 400 { _ = resp.Body.Close() resp = fallbackResp + billingModel = fallbackModel } else if fallbackResp != nil { _ = fallbackResp.Body.Close() } @@ -431,17 +433,19 @@ handleSuccess: } return &ForwardResult{ - RequestID: requestID, - Usage: *usage, - Model: originalModel, - UpstreamModel: billingModel, - Stream: stream, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, - ClientDisconnect: clientDisconnect, - ImageCount: imageCount, - ImageSize: imageSize, - ImageInputSize: imageInputSize, + RequestID: requestID, + Usage: *usage, + Model: originalModel, + UpstreamModel: billingModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: stream, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnect, + ImageCount: imageCount, + ImageSize: imageSize, + ImageInputSize: imageInputSize, }, nil } diff --git a/backend/internal/service/antigravity_gateway_streaming.go b/backend/internal/service/antigravity_gateway_streaming.go index 2810162c45..0b3a672853 100644 --- a/backend/internal/service/antigravity_gateway_streaming.go +++ b/backend/internal/service/antigravity_gateway_streaming.go @@ -22,6 +22,26 @@ type antigravityStreamResult struct { clientDisconnect bool // 客户端是否在流式传输过程中断开 } +func (s *AntigravityGatewayService) observeAntigravityGeminiSSELine(c *gin.Context, line string) { + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + trimmed := strings.TrimSpace(line) + if !strings.HasPrefix(trimmed, "data:") { + return + } + payload := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) + if payload == "" || payload == "[DONE]" { + return + } + raw := []byte(payload) + if inner, err := s.unwrapV1InternalResponse(raw); err == nil && len(inner) > 0 { + raw = inner + } + observer.ObserveGemini(raw) +} + // antigravityClientWriter 封装流式响应的客户端写入,自动检测断开并标记。 // 断开后所有写入操作变为 no-op,调用方通过 Disconnected() 判断是否继续 drain 上游。 type antigravityClientWriter struct { @@ -95,6 +115,9 @@ func handleStreamReadError(err error, clientDisconnected bool, prefix string) (d } func (s *AntigravityGatewayService) handleGeminiStreamingResponse(c *gin.Context, resp *http.Response, startTime time.Time) (*antigravityStreamResult, error) { + if upstreamResponseModelObserverFromContext(c) == nil { + beginUpstreamResponseModelObservation(c) + } c.Status(resp.StatusCode) c.Header("Cache-Control", "no-cache") c.Header("Connection", "keep-alive") @@ -220,6 +243,7 @@ func (s *AntigravityGatewayService) handleGeminiStreamingResponse(c *gin.Context lastDataAt = time.Now() line := ev.line + s.observeAntigravityGeminiSSELine(c, line) trimmed := strings.TrimRight(line, "\r\n") if strings.HasPrefix(trimmed, "data:") { payload := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) @@ -298,6 +322,9 @@ func (s *AntigravityGatewayService) handleGeminiStreamingResponse(c *gin.Context // handleGeminiStreamToNonStreaming 读取上游流式响应,合并为非流式响应返回给客户端 // Gemini 流式响应是增量的,需要累积所有 chunk 的内容 func (s *AntigravityGatewayService) handleGeminiStreamToNonStreaming(c *gin.Context, resp *http.Response, startTime time.Time) (*antigravityStreamResult, error) { + if upstreamResponseModelObserverFromContext(c) == nil { + beginUpstreamResponseModelObservation(c) + } scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 { @@ -377,6 +404,7 @@ func (s *AntigravityGatewayService) handleGeminiStreamToNonStreaming(c *gin.Cont } line := ev.line + s.observeAntigravityGeminiSSELine(c, line) trimmed := strings.TrimRight(line, "\r\n") if !strings.HasPrefix(trimmed, "data:") { @@ -765,7 +793,10 @@ func (s *AntigravityGatewayService) writeGoogleError(c *gin.Context, status int, // collectClaudeStreamResponse 收集上游流式响应,转换为 Claude 非流式格式返回 // 用于处理客户端非流式请求但上游只支持流式的情况 -func (s *AntigravityGatewayService) collectClaudeStreamResponse(resp *http.Response, startTime time.Time, originalModel string) ([]byte, *antigravityStreamResult, error) { +func (s *AntigravityGatewayService) collectClaudeStreamResponse(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) ([]byte, *antigravityStreamResult, error) { + if upstreamResponseModelObserverFromContext(c) == nil { + beginUpstreamResponseModelObservation(c) + } scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.settingService.cfg != nil && s.settingService.cfg.Gateway.MaxLineSize > 0 { @@ -860,6 +891,7 @@ func (s *AntigravityGatewayService) collectClaudeStreamResponse(resp *http.Respo if parseErr != nil { continue } + upstreamResponseModelObserverFromContext(c).ObserveGemini(inner) var parsed map[string]any if err := json.Unmarshal(inner, &parsed); err != nil { @@ -941,7 +973,7 @@ returnResponse: // handleClaudeStreamToNonStreaming 收集上游流式响应,转换为 Claude 非流式格式返回 // 用于处理客户端非流式请求但上游只支持流式的情况 func (s *AntigravityGatewayService) handleClaudeStreamToNonStreaming(c *gin.Context, resp *http.Response, startTime time.Time, originalModel string) (*antigravityStreamResult, error) { - claudeResp, streamRes, err := s.collectClaudeStreamResponse(resp, startTime, originalModel) + claudeResp, streamRes, err := s.collectClaudeStreamResponse(c, resp, startTime, originalModel) if err != nil { var failoverErr *UpstreamFailoverError if errors.As(err, &failoverErr) { @@ -1120,6 +1152,7 @@ func (s *AntigravityGatewayService) handleClaudeStreamingResponse(c *gin.Context } lastDataAt = time.Now() + s.observeAntigravityGeminiSSELine(c, ev.line) // 处理 SSE 行,转换为 Claude 格式 claudeEvents := processor.ProcessLine(strings.TrimRight(ev.line, "\r\n")) diff --git a/backend/internal/service/antigravity_gateway_upstream.go b/backend/internal/service/antigravity_gateway_upstream.go index 2914863131..7254745698 100644 --- a/backend/internal/service/antigravity_gateway_upstream.go +++ b/backend/internal/service/antigravity_gateway_upstream.go @@ -19,6 +19,7 @@ import ( // ForwardUpstream 使用 base_url + /v1/messages + 双 header 认证透传上游 Claude 请求 func (s *AntigravityGatewayService) ForwardUpstream(ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + beginUpstreamResponseModelObservation(c) startTime := time.Now() sessionID := getSessionID(c) prefix := logPrefix(sessionID, account.Name) @@ -129,6 +130,7 @@ func (s *AntigravityGatewayService) ForwardUpstream(ctx context.Context, c *gin. } // 提取 usage + upstreamResponseModelObserverFromContext(c).ObserveAnthropic(respBody) usage = s.extractClaudeUsage(respBody) c.Header("Content-Type", resp.Header.Get("Content-Type")) @@ -141,11 +143,13 @@ func (s *AntigravityGatewayService) ForwardUpstream(ctx context.Context, c *gin. logger.LegacyPrintf("service.antigravity_gateway", "%s status=success duration_ms=%d", prefix, duration.Milliseconds()) return &ForwardResult{ - Model: originalModel, - Stream: claudeReq.Stream, - Duration: duration, - FirstTokenMs: firstTokenMs, - ClientDisconnect: clientDisconnect, + Model: originalModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: claudeReq.Stream, + Duration: duration, + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnect, Usage: ClaudeUsage{ InputTokens: usage.InputTokens, OutputTokens: usage.OutputTokens, @@ -247,6 +251,9 @@ func (s *AntigravityGatewayService) streamUpstreamResponse(c *gin.Context, resp lastDataAt = time.Now() line := ev.line + if data, ok := extractAnthropicSSEDataLine(line); ok { + upstreamResponseModelObserverFromContext(c).ObserveAnthropic([]byte(strings.TrimSpace(data))) + } // 记录首 token 时间 if firstTokenMs == nil && len(line) > 0 { diff --git a/backend/internal/service/api_key_auth_cache.go b/backend/internal/service/api_key_auth_cache.go index a665cdad3c..5fdc96b793 100644 --- a/backend/internal/service/api_key_auth_cache.go +++ b/backend/internal/service/api_key_auth_cache.go @@ -56,32 +56,37 @@ type APIKeyAuthUserSnapshot struct { // APIKeyAuthGroupSnapshot 分组快照 type APIKeyAuthGroupSnapshot struct { - ID int64 `json:"id"` - Name string `json:"name"` - Platform string `json:"platform"` - IsExclusive bool `json:"is_exclusive"` - Status string `json:"status"` - SubscriptionType string `json:"subscription_type"` - RateMultiplier float64 `json:"rate_multiplier"` - DailyLimitUSD *float64 `json:"daily_limit_usd,omitempty"` - WeeklyLimitUSD *float64 `json:"weekly_limit_usd,omitempty"` - MonthlyLimitUSD *float64 `json:"monthly_limit_usd,omitempty"` - AllowImageGeneration bool `json:"allow_image_generation"` - AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` - ImageRateIndependent bool `json:"image_rate_independent"` - ImageRateMultiplier float64 `json:"image_rate_multiplier"` - ImagePrice1K *float64 `json:"image_price_1k,omitempty"` - ImagePrice2K *float64 `json:"image_price_2k,omitempty"` - ImagePrice4K *float64 `json:"image_price_4k,omitempty"` - VideoRateIndependent bool `json:"video_rate_independent"` - VideoRateMultiplier float64 `json:"video_rate_multiplier"` - VideoPrice480P *float64 `json:"video_price_480p,omitempty"` - VideoPrice720P *float64 `json:"video_price_720p,omitempty"` - VideoPrice1080P *float64 `json:"video_price_1080p,omitempty"` - WebSearchPricePerCall *float64 `json:"web_search_price_per_call,omitempty"` - ClaudeCodeOnly bool `json:"claude_code_only"` - FallbackGroupID *int64 `json:"fallback_group_id,omitempty"` - FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request,omitempty"` + ID int64 `json:"id"` + Name string `json:"name"` + Platform string `json:"platform"` + IsExclusive bool `json:"is_exclusive"` + Status string `json:"status"` + SubscriptionType string `json:"subscription_type"` + RateMultiplier float64 `json:"rate_multiplier"` + DailyLimitUSD *float64 `json:"daily_limit_usd,omitempty"` + WeeklyLimitUSD *float64 `json:"weekly_limit_usd,omitempty"` + MonthlyLimitUSD *float64 `json:"monthly_limit_usd,omitempty"` + AllowImageGeneration bool `json:"allow_image_generation"` + AllowBatchImageGeneration bool `json:"allow_batch_image_generation"` + ImageRateIndependent bool `json:"image_rate_independent"` + ImageRateMultiplier float64 `json:"image_rate_multiplier"` + ImagePrice1K *float64 `json:"image_price_1k,omitempty"` + ImagePrice2K *float64 `json:"image_price_2k,omitempty"` + ImagePrice4K *float64 `json:"image_price_4k,omitempty"` + VideoRateIndependent bool `json:"video_rate_independent"` + VideoRateMultiplier float64 `json:"video_rate_multiplier"` + VideoPrice480P *float64 `json:"video_price_480p,omitempty"` + VideoPrice720P *float64 `json:"video_price_720p,omitempty"` + VideoPrice1080P *float64 `json:"video_price_1080p,omitempty"` + VideoModelPrices map[string]map[string]float64 `json:"video_model_prices,omitempty"` + WebSearchPricePerCall *float64 `json:"web_search_price_per_call,omitempty"` + SearchPricePer1k *float64 `json:"search_price_per_1k,omitempty"` + AudioRealtimePricePerMin *float64 `json:"audio_realtime_price_per_min,omitempty"` + AudioTTSPricePerMillionChars *float64 `json:"audio_tts_price_per_million_chars,omitempty"` + AudioSTTPricePerHour *float64 `json:"audio_stt_price_per_hour,omitempty"` + ClaudeCodeOnly bool `json:"claude_code_only"` + FallbackGroupID *int64 `json:"fallback_group_id,omitempty"` + FallbackGroupIDOnInvalidRequest *int64 `json:"fallback_group_id_on_invalid_request,omitempty"` // Model routing is used by gateway account selection, so it must be part of auth cache snapshot. // Only anthropic groups use these fields; others may leave them empty. diff --git a/backend/internal/service/api_key_auth_cache_impl.go b/backend/internal/service/api_key_auth_cache_impl.go index e834ce62de..696f21bb25 100644 --- a/backend/internal/service/api_key_auth_cache_impl.go +++ b/backend/internal/service/api_key_auth_cache_impl.go @@ -14,7 +14,7 @@ import ( "github.com/dgraph-io/ristretto" ) -const apiKeyAuthSnapshotVersion = 18 // v18: include group profit control fields (force refresh of pre-fix snapshots) +const apiKeyAuthSnapshotVersion = 19 // v19: group search/audio/video_model_prices billing fields (force refresh of pre-fix snapshots) type apiKeyAuthCacheConfig struct { l1Size int @@ -400,7 +400,12 @@ func (s *APIKeyService) snapshotFromAPIKey(ctx context.Context, apiKey *APIKey) VideoPrice480P: apiKey.Group.VideoPrice480P, VideoPrice720P: apiKey.Group.VideoPrice720P, VideoPrice1080P: apiKey.Group.VideoPrice1080P, + VideoModelPrices: NormalizeVideoModelPrices(apiKey.Group.VideoModelPrices), WebSearchPricePerCall: apiKey.Group.WebSearchPricePerCall, + SearchPricePer1k: apiKey.Group.SearchPricePer1k, + AudioRealtimePricePerMin: apiKey.Group.AudioRealtimePricePerMin, + AudioTTSPricePerMillionChars: apiKey.Group.AudioTTSPricePerMillionChars, + AudioSTTPricePerHour: apiKey.Group.AudioSTTPricePerHour, ClaudeCodeOnly: apiKey.Group.ClaudeCodeOnly, FallbackGroupID: apiKey.Group.FallbackGroupID, FallbackGroupIDOnInvalidRequest: apiKey.Group.FallbackGroupIDOnInvalidRequest, @@ -490,7 +495,12 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho VideoPrice480P: snapshot.Group.VideoPrice480P, VideoPrice720P: snapshot.Group.VideoPrice720P, VideoPrice1080P: snapshot.Group.VideoPrice1080P, + VideoModelPrices: NormalizeVideoModelPrices(snapshot.Group.VideoModelPrices), WebSearchPricePerCall: snapshot.Group.WebSearchPricePerCall, + SearchPricePer1k: snapshot.Group.SearchPricePer1k, + AudioRealtimePricePerMin: snapshot.Group.AudioRealtimePricePerMin, + AudioTTSPricePerMillionChars: snapshot.Group.AudioTTSPricePerMillionChars, + AudioSTTPricePerHour: snapshot.Group.AudioSTTPricePerHour, ClaudeCodeOnly: snapshot.Group.ClaudeCodeOnly, FallbackGroupID: snapshot.Group.FallbackGroupID, FallbackGroupIDOnInvalidRequest: snapshot.Group.FallbackGroupIDOnInvalidRequest, diff --git a/backend/internal/service/api_key_auth_cache_profit_test.go b/backend/internal/service/api_key_auth_cache_profit_test.go index 65bef3176f..40519f2896 100644 --- a/backend/internal/service/api_key_auth_cache_profit_test.go +++ b/backend/internal/service/api_key_auth_cache_profit_test.go @@ -53,7 +53,7 @@ func TestAPIKeyAuthSnapshotProfitControlRoundtrip(t *testing.T) { snapshot := svc.snapshotFromAPIKey(context.Background(), apiKey) require.NotNil(t, snapshot) require.Equal(t, apiKeyAuthSnapshotVersion, snapshot.Version) - require.Equal(t, 18, snapshot.Version, "v18 起认证快照携带利润控制字段") + require.Equal(t, 19, snapshot.Version, "v19 起认证快照携带 search/audio/video_model_prices 计费字段") // 模拟 L2 缓存的完整 JSON 往返(与 apiKeyCache.SetAuthCache/GetAuthCache 同构)。 payload, err := json.Marshal(&APIKeyAuthCacheEntry{Snapshot: snapshot}) diff --git a/backend/internal/service/billing_search_audio_cost_test.go b/backend/internal/service/billing_search_audio_cost_test.go new file mode 100644 index 0000000000..94d2e5c659 --- /dev/null +++ b/backend/internal/service/billing_search_audio_cost_test.go @@ -0,0 +1,42 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCalculateSearchCost(t *testing.T) { + t.Parallel() + s := &BillingService{} + require.Equal(t, 0.0, s.CalculateSearchCost(0, floatPtr(10), 1).ActualCost) + // nil price → default $10/1k: 5 calls = 0.05 + require.InDelta(t, 0.05, s.CalculateSearchCost(5, nil, 1).ActualCost, 1e-9) + // explicit 0 → free + require.Equal(t, 0.0, s.CalculateSearchCost(5, floatPtr(0), 1).ActualCost) + price := 10.0 + cost := s.CalculateSearchCost(100, &price, 1.5) + // 10 / 1000 * 100 * 1.5 = 1.5 + require.InDelta(t, 1.0, cost.TotalCost, 1e-9) + require.InDelta(t, 1.5, cost.ActualCost, 1e-9) +} + +func TestCalculateAudioCost(t *testing.T) { + t.Parallel() + s := &BillingService{} + rt, tts, stt := 0.10, 15.0, 0.50 + cfg := &audioPriceConfig{RealtimePerMin: &rt, TTSPerMChars: &tts, STTPerHour: &stt} + require.InDelta(t, 0.20, s.CalculateAudioCost("realtime", 2, cfg, 1).ActualCost, 1e-9) + require.InDelta(t, 1.5, s.CalculateAudioCost("tts", 0.1, cfg, 1).ActualCost, 1e-9) + require.InDelta(t, 0.25, s.CalculateAudioCost("stt", 0.5, cfg, 1).ActualCost, 1e-9) + require.Equal(t, 0.0, s.CalculateAudioCost("unknown", 1, cfg, 1).ActualCost) + // nil config → defaults (realtime $0.10/min, tts $15/M, stt $0.36/hr) + require.InDelta(t, 0.10, s.CalculateAudioCost("realtime", 1, nil, 1).ActualCost, 1e-9) + require.InDelta(t, 15.0, s.CalculateAudioCost("tts", 1, nil, 1).ActualCost, 1e-9) + require.InDelta(t, 0.36, s.CalculateAudioCost("stt", 1, nil, 1).ActualCost, 1e-9) + // explicit 0 → free + zero := 0.0 + require.Equal(t, 0.0, s.CalculateAudioCost("realtime", 1, &audioPriceConfig{RealtimePerMin: &zero}, 1).ActualCost) +} + +func floatPtr(v float64) *float64 { return &v } diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index a43f26a8ad..2b11fc978d 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -1415,6 +1415,9 @@ type VideoPriceConfig struct { Price480P *float64 // 480p 每秒价格(nil 表示使用默认值) Price720P *float64 // 720p 每秒价格(nil 表示使用默认值) Price1080P *float64 // 1080p 每秒价格(nil 表示使用默认值) + // ModelPrices is optional per-model-family override: family → resolution → USD/s. + // When set for a model, it wins over Price* flat columns for that model only. + ModelPrices map[string]map[string]float64 } const ( @@ -1434,6 +1437,15 @@ const ( // Codex alpha/search 网页搜索单次默认价:OpenAI 官方 web search 定价 $10/1000 次。 defaultWebSearchPricePerCall = 0.01 + + // Grok /v1/web_search 与 SearchCount 附加费:与 Codex 对齐 $10/1000 次(按 1k 计价字段存储)。 + defaultSearchPricePer1k = 10.0 + + // Grok Voice 默认价(分组列 NULL 时使用;显式配 0 表示免费)。 + // 保守运营占位,运维可通过 groups.audio_* 覆盖。 + defaultAudioRealtimePricePerMin = 0.10 + defaultAudioTTSPricePerMillionChars = 15.0 + defaultAudioSTTPricePerHour = 0.36 ) // CalculateWebSearchCost 计算 Codex alpha/search 网页搜索按次费用。 @@ -1461,6 +1473,80 @@ func (s *BillingService) CalculateWebSearchCost(callCount int, groupPrice *float } } +// CalculateSearchCost bills search/tool invocations (e.g. web_search) per 1k calls. +// groupPricePer1k: nil → defaultSearchPricePer1k; explicit 0 → free; >0 → that rate. +func (s *BillingService) CalculateSearchCost(numCalls int, groupPricePer1k *float64, rateMultiplier float64) *CostBreakdown { + if numCalls <= 0 { + return &CostBreakdown{} + } + pricePer1k := defaultSearchPricePer1k + if groupPricePer1k != nil { + if *groupPricePer1k < 0 { + return &CostBreakdown{} + } + pricePer1k = *groupPricePer1k + } + if pricePer1k == 0 { + return &CostBreakdown{} + } + if rateMultiplier < 0 { + rateMultiplier = 0 + } + unit := pricePer1k / 1000.0 + total := unit * float64(numCalls) + return &CostBreakdown{ + TotalCost: total, + ActualCost: total * rateMultiplier, + BillingMode: string(BillingModePerRequest), + } +} + +type audioPriceConfig struct { + RealtimePerMin *float64 + TTSPerMChars *float64 + STTPerHour *float64 +} + +// CalculateAudioCost supports realtime (per min), tts (per M chars), stt (per hr). +// Missing group prices use defaults; explicit 0 means free for that mode. +func (s *BillingService) CalculateAudioCost(mode string, durationOrUnits float64, groupConfig *audioPriceConfig, rateMultiplier float64) *CostBreakdown { + if durationOrUnits <= 0 { + return &CostBreakdown{} + } + var unitPrice float64 + switch strings.ToLower(mode) { + case "realtime": + unitPrice = defaultAudioRealtimePricePerMin + if groupConfig != nil && groupConfig.RealtimePerMin != nil { + unitPrice = *groupConfig.RealtimePerMin + } + case "tts": + unitPrice = defaultAudioTTSPricePerMillionChars + if groupConfig != nil && groupConfig.TTSPerMChars != nil { + unitPrice = *groupConfig.TTSPerMChars + } + case "stt": + unitPrice = defaultAudioSTTPricePerHour + if groupConfig != nil && groupConfig.STTPerHour != nil { + unitPrice = *groupConfig.STTPerHour + } + default: + return &CostBreakdown{} + } + if unitPrice <= 0 { + return &CostBreakdown{} + } + if rateMultiplier < 0 { + rateMultiplier = 0 + } + total := unitPrice * durationOrUnits + return &CostBreakdown{ + TotalCost: total, + ActualCost: total * rateMultiplier, + BillingMode: string(BillingModePerRequest), + } +} + // CalculateImageCost 计算图片生成费用 // model: 请求的模型名称(用于获取 LiteLLM 默认价格) // imageSize: 图片尺寸 "1K", "2K", "4K" @@ -1546,8 +1632,12 @@ func (s *BillingService) getImageUnitPrice(model string, imageSize string, group } func (s *BillingService) getVideoUnitPrice(model string, resolution string, groupConfig *VideoPriceConfig) float64 { + // Order: (a) per-model map (b) flat group video_price_* (c) model-aware code defaults. if groupConfig != nil { - switch resolution { + if price := LookupVideoModelPrice(groupConfig.ModelPrices, model, resolution); price != nil { + return *price + } + switch NormalizeVideoBillingResolutionOrDefault(resolution) { case VideoBillingResolution480P: if groupConfig.Price480P != nil { return *groupConfig.Price480P diff --git a/backend/internal/service/channel_plaza.go b/backend/internal/service/channel_plaza.go index a71a6e4f1e..1e4dee2717 100644 --- a/backend/internal/service/channel_plaza.go +++ b/backend/internal/service/channel_plaza.go @@ -28,7 +28,8 @@ type PlazaModel struct { // PlazaGroup 模型广场中以分组为顶层的条目。 // // 与 AvailableGroupRef 相比多了 Description 与 Models;Models 来自该分组关联渠道的 -// 支持模型(按分组平台隔离,防跨平台泄漏),与「可用渠道」页口径一致。 +// 支持模型(普通分组按分组平台隔离,Composite 分组展开关联渠道已配置的 +// 具体平台),与「可用渠道」页口径一致。 type PlazaGroup struct { ID int64 Name string @@ -99,8 +100,12 @@ func (s *ChannelService) ListPlazaGroups(ctx context.Context) ([]PlazaGroup, err order = append(order, g.ID) } - // modelIdx[groupID][modelName] = index into byGroup[groupID].Models - modelIdx := make(map[int64]map[string]int, len(groups)) + type modelKey struct { + platform string + name string + } + // modelIdx[groupID][platform+modelName] = index into byGroup[groupID].Models + modelIdx := make(map[int64]map[modelKey]int, len(groups)) for i := range channels { ch := &channels[i] if ch.Status != StatusActive { @@ -117,23 +122,28 @@ func (s *ChannelService) ListPlazaGroups(ctx context.Context) ([]PlazaGroup, err } idx := modelIdx[gid] if idx == nil { - idx = make(map[string]int, len(supported)) + idx = make(map[modelKey]int, len(supported)) modelIdx[gid] = idx } for j := range supported { m := supported[j] - if m.Platform != pg.Platform { + if pg.Platform == PlatformComposite { + if !isConcreteRequestPlatform(m.Platform) { + continue + } + } else if m.Platform != pg.Platform { continue } pricing := plazaImageDisplayPricing(m.Pricing, groupEnt[gid]) - if at, seen := idx[m.Name]; seen { + key := modelKey{platform: m.Platform, name: m.Name} + if at, seen := idx[key]; seen { // 先见者胜;仅当已存条目无定价而新条目有定价时升级。 if pg.Models[at].Pricing == nil && pricing != nil { pg.Models[at].Pricing = pricing } continue } - idx[m.Name] = len(pg.Models) + idx[key] = len(pg.Models) pg.Models = append(pg.Models, PlazaModel{ Name: m.Name, Platform: m.Platform, @@ -150,7 +160,12 @@ func (s *ChannelService) ListPlazaGroups(ctx context.Context) ([]PlazaGroup, err if len(pg.Models) == 0 { continue } - sort.SliceStable(pg.Models, func(i, j int) bool { return pg.Models[i].Name < pg.Models[j].Name }) + sort.SliceStable(pg.Models, func(i, j int) bool { + if pg.Models[i].Name != pg.Models[j].Name { + return pg.Models[i].Name < pg.Models[j].Name + } + return pg.Models[i].Platform < pg.Models[j].Platform + }) for j := range pg.Models { pg.Models[j].OfficialPricing = s.lookupOfficialPricing(pg.Models[j].Name, officialMemo) } diff --git a/backend/internal/service/channel_plaza_test.go b/backend/internal/service/channel_plaza_test.go index 82d426ce99..554654149c 100644 --- a/backend/internal/service/channel_plaza_test.go +++ b/backend/internal/service/channel_plaza_test.go @@ -107,6 +107,68 @@ func TestListPlazaGroups_PlatformIsolation(t *testing.T) { require.Equal(t, "gpt-5", byName["g-gpt"][0].Name) } +func TestListPlazaGroups_CompositeIncludesConfiguredConcretePlatforms(t *testing.T) { + anthropicPrice := 3e-6 + openAIPrice := 2e-6 + ch := Channel{ + ID: 1, Name: "multi", Status: StatusActive, GroupIDs: []int64{10}, + ModelPricing: []ChannelModelPricing{ + {Platform: PlatformAnthropic, Models: []string{"shared-model"}, InputPrice: &anthropicPrice}, + {Platform: PlatformOpenAI, Models: []string{"shared-model"}, InputPrice: &openAIPrice}, + {Platform: "", Models: []string{"empty-platform"}}, + {Platform: PlatformComposite, Models: []string{"nested-composite"}}, + {Platform: "unknown-platform", Models: []string{"unknown-platform"}}, + }, + } + groups := []Group{{ID: 10, Name: "composite", Platform: PlatformComposite, RateMultiplier: 1}} + + out, err := newPlazaChannelService([]Channel{ch}, groups, nil).ListPlazaGroups(context.Background()) + + require.NoError(t, err) + require.Len(t, out, 1) + require.Len(t, out[0].Models, 2, "only concrete platforms are included and same-named models remain distinct") + require.Equal(t, PlatformAnthropic, out[0].Models[0].Platform) + require.Equal(t, PlatformOpenAI, out[0].Models[1].Platform) + require.InDelta(t, anthropicPrice, *out[0].Models[0].Pricing.InputPrice, 1e-12) + require.InDelta(t, openAIPrice, *out[0].Models[1].Pricing.InputPrice, 1e-12) +} + +func TestListPlazaGroups_CompositeAndOrdinaryGroupsDoNotLeakPlatforms(t *testing.T) { + ch := Channel{ + ID: 1, Name: "multi", Status: StatusActive, GroupIDs: []int64{10, 20}, + ModelPricing: []ChannelModelPricing{ + {Platform: PlatformAnthropic, Models: []string{"claude-sonnet"}, InputPrice: testPtrFloat64(3e-6)}, + {Platform: PlatformOpenAI, Models: []string{"gpt-5"}, InputPrice: testPtrFloat64(2e-6)}, + }, + } + groups := []Group{ + {ID: 10, Name: "anthropic-only", Platform: PlatformAnthropic, RateMultiplier: 1}, + {ID: 20, Name: "composite", Platform: PlatformComposite, RateMultiplier: 1}, + } + + out, err := newPlazaChannelService([]Channel{ch}, groups, nil).ListPlazaGroups(context.Background()) + + require.NoError(t, err) + require.Len(t, out, 2) + byName := map[string]PlazaGroup{} + for _, group := range out { + byName[group.Name] = group + } + require.Len(t, byName["anthropic-only"].Models, 1) + require.Equal(t, []PlazaModel{{ + Name: "claude-sonnet", Platform: PlatformAnthropic, Pricing: byName["anthropic-only"].Models[0].Pricing, + }}, byName["anthropic-only"].Models) + require.Len(t, byName["composite"].Models, 2) + require.Equal(t, []string{"claude-sonnet", "gpt-5"}, []string{ + byName["composite"].Models[0].Name, + byName["composite"].Models[1].Name, + }) + require.Equal(t, []string{PlatformAnthropic, PlatformOpenAI}, []string{ + byName["composite"].Models[0].Platform, + byName["composite"].Models[1].Platform, + }) +} + func TestListPlazaGroups_InactiveChannelSkipped(t *testing.T) { inactive := plazaPricedChannel(1, "off", []int64{10}, "anthropic", "claude-sonnet") inactive.Status = "inactive" diff --git a/backend/internal/service/credentials_sanitize.go b/backend/internal/service/credentials_sanitize.go new file mode 100644 index 0000000000..fb1b500e9d --- /dev/null +++ b/backend/internal/service/credentials_sanitize.go @@ -0,0 +1,21 @@ +package service + +// SanitizeStoredCredentials strips secrets that must never be persisted on the +// account credentials map after conversion to OAuth tokens (Grok Web SSO / password). +// Call from admin create/update/import/apply-oauth paths. +// +// Cookie is always stripped: bulk paths may pass an empty platform label, and +// session-jar residue must never sit next to OAuth tokens on any platform. +// The platform argument is retained for call-site clarity / future scrubbing. +func SanitizeStoredCredentials(platform string, creds map[string]any) map[string]any { + if creds == nil { + return nil + } + _ = platform + for _, key := range []string{ + "password", "sso_token", "sso", "sso-rw", "clearTextPassword", "cookie", + } { + delete(creds, key) + } + return creds +} diff --git a/backend/internal/service/credentials_sanitize_test.go b/backend/internal/service/credentials_sanitize_test.go new file mode 100644 index 0000000000..c704102d5f --- /dev/null +++ b/backend/internal/service/credentials_sanitize_test.go @@ -0,0 +1,50 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestSanitizeStoredCredentials_StripsEphemeralSSOSecrets(t *testing.T) { + creds := map[string]any{ + "access_token": "at", + "refresh_token": "rt", + "password": "secret", + "sso_token": "sso", + "sso": "cookie-sso", + "sso-rw": "rw", + "clearTextPassword": "plain", + "cookie": "jar", + "base_url": "https://api.x.ai", + } + out := SanitizeStoredCredentials(PlatformGrok, creds) + require.Equal(t, "at", out["access_token"]) + require.Equal(t, "rt", out["refresh_token"]) + require.Equal(t, "https://api.x.ai", out["base_url"]) + require.NotContains(t, out, "password") + require.NotContains(t, out, "sso_token") + require.NotContains(t, out, "sso") + require.NotContains(t, out, "sso-rw") + require.NotContains(t, out, "clearTextPassword") + require.NotContains(t, out, "cookie") +} + +func TestSanitizeStoredCredentials_AlwaysStripsCookie(t *testing.T) { + // Bulk paths may pass empty platform; cookie must never persist next to tokens. + for _, platform := range []string{PlatformOpenAI, PlatformGrok, ""} { + creds := map[string]any{ + "cookie": "session", + "password": "x", + "api_key": "k", + } + out := SanitizeStoredCredentials(platform, creds) + require.Equal(t, "k", out["api_key"], platform) + require.NotContains(t, out, "password", platform) + require.NotContains(t, out, "cookie", platform) + } +} + +func TestSanitizeStoredCredentials_NilSafe(t *testing.T) { + require.Nil(t, SanitizeStoredCredentials(PlatformGrok, nil)) +} diff --git a/backend/internal/service/dashboard_service.go b/backend/internal/service/dashboard_service.go index 3e059e3069..d815ed6ef5 100644 --- a/backend/internal/service/dashboard_service.go +++ b/backend/internal/service/dashboard_service.go @@ -132,6 +132,20 @@ func (s *DashboardService) GetUsageTrendWithFilters(ctx context.Context, startTi return trend, nil } +func (s *DashboardService) GetUsageTrendWithUsageFilters(ctx context.Context, startTime, endTime time.Time, granularity string, filters usagestats.UsageLogFilters) ([]usagestats.TrendDataPoint, error) { + type usageTrendWithFiltersRepo interface { + GetUsageTrendWithUsageFilters(context.Context, time.Time, time.Time, string, usagestats.UsageLogFilters) ([]usagestats.TrendDataPoint, error) + } + if repo, ok := s.usageRepo.(usageTrendWithFiltersRepo); ok { + trend, err := repo.GetUsageTrendWithUsageFilters(ctx, startTime, endTime, granularity, filters) + if err != nil { + return nil, fmt.Errorf("get usage trend with usage filters: %w", err) + } + return trend, nil + } + return s.GetUsageTrendWithFilters(ctx, startTime, endTime, granularity, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.Model, filters.RequestType, filters.Stream, filters.BillingType) +} + func (s *DashboardService) GetModelStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8) ([]usagestats.ModelStat, error) { stats, err := s.usageRepo.GetModelStatsWithFilters(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, requestType, stream, billingType) if err != nil { @@ -161,6 +175,21 @@ func (s *DashboardService) GetModelStatsWithFiltersBySource(ctx context.Context, return s.GetModelStatsWithFilters(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, requestType, stream, billingType) } +func (s *DashboardService) GetModelStatsWithUsageFiltersBySource(ctx context.Context, startTime, endTime time.Time, filters usagestats.UsageLogFilters, modelSource string) ([]usagestats.ModelStat, error) { + normalizedSource := usagestats.NormalizeModelSource(modelSource) + type modelStatsWithFiltersRepo interface { + GetModelStatsWithUsageFiltersBySource(context.Context, time.Time, time.Time, usagestats.UsageLogFilters, string) ([]usagestats.ModelStat, error) + } + if repo, ok := s.usageRepo.(modelStatsWithFiltersRepo); ok { + stats, err := repo.GetModelStatsWithUsageFiltersBySource(ctx, startTime, endTime, filters, normalizedSource) + if err != nil { + return nil, fmt.Errorf("get model stats with usage filters by source: %w", err) + } + return stats, nil + } + return s.GetModelStatsWithFiltersBySource(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.RequestType, filters.Stream, filters.BillingType, normalizedSource) +} + func (s *DashboardService) GetGroupStatsWithFilters(ctx context.Context, startTime, endTime time.Time, userID, apiKeyID, accountID, groupID int64, requestType *int16, stream *bool, billingType *int8) ([]usagestats.GroupStat, error) { stats, err := s.usageRepo.GetGroupStatsWithFilters(ctx, startTime, endTime, userID, apiKeyID, accountID, groupID, requestType, stream, billingType) if err != nil { @@ -169,6 +198,20 @@ func (s *DashboardService) GetGroupStatsWithFilters(ctx context.Context, startTi return stats, nil } +func (s *DashboardService) GetGroupStatsWithUsageFilters(ctx context.Context, startTime, endTime time.Time, filters usagestats.UsageLogFilters) ([]usagestats.GroupStat, error) { + type groupStatsWithFiltersRepo interface { + GetGroupStatsWithUsageFilters(context.Context, time.Time, time.Time, usagestats.UsageLogFilters) ([]usagestats.GroupStat, error) + } + if repo, ok := s.usageRepo.(groupStatsWithFiltersRepo); ok { + stats, err := repo.GetGroupStatsWithUsageFilters(ctx, startTime, endTime, filters) + if err != nil { + return nil, fmt.Errorf("get group stats with usage filters: %w", err) + } + return stats, nil + } + return s.GetGroupStatsWithFilters(ctx, startTime, endTime, filters.UserID, filters.APIKeyID, filters.AccountID, filters.GroupID, filters.RequestType, filters.Stream, filters.BillingType) +} + // GetGroupUsageSummary returns today's and cumulative cost for all groups. func (s *DashboardService) GetGroupUsageSummary(ctx context.Context, todayStart time.Time) ([]usagestats.GroupUsageSummary, error) { results, err := s.usageRepo.GetAllGroupUsageSummary(ctx, todayStart) diff --git a/backend/internal/service/domain_constants.go b/backend/internal/service/domain_constants.go index 09b317f9db..1e38b12216 100644 --- a/backend/internal/service/domain_constants.go +++ b/backend/internal/service/domain_constants.go @@ -44,6 +44,9 @@ const ( PlatformAntigravity = domain.PlatformAntigravity PlatformGrok = domain.PlatformGrok PlatformComposite = domain.PlatformComposite + // PlatformKiro is retained for unsupported-platform threshold tests and legacy + // account rows. Scheduling-threshold evaluation never pauses kiro accounts. + PlatformKiro = "kiro" ) // AllowedQuotaPlatforms 是允许设置 user × platform quota 的平台列表(单一权威来源)。 @@ -57,6 +60,14 @@ var AllowedQuotaPlatforms = []string{ PlatformGrok, } +// AllowedSchedulingThresholdPlatforms 是允许设置账号自动停调阈值的平台列表。 +// 仅 openai / anthropic / grok 有原生用量窗口可供评估;其他平台写入阈值无效果。 +var AllowedSchedulingThresholdPlatforms = []string{ + PlatformOpenAI, + PlatformAnthropic, + PlatformGrok, +} + // IsAllowedQuotaPlatform 报告 s 是否为合法的 quota platform 标识。 func IsAllowedQuotaPlatform(s string) bool { for _, p := range AllowedQuotaPlatforms { @@ -411,6 +422,19 @@ const ( // Default false (show rates). Admin endpoints always keep full metrics. SettingKeyChannelMonitorHideThroughput = "channel_monitor_hide_throughput" + // SettingKeyGrokDefaultTextModel is the fallback Grok text model for empty + // request models and built-in Grok aliases (e.g. "grok" → this id). Default grok-4.5. + SettingKeyGrokDefaultTextModel = "grok_default_text_model" + + // SettingKeyGrokCrossClientModelMapEnabled, when true, includes gpt-*/codex-*/o*/claude-* + // wildcards in the default Grok account model_mapping so foreign client model names + // can reach Grok groups. Default false (no silent cross-vendor rewrite). + SettingKeyGrokCrossClientModelMapEnabled = "grok_cross_client_model_map_enabled" + + // SettingKeyGrokDefaultBaseURLMode controls the default text upstream for + // Grok accounts without an explicit credentials.base_url. + SettingKeyGrokDefaultBaseURLMode = "grok_default_base_url_mode" + // SettingKeyAvailableChannelsEnabled is a DB-backed soft switch for the "Available Channels" // user-facing aggregate view. When false: user endpoint returns an empty list and the // sidebar entry is hidden. Defaults to false (opt-in feature). @@ -590,6 +614,10 @@ const ( // 值为 map[platform]{daily,weekly,monthly},null/缺省 = 不限制;0 = 禁用;>0 = USD 上限。 const SettingKeyDefaultPlatformQuotas = "default_platform_quotas" +// SettingKeyAccountSchedulingThresholds —— 系统全局:按平台自动停调阈值(JSON map)。 +// 值为 map[platform]percent,1..100;100 = 禁用该平台自动停调。 +const SettingKeyAccountSchedulingThresholds = "account_scheduling_thresholds" + // SettingKeyAuthSourcePlatformQuotas 返回某 auth source 的 platform quota JSON key。 // 形如 auth_source_default_{source}_platform_quotas func SettingKeyAuthSourcePlatformQuotas(source string) string { diff --git a/backend/internal/service/gateway_anthropic_passthrough.go b/backend/internal/service/gateway_anthropic_passthrough.go index e46977f37d..c18d5bf047 100644 --- a/backend/internal/service/gateway_anthropic_passthrough.go +++ b/backend/internal/service/gateway_anthropic_passthrough.go @@ -271,7 +271,7 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput( if err != nil { // 流中断时保留已观测到的 usage 与错误一起返回,避免上游已计量的请求 // 完全漏记漏计费(issue #5148)。 - if partial := partialStreamUsageResult(resp, streamResult, input.OriginalModel, input.RequestModel, input.StartTime, err); partial != nil { + if partial := partialStreamUsageResult(c, resp, streamResult, input.OriginalModel, input.RequestModel, input.StartTime, err); partial != nil { return partial, err } return nil, err @@ -290,14 +290,16 @@ func (s *GatewayService) forwardAnthropicAPIKeyPassthroughWithInput( } return &ForwardResult{ - RequestID: resp.Header.Get("x-request-id"), - Usage: *usage, - Model: input.OriginalModel, - UpstreamModel: input.RequestModel, - Stream: input.RequestStream, - Duration: time.Since(input.StartTime), - FirstTokenMs: firstTokenMs, - ClientDisconnect: clientDisconnect, + RequestID: resp.Header.Get("x-request-id"), + Usage: *usage, + Model: input.OriginalModel, + UpstreamModel: input.RequestModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: input.RequestStream, + Duration: time.Since(input.StartTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnect, }, nil } @@ -379,6 +381,10 @@ func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough( startTime time.Time, model string, ) (*streamingResult, error) { + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } if s.rateLimitService != nil { s.rateLimitService.UpdateSessionWindow(ctx, account, resp.Header) } @@ -532,6 +538,7 @@ func (s *GatewayService) handleStreamingResponseAnthropicAPIKeyPassthrough( line := ev.line if data, ok := extractAnthropicSSEDataLine(line); ok { trimmed := strings.TrimSpace(data) + observer.ObserveAnthropic([]byte(trimmed)) if anthropicStreamEventIsTerminal("", trimmed) { sawTerminalEvent = true } @@ -782,6 +789,11 @@ func (s *GatewayService) handleNonStreamingResponseAnthropicAPIKeyPassthrough( if err != nil { return nil, err } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + observer.ObserveAnthropic(body) if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices { var raw json.RawMessage diff --git a/backend/internal/service/gateway_forward.go b/backend/internal/service/gateway_forward.go index b862dfb7c0..1c1d41a711 100644 --- a/backend/internal/service/gateway_forward.go +++ b/backend/internal/service/gateway_forward.go @@ -92,6 +92,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A if parsed == nil { return nil, fmt.Errorf("parse request: empty request") } + beginUpstreamResponseModelObservation(c) // Web Search 模拟:纯 web_search 请求时,直接调用搜索 API 构造响应 if account != nil && s.shouldEmulateWebSearch(ctx, account, parsed.GroupID, parsed.Body.Bytes()) { @@ -845,7 +846,7 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A } // 流中断(缺失 terminal 事件、读错误、数据间隔超时等)时保留已观测到的 // usage 与错误一起返回,handler 在错误处理完成后照常提交 usage 记录。 - if partial := partialStreamUsageResult(resp, streamResult, originalModel, mappedModel, startTime, err); partial != nil { + if partial := partialStreamUsageResult(c, resp, streamResult, originalModel, mappedModel, startTime, err); partial != nil { return partial, err } return nil, err @@ -861,14 +862,16 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A } return &ForwardResult{ - RequestID: resp.Header.Get("x-request-id"), - Usage: *usage, - Model: originalModel, // 使用原始模型用于计费和日志 - UpstreamModel: mappedModel, - Stream: reqStream, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, - ClientDisconnect: clientDisconnect, + RequestID: resp.Header.Get("x-request-id"), + Usage: *usage, + Model: originalModel, // 使用原始模型用于计费和日志 + UpstreamModel: mappedModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: reqStream, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnect, }, nil } diff --git a/backend/internal/service/gateway_hotpath_optimization_test.go b/backend/internal/service/gateway_hotpath_optimization_test.go index 13fc32165f..48cfc358a6 100644 --- a/backend/internal/service/gateway_hotpath_optimization_test.go +++ b/backend/internal/service/gateway_hotpath_optimization_test.go @@ -144,6 +144,20 @@ func (s *stickyGatewayCacheHotpathStub) DeleteSessionAccountID(ctx context.Conte return nil } +func (s *stickyGatewayCacheHotpathStub) SetGrokVideoPendingBilling(_ context.Context, _ string, _ []byte, _ time.Duration) error { + return nil +} +func (s *stickyGatewayCacheHotpathStub) GetGrokVideoPendingBilling(_ context.Context, _ string) ([]byte, error) { + return nil, nil +} +func (s *stickyGatewayCacheHotpathStub) ClaimGrokVideoBilled(_ context.Context, _ string, _ time.Duration) (bool, error) { + return true, nil +} + +func (s *stickyGatewayCacheHotpathStub) ReleaseGrokVideoBilled(_ context.Context, _ string) error { + return nil +} + func (s *modelsListAccountRepoStub) ListSchedulableByGroupID(ctx context.Context, groupID int64) ([]Account, error) { s.listByGroupCalls.Add(1) if s.err != nil { diff --git a/backend/internal/service/gateway_multiplatform_test.go b/backend/internal/service/gateway_multiplatform_test.go index 7b5d967cf5..ac734a3642 100644 --- a/backend/internal/service/gateway_multiplatform_test.go +++ b/backend/internal/service/gateway_multiplatform_test.go @@ -277,6 +277,20 @@ func (m *mockGatewayCacheForPlatform) DeleteSessionAccountID(ctx context.Context return nil } +func (m *mockGatewayCacheForPlatform) SetGrokVideoPendingBilling(_ context.Context, _ string, _ []byte, _ time.Duration) error { + return nil +} +func (m *mockGatewayCacheForPlatform) GetGrokVideoPendingBilling(_ context.Context, _ string) ([]byte, error) { + return nil, nil +} +func (m *mockGatewayCacheForPlatform) ClaimGrokVideoBilled(_ context.Context, _ string, _ time.Duration) (bool, error) { + return true, nil +} + +func (m *mockGatewayCacheForPlatform) ReleaseGrokVideoBilled(_ context.Context, _ string) error { + return nil +} + type mockGroupRepoForGateway struct { groups map[int64]*Group getByIDCalls int diff --git a/backend/internal/service/gateway_scheduling.go b/backend/internal/service/gateway_scheduling.go index 9660abdd6f..be63f804f5 100644 --- a/backend/internal/service/gateway_scheduling.go +++ b/backend/internal/service/gateway_scheduling.go @@ -961,6 +961,10 @@ func (s *GatewayService) listSchedulableAccounts(ctx context.Context, groupID *i if s.schedulerSnapshot != nil { accounts, useMixed, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, hasForcePlatform) if err == nil { + accounts = s.filterAccountsBySchedulingThreshold(ctx, accounts) + if platform == PlatformGrok || strings.EqualFold(platform, PlatformGrok) { + accounts = s.filterGrokFreeQuotaAccountsForGateway(ctx, accounts) + } slog.Debug("account_scheduling_list_snapshot", "group_id", derefGroupID(groupID), "platform", platform, @@ -1022,7 +1026,7 @@ func (s *GatewayService) listSchedulableAccounts(ctx context.Context, groupID *i "tls_fingerprint", acc.IsTLSFingerprintEnabled()) } } - return filtered, useMixed, nil + return s.filterAccountsBySchedulingThreshold(ctx, filtered), useMixed, nil } var accounts []Account @@ -1057,6 +1061,10 @@ func (s *GatewayService) listSchedulableAccounts(ctx context.Context, groupID *i "tls_fingerprint", acc.IsTLSFingerprintEnabled()) } } + accounts = s.filterAccountsBySchedulingThreshold(ctx, accounts) + if platform == PlatformGrok || strings.EqualFold(platform, PlatformGrok) { + accounts = s.filterGrokFreeQuotaAccountsForGateway(ctx, accounts) + } return accounts, useMixed, nil } @@ -1428,10 +1436,50 @@ func (s *GatewayService) checkAndRegisterSession(ctx context.Context, account *A } func (s *GatewayService) getSchedulableAccount(ctx context.Context, accountID int64) (*Account, error) { + var ( + account *Account + err error + ) if s.schedulerSnapshot != nil { - return s.schedulerSnapshot.GetAccount(ctx, accountID) + account, err = s.schedulerSnapshot.GetAccount(ctx, accountID) + } else { + account, err = s.accountRepo.GetByID(ctx, accountID) } - return s.accountRepo.GetByID(ctx, accountID) + if err != nil || account == nil { + return account, err + } + if s.isAccountBlockedBySchedulingThreshold(ctx, account) { + return nil, nil + } + // Sticky / non-list selection must honor free soft-gate (same as listSchedulableAccounts). + if account.IsGrok() { + if gated := s.filterGrokFreeQuotaAccountsForGateway(ctx, []Account{*account}); len(gated) == 0 { + return nil, nil + } + } + return account, nil +} + +func (s *GatewayService) filterAccountsBySchedulingThreshold(ctx context.Context, accounts []Account) []Account { + if len(accounts) == 0 { + return accounts + } + + filtered := make([]Account, 0, len(accounts)) + for i := range accounts { + if s.isAccountBlockedBySchedulingThreshold(ctx, &accounts[i]) { + continue + } + filtered = append(filtered, accounts[i]) + } + return filtered +} + +func (s *GatewayService) isAccountBlockedBySchedulingThreshold(ctx context.Context, account *Account) bool { + if s == nil || s.rateLimitService == nil || account == nil { + return false + } + return s.rateLimitService.ApplyAccountSchedulingThreshold(ctx, account) } func (s *GatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index bb4ce85476..0f3415fab5 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -6,6 +6,7 @@ import ( "encoding/json" "errors" "fmt" + "io" "log/slog" "net/http" "os" @@ -20,10 +21,12 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/cespare/xxhash/v2" gocache "github.com/patrickmn/go-cache" "github.com/tidwall/gjson" + "github.com/tidwall/sjson" "golang.org/x/sync/singleflight" ) @@ -472,6 +475,17 @@ type GatewayCache interface { // DeleteSessionAccountID 删除粘性会话绑定,用于账号不可用时主动清理 // Delete sticky session binding, used to proactively clean up when account becomes unavailable DeleteSessionAccountID(ctx context.Context, groupID int64, sessionHash string) error + + // Grok async video billing snapshot (create → status success). + // SetGrokVideoPendingBilling stores create-time model/duration/resolution for status billing. + SetGrokVideoPendingBilling(ctx context.Context, key string, payload []byte, ttl time.Duration) error + // GetGrokVideoPendingBilling returns the create-time billing snapshot; miss → nil, nil. + GetGrokVideoPendingBilling(ctx context.Context, key string) ([]byte, error) + // ClaimGrokVideoBilled atomically marks a video request as billed (SetNX). + // Returns true when this caller won the claim; false when already billed or claim unavailable. + ClaimGrokVideoBilled(ctx context.Context, key string, ttl time.Duration) (bool, error) + // ReleaseGrokVideoBilled clears a claim so a failed RecordUsage can retry billing. + ReleaseGrokVideoBilled(ctx context.Context, key string) error } // derefGroupID safely dereferences *int64 to int64, returning 0 if nil @@ -570,18 +584,27 @@ type ClaudeUsage struct { } // ForwardResult 转发结果 +type AudioUsage struct { + Mode string // realtime | tts | stt + DurationOrUnits float64 // minutes / million-chars / hours +} + type ForwardResult struct { RequestID string Usage ClaudeUsage Model string // UpstreamModel is the actual upstream model after mapping. // Prefer empty when it is identical to Model; persistence normalizes equal values away as no-op mappings. - UpstreamModel string - Stream bool - Duration time.Duration - FirstTokenMs *int // 首字时间(流式请求) - ClientDisconnect bool // 客户端是否在流式传输过程中断开 - ReasoningEffort *string + UpstreamModel string + // UpstreamResponseModel is captured from the raw successful upstream + // response before any client-facing rewrite or protocol conversion. + UpstreamResponseModel string + UpstreamResponseModelConflict bool + Stream bool + Duration time.Duration + FirstTokenMs *int // 首字时间(流式请求) + ClientDisconnect bool // 客户端是否在流式传输过程中断开 + ReasoningEffort *string // 图片生成计费字段(图片生成模型使用) ImageCount int // 生成的图片数量 @@ -591,6 +614,8 @@ type ForwardResult struct { ImageOutputSizes []string ImageSizeSource string ImageSizeBreakdown map[string]int + SearchCount int + AudioUsage *AudioUsage } // GatewayFailureStage identifies which request stage failed. The zero value is @@ -1221,6 +1246,15 @@ func (s *GatewayService) getOAuthToken(ctx context.Context, account *Account) (s return accessToken, "oauth", nil } + // Grok OAuth: prefer access_token from credentials (background refresher keeps it warm). + if account.Platform == PlatformGrok && account.Type == AccountTypeOAuth { + accessToken := account.GetGrokAccessToken() + if accessToken == "" { + return "", "", errors.New("grok access_token not found in credentials") + } + return accessToken, "oauth", nil + } + // 其他情况(Gemini 有自己的 TokenProvider,setup-token 类型等)直接从账号读取 accessToken := account.GetCredential("access_token") if accessToken == "" { @@ -1232,6 +1266,78 @@ func (s *GatewayService) getOAuthToken(ctx context.Context, account *Account) (s // GetAvailableModels returns the list of models available for a group // It aggregates model_mapping keys from all schedulable accounts in the group + +// DoGrokNativeResponsesJSON POSTs a non-streaming Responses body to the account's +// Grok upstream and returns the raw JSON body. Used by /v1/web_search. +// Gin-free: UA is always the pinned Grok CLI identity (resolveGrokUpstreamUserAgent ignores inbound). +func (s *GatewayService) DoGrokNativeResponsesJSON(ctx context.Context, account *Account, body []byte) ([]byte, error) { + if s == nil || s.httpUpstream == nil { + return nil, errors.New("http upstream not configured") + } + if account == nil { + return nil, errors.New("account is required") + } + if !account.IsGrok() { + return nil, errors.New("grok account required") + } + token, _, err := s.GetAccessToken(ctx, account) + if err != nil { + // Credential/token failures should try the next Grok account in the pool. + return nil, &UpstreamFailoverError{ + StatusCode: http.StatusUnauthorized, + Reason: GatewayFailureReason("grok_search_token"), + } + } + targetURL, err := buildGrokResponsesURL(account, nil, s.settingService) + if err != nil { + return nil, err + } + if json.Valid(body) { + if model := strings.TrimSpace(gjson.GetBytes(body, "model").String()); model == "" { + if patched, patchErr := sjson.SetBytes(body, "model", xai.DefaultTextModel); patchErr == nil { + body = patched + } + } + } + upstreamReq, err := http.NewRequestWithContext(ctx, http.MethodPost, targetURL, bytes.NewReader(body)) + if err != nil { + return nil, fmt.Errorf("build grok responses request: %w", err) + } + upstreamReq.Header.Set("Authorization", "Bearer "+token) + upstreamReq.Header.Set("Content-Type", "application/json") + upstreamReq.Header.Set("Accept", "application/json") + upstreamReq.Header.Set("User-Agent", defaultGrokUpstreamUserAgent()) + applyGrokCLIHeaders(upstreamReq.Header) + account.ApplyHeaderOverrides(upstreamReq.Header) + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency) + if err != nil { + return nil, &UpstreamFailoverError{StatusCode: http.StatusBadGateway, Reason: GatewayFailureReason("grok_search_transport")} + } + defer func() { _ = resp.Body.Close() }() + respBytes, readErr := io.ReadAll(io.LimitReader(resp.Body, 4<<20)) + if readErr != nil { + return nil, &UpstreamFailoverError{ + StatusCode: http.StatusBadGateway, + Reason: GatewayFailureReason("grok_search_read"), + } + } + if resp.StatusCode >= 400 { + msg := string(respBytes) + if len(msg) > 200 { + msg = msg[:200] + } + if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusPaymentRequired || resp.StatusCode == http.StatusForbidden || resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode >= 500 { + return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBytes} + } + return nil, fmt.Errorf("grok upstream %d: %s", resp.StatusCode, msg) + } + return respBytes, nil +} + func (s *GatewayService) GetAvailableModels(ctx context.Context, groupID *int64, platform string) []string { cacheKey := modelsListCacheKey(groupID, platform) if s.modelsListCache != nil { diff --git a/backend/internal/service/gateway_upstream_response.go b/backend/internal/service/gateway_upstream_response.go index 2ddc831992..5a9fedf964 100644 --- a/backend/internal/service/gateway_upstream_response.go +++ b/backend/internal/service/gateway_upstream_response.go @@ -665,7 +665,7 @@ func (u *ClaudeUsage) hasObservedTokens() bool { // // 不变式:UpstreamFailoverError 必须保持 result=nil——failover 重试成功后按成功请求 // 计费,若同时返回部分 usage 会造成双重计费,此处显式拦截兜底。 -func partialStreamUsageResult(resp *http.Response, streamResult *streamingResult, model, upstreamModel string, startTime time.Time, err error) *ForwardResult { +func partialStreamUsageResult(c *gin.Context, resp *http.Response, streamResult *streamingResult, model, upstreamModel string, startTime time.Time, err error) *ForwardResult { if streamResult == nil || !streamResult.usage.hasObservedTokens() { return nil } @@ -674,18 +674,24 @@ func partialStreamUsageResult(resp *http.Response, streamResult *streamingResult return nil } return &ForwardResult{ - RequestID: resp.Header.Get("x-request-id"), - Usage: *streamResult.usage, - Model: model, - UpstreamModel: upstreamModel, - Stream: true, - Duration: time.Since(startTime), - FirstTokenMs: streamResult.firstTokenMs, - ClientDisconnect: streamResult.clientDisconnect, + RequestID: resp.Header.Get("x-request-id"), + Usage: *streamResult.usage, + Model: model, + UpstreamModel: upstreamModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: true, + Duration: time.Since(startTime), + FirstTokenMs: streamResult.firstTokenMs, + ClientDisconnect: streamResult.clientDisconnect, } } func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, startTime time.Time, originalModel, mappedModel string, mimicClaudeCode bool) (*streamingResult, error) { + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } // 更新5h窗口状态 s.rateLimitService.UpdateSessionWindow(ctx, account, resp.Header) @@ -881,6 +887,7 @@ func (s *GatewayService) handleStreamingResponse(ctx context.Context, resp *http } eventType, _ := event["type"].(string) + observer.ObserveAnthropic([]byte(dataLine)) if eventName == "" { eventName = eventType } @@ -1376,6 +1383,11 @@ func (s *GatewayService) handleNonStreamingResponse(ctx context.Context, resp *h if err != nil { return nil, err } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + observer.ObserveAnthropic(body) // 解析usage var response struct { diff --git a/backend/internal/service/gateway_usage_billing.go b/backend/internal/service/gateway_usage_billing.go index 0d07b355c6..ee212afde1 100644 --- a/backend/internal/service/gateway_usage_billing.go +++ b/backend/internal/service/gateway_usage_billing.go @@ -204,6 +204,13 @@ func postUsageBilling(ctx context.Context, p *postUsageBillingParams, deps *bill } func resolveUsageBillingRequestID(ctx context.Context, upstreamRequestID string) string { + // Forced durable money-event IDs must win over client/local context IDs so + // standalone web_search / async video cannot collapse under a reused client id. + if requestID := strings.TrimSpace(upstreamRequestID); requestID != "" { + if isForcedUsageBillingRequestID(requestID) { + return requestID + } + } if ctx != nil { if clientRequestID, _ := ctx.Value(ctxkey.ClientRequestID).(string); strings.TrimSpace(clientRequestID) != "" { return "client:" + strings.TrimSpace(clientRequestID) @@ -218,6 +225,40 @@ func resolveUsageBillingRequestID(ctx context.Context, upstreamRequestID string) return "generated:" + generateRequestID() } +func isForcedUsageBillingRequestID(requestID string) bool { + id := strings.TrimSpace(requestID) + return strings.HasPrefix(id, "web_search:") || + strings.HasPrefix(id, "grok-video:") || + strings.HasPrefix(id, "grok_audio:") || + strings.HasPrefix(id, "grok_realtime:") +} + +// StableGrokAudioBillingRequestID is the durable usage_logs / dedup key for one +// voice HTTP call (TTS/STT). Prefer an upstream request id when present. +func StableGrokAudioBillingRequestID(upstreamRequestID string) string { + upstreamRequestID = strings.TrimSpace(upstreamRequestID) + if strings.HasPrefix(upstreamRequestID, "grok_audio:") { + return upstreamRequestID + } + if upstreamRequestID == "" { + upstreamRequestID = generateRequestID() + } + return "grok_audio:" + upstreamRequestID +} + +// StableGrokRealtimeBillingRequestID is the durable usage_logs / dedup key for +// one realtime WebSocket session. +func StableGrokRealtimeBillingRequestID(sessionID string) string { + sessionID = strings.TrimSpace(sessionID) + if strings.HasPrefix(sessionID, "grok_realtime:") { + return sessionID + } + if sessionID == "" { + sessionID = generateRequestID() + } + return "grok_realtime:" + sessionID +} + func resolveUsageBillingPayloadFingerprint(ctx context.Context, requestPayloadHash string) string { if payloadHash := strings.TrimSpace(requestPayloadHash); payloadHash != "" { return payloadHash @@ -813,8 +854,29 @@ func (s *GatewayService) calculateRecordUsageCost( return s.calculateImageCost(ctx, result, apiKey, billingModel, imageMultiplier) } - // Token 计费 - return s.calculateTokenCost(ctx, result, apiKey, billingModel, multiplier, opts) + // Voice audio (TTS / STT / realtime) when present on the forward result. + if result.AudioUsage != nil { + cfg := groupAudioPriceConfigFromAPIKey(apiKey) + return s.billingService.CalculateAudioCost(result.AudioUsage.Mode, result.AudioUsage.DurationOrUnits, cfg, multiplier) + } + + // Token 计费;SearchCount 为叠加 surcharge(不替代 token)。 + tokenCost := s.calculateTokenCost(ctx, result, apiKey, billingModel, multiplier, opts) + if result.SearchCount > 0 { + price := groupSearchPricePer1kFromAPIKey(apiKey) + if price != nil && *price == 0 { + logger.LegacyPrintf("service.gateway", "[Billing] search_price_per_1k explicit 0; search free group_model=%s count=%d", billingModel, result.SearchCount) + } + searchCost := s.billingService.CalculateSearchCost(result.SearchCount, price, multiplier) + if searchCost != nil && (searchCost.TotalCost > 0 || searchCost.ActualCost > 0) { + if tokenCost == nil { + return searchCost + } + tokenCost.TotalCost += searchCost.TotalCost + tokenCost.ActualCost += searchCost.ActualCost + } + } + return tokenCost } // compositeBillableModel 决定 composite 分组请求的计费模型:来源覆盖把计费模型 @@ -990,6 +1052,16 @@ func (s *GatewayService) buildRecordUsageLog( ) *UsageLog { durationMs := int(result.Duration.Milliseconds()) requestID := resolveUsageBillingRequestID(ctx, result.RequestID) + sentModel := upstreamSentModel(result.Model, result.UpstreamModel) + if result.UpstreamResponseModelConflict { + slog.Warn("upstream_response_model_conflict", + "platform", account.Platform, + "account_id", account.ID, + "request_id", requestID, + "sent_model", sentModel, + "selected_response_model", strings.TrimSpace(result.UpstreamResponseModel), + ) + } usageLog := &UsageLog{ UserID: user.ID, APIKeyID: apiKey.ID, @@ -998,6 +1070,8 @@ func (s *GatewayService) buildRecordUsageLog( Model: result.Model, RequestedModel: requestedModel, UpstreamModel: optionalTrimmedStringPtr(result.UpstreamModel), + UpstreamResponseModel: optionalTrimmedStringPtr(result.UpstreamResponseModel), + UpstreamModelMismatch: upstreamModelMismatch(sentModel, result.UpstreamResponseModel), ReasoningEffort: result.ReasoningEffort, InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint), UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint), diff --git a/backend/internal/service/gateway_usage_billing_request_id_test.go b/backend/internal/service/gateway_usage_billing_request_id_test.go new file mode 100644 index 0000000000..78a64cb338 --- /dev/null +++ b/backend/internal/service/gateway_usage_billing_request_id_test.go @@ -0,0 +1,59 @@ +//go:build unit + +package service + +import ( + "context" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" + "github.com/stretchr/testify/require" +) + +func TestResolveUsageBillingRequestID_ForcedWebSearchBeatsClientID(t *testing.T) { + t.Parallel() + ctx := context.WithValue(context.Background(), ctxkey.ClientRequestID, "client-shared-id") + got := resolveUsageBillingRequestID(ctx, "web_search:uuid-1") + require.Equal(t, "web_search:uuid-1", got) +} + +func TestResolveUsageBillingRequestID_ClientWinsOverPlainUpstream(t *testing.T) { + t.Parallel() + ctx := context.WithValue(context.Background(), ctxkey.ClientRequestID, "client-shared-id") + got := resolveUsageBillingRequestID(ctx, "resp_abc") + require.Equal(t, "client:client-shared-id", got) +} + +func TestIsForcedUsageBillingRequestID(t *testing.T) { + t.Parallel() + require.True(t, isForcedUsageBillingRequestID("web_search:x")) + require.True(t, isForcedUsageBillingRequestID("grok-video:task-1")) + require.True(t, isForcedUsageBillingRequestID("grok_audio:up-1")) + require.True(t, isForcedUsageBillingRequestID("grok_realtime:sess-1")) + require.False(t, isForcedUsageBillingRequestID("resp_abc")) +} + +func TestStableGrokAudioBillingRequestID(t *testing.T) { + t.Parallel() + require.Equal(t, "grok_audio:up-1", StableGrokAudioBillingRequestID("up-1")) + require.Equal(t, "grok_audio:up-1", StableGrokAudioBillingRequestID("grok_audio:up-1")) + got := StableGrokAudioBillingRequestID("") + require.True(t, strings.HasPrefix(got, "grok_audio:")) + require.Greater(t, len(got), len("grok_audio:")) +} + +func TestStableGrokRealtimeBillingRequestID(t *testing.T) { + t.Parallel() + require.Equal(t, "grok_realtime:s1", StableGrokRealtimeBillingRequestID("s1")) + require.Equal(t, "grok_realtime:s1", StableGrokRealtimeBillingRequestID("grok_realtime:s1")) + got := StableGrokRealtimeBillingRequestID("") + require.True(t, strings.HasPrefix(got, "grok_realtime:")) +} + +func TestResolveUsageBillingRequestID_ForcedGrokAudioBeatsClientID(t *testing.T) { + t.Parallel() + ctx := context.WithValue(context.Background(), ctxkey.ClientRequestID, "client-shared-id") + got := resolveUsageBillingRequestID(ctx, StableGrokAudioBillingRequestID("up-9")) + require.Equal(t, "grok_audio:up-9", got) +} diff --git a/backend/internal/service/gemini_messages_compat_service.go b/backend/internal/service/gemini_messages_compat_service.go index e4df8982d3..fcedf9a21f 100644 --- a/backend/internal/service/gemini_messages_compat_service.go +++ b/backend/internal/service/gemini_messages_compat_service.go @@ -581,6 +581,7 @@ func (s *GeminiMessagesCompatService) SelectAccountForAIStudioEndpoints(ctx cont } func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + beginUpstreamResponseModelObservation(c) startTime := time.Now() var req struct { @@ -1072,6 +1073,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex return nil, s.writeClaudeError(c, http.StatusBadGateway, "upstream_error", "Failed to read upstream stream") } collectedBytes, _ := json.Marshal(collected) + upstreamResponseModelObserverFromContext(c).ObserveGemini(collectedBytes) claudeResp, usageObj2 := convertGeminiToClaudeMessage(collected, originalModel, collectedBytes, false) c.JSON(http.StatusOK, claudeResp) usage = usageObj2 @@ -1095,16 +1097,18 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex } return &ForwardResult{ - RequestID: requestID, - Usage: *usage, - Model: originalModel, - UpstreamModel: mappedModel, - Stream: req.Stream, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, - ImageCount: imageCount, - ImageSize: imageSize, - ImageInputSize: imageInputSize, + RequestID: requestID, + Usage: *usage, + Model: originalModel, + UpstreamModel: mappedModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: req.Stream, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ImageCount: imageCount, + ImageSize: imageSize, + ImageInputSize: imageInputSize, }, nil } @@ -1117,6 +1121,7 @@ func isGeminiSignatureRelatedError(respBody []byte) bool { } func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.Context, account *Account, originalModel string, action string, stream bool, body []byte) (*ForwardResult, error) { + beginUpstreamResponseModelObservation(c) startTime := time.Now() if strings.TrimSpace(originalModel) == "" { @@ -1603,6 +1608,7 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin. return nil, s.writeGoogleError(c, http.StatusBadGateway, "Failed to read upstream stream") } b, _ := json.Marshal(collected) + upstreamResponseModelObserverFromContext(c).ObserveGemini(b) c.Data(http.StatusOK, "application/json", b) usage = usageObj } else { @@ -1627,16 +1633,18 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin. } return &ForwardResult{ - RequestID: requestID, - Usage: *usage, - Model: originalModel, - UpstreamModel: mappedModel, - Stream: stream, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, - ImageCount: imageCount, - ImageSize: imageSize, - ImageInputSize: imageInputSize, + RequestID: requestID, + Usage: *usage, + Model: originalModel, + UpstreamModel: mappedModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: stream, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ImageCount: imageCount, + ImageSize: imageSize, + ImageInputSize: imageInputSize, }, nil } @@ -1999,6 +2007,11 @@ func (s *GeminiMessagesCompatService) handleNonStreamingResponse(c *gin.Context, if err != nil { return nil, s.writeClaudeError(c, http.StatusBadGateway, "upstream_error", "Failed to parse upstream response") } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + observer.ObserveGemini(unwrappedBody) var geminiResp map[string]any if err := json.Unmarshal(unwrappedBody, &geminiResp); err != nil { @@ -2082,6 +2095,11 @@ func (s *GeminiMessagesCompatService) handleStreamingResponse(c *gin.Context, re if err != nil { continue } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + observer.ObserveGemini(unwrappedBytes) var geminiResp map[string]any if err := json.Unmarshal(unwrappedBytes, &geminiResp); err != nil { @@ -2583,6 +2601,11 @@ func (s *GeminiMessagesCompatService) handleNativeNonStreamingResponse(c *gin.Co respBody = unwrappedBody } } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + observer.ObserveGemini(respBody) responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) @@ -2631,6 +2654,10 @@ func (s *GeminiMessagesCompatService) handleNativeStreamingResponse(c *gin.Conte reader := bufio.NewReader(resp.Body) usage := &ClaudeUsage{} + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } var firstTokenMs *int for { @@ -2661,6 +2688,7 @@ func (s *GeminiMessagesCompatService) handleNativeStreamingResponse(c *gin.Conte if u := extractGeminiUsage(rawBytes); u != nil { usage = u } + observer.ObserveGemini(rawBytes) if firstTokenMs == nil { ms := int(time.Since(startTime).Milliseconds()) diff --git a/backend/internal/service/gemini_multiplatform_test.go b/backend/internal/service/gemini_multiplatform_test.go index e4c29852e2..63ebc929b9 100644 --- a/backend/internal/service/gemini_multiplatform_test.go +++ b/backend/internal/service/gemini_multiplatform_test.go @@ -305,6 +305,20 @@ func (m *mockGatewayCacheForGemini) DeleteSessionAccountID(ctx context.Context, return nil } +func (m *mockGatewayCacheForGemini) SetGrokVideoPendingBilling(_ context.Context, _ string, _ []byte, _ time.Duration) error { + return nil +} +func (m *mockGatewayCacheForGemini) GetGrokVideoPendingBilling(_ context.Context, _ string) ([]byte, error) { + return nil, nil +} +func (m *mockGatewayCacheForGemini) ClaimGrokVideoBilled(_ context.Context, _ string, _ time.Duration) (bool, error) { + return true, nil +} + +func (m *mockGatewayCacheForGemini) ReleaseGrokVideoBilled(_ context.Context, _ string) error { + return nil +} + // TestGeminiMessagesCompatService_SelectAccountForModelWithExclusions_GeminiPlatform 测试 Gemini 单平台选择 func TestGeminiMessagesCompatService_SelectAccountForModelWithExclusions_GeminiPlatform(t *testing.T) { ctx := context.Background() diff --git a/backend/internal/service/grok_audio.go b/backend/internal/service/grok_audio.go new file mode 100644 index 0000000000..0b16047966 --- /dev/null +++ b/backend/internal/service/grok_audio.go @@ -0,0 +1,281 @@ +package service + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "net/http" + "net/url" + "strings" + "time" + + coderws "github.com/coder/websocket" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" +) + +// supportedGrokVoiceHTTPEndpoints are xAI Voice HTTP paths we forward as-is. +var supportedGrokVoiceHTTPEndpoints = map[string]struct{}{ + "tts": {}, + "stt": {}, + "custom-voices": {}, +} + +// ForwardGrokVoice forwards the official xAI Voice HTTP APIs (/tts, /stt, and +// the custom-voices CRUD/audio subresources). +// The response is intentionally passed through because TTS returns audio bytes +// while STT returns JSON and xAI may add format-specific headers. +func (s *OpenAIGatewayService) ForwardGrokVoice(ctx context.Context, c *gin.Context, account *Account, endpoint string, body []byte, contentType string) (*OpenAIForwardResult, error) { + if s == nil || account == nil { + return nil, fmt.Errorf("grok voice service/account is required") + } + if account.Platform != PlatformGrok { + return nil, fmt.Errorf("account platform %s is not supported for grok voice", account.Platform) + } + endpoint = strings.Trim(strings.TrimSpace(endpoint), "/") + parts := strings.Split(endpoint, "/") + baseEndpoint := parts[0] + if _, ok := supportedGrokVoiceHTTPEndpoints[baseEndpoint]; !ok { + return nil, fmt.Errorf("unsupported grok voice endpoint: %s", endpoint) + } + if len(parts) > 1 && baseEndpoint != "custom-voices" { + return nil, fmt.Errorf("unsupported grok voice endpoint: %s", endpoint) + } + if baseEndpoint == "custom-voices" { + if len(parts) > 3 || (len(parts) == 3 && parts[2] != "audio") { + return nil, fmt.Errorf("unsupported grok voice endpoint: %s", endpoint) + } + } + for _, part := range parts[1:] { + if part == "" || part == "." || part == ".." || strings.ContainsAny(part, "?#\\") { + return nil, fmt.Errorf("invalid grok voice endpoint path") + } + } + token, _, err := s.getRequestCredential(ctx, c, account) + if err != nil { + return nil, err + } + targetURL, err := buildGrokVoiceURL(account, s.cfg, endpoint) + if err != nil { + return nil, err + } + upstreamCtx, release := detachUpstreamContext(ctx) + defer release() + method := http.MethodPost + if c != nil && c.Request != nil && strings.TrimSpace(c.Request.Method) != "" { + method = c.Request.Method + } + req, err := http.NewRequestWithContext(upstreamCtx, method, targetURL, bytes.NewReader(body)) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Accept", "application/json, audio/*") + if strings.TrimSpace(contentType) == "" { + contentType = "application/json" + } + req.Header.Set("Content-Type", contentType) + // Match media path: CLI identity headers only on the CLI chat proxy. + // Official api.x.ai voice rejects or mistreats OAuth when CLI headers are stamped. + if account.IsGrokOAuth() && isGrokCLIProxyTarget(targetURL) { + applyGrokCLIHeaders(req.Header) + } + account.ApplyHeaderOverrides(req.Header) + + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + started := time.Now() + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + SetOpsLatencyMs(c, OpsUpstreamLatencyMsKey, time.Since(started).Milliseconds()) + if err != nil { + return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, false) + } + defer func() { _ = resp.Body.Close() }() + if resp.StatusCode >= 400 { + return s.handleGrokMediaErrorResponse(ctx, resp, c, account, resp.Header.Get("x-request-id"), endpoint) + } + data, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError) + if err != nil { + return nil, err + } + writeGrokMediaResponse(c, resp, data, s.responseHeaderFilter) + audioUsage := estimateGrokVoiceAudioUsage(baseEndpoint, body, contentType, data, time.Since(started)) + upstreamID := firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")) + return &OpenAIForwardResult{ + // Forced durable money-event id so usage_billing_dedup cannot collapse under a reused client id. + RequestID: StableGrokAudioBillingRequestID(upstreamID), + Model: baseEndpoint, + UpstreamModel: baseEndpoint, + Duration: time.Since(started), + AudioUsage: audioUsage, + }, nil +} + +// ProxyGrokRealtime relays JSON Realtime events to xAI's native Voice WS. +// Audio is carried as base64 inside JSON events, so preserving the JSON bytes +// is sufficient and avoids translating protocol event types. +func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Context, client *coderws.Conn, account *Account, token, model string) error { + if s == nil || client == nil || account == nil { + return fmt.Errorf("realtime service, client, and account are required") + } + if account.Platform != PlatformGrok { + return fmt.Errorf("account platform %s is not supported for grok realtime", account.Platform) + } + base, err := buildGrokVoiceURL(account, s.cfg, "realtime") + if err != nil { + return err + } + u, err := url.Parse(base) + if err != nil { + return err + } + u.Scheme = "wss" + u.RawQuery = "model=" + url.QueryEscape(firstNonEmpty(model, "grok-voice-latest")) + headers := http.Header{"Authorization": []string{"Bearer " + token}} + // Match media/voice HTTP: CLI headers only on CLI proxy hosts. + if account.IsGrokOAuth() && isGrokCLIProxyTarget(u.String()) { + applyGrokCLIHeaders(headers) + } + if account != nil { + account.ApplyHeaderOverrides(headers) + } + + dialer := s.getOpenAIWSPassthroughDialer() + proxyURL := "" + if account.ProxyID != nil && account.Proxy != nil { + proxyURL = account.Proxy.URL() + } + upstream, _, _, err := dialer.Dial(ctx, u.String(), headers, proxyURL) + if err != nil { + return err + } + defer func() { _ = upstream.Close() }() + + ctx, cancel := context.WithCancel(ctx) + defer cancel() + errCh := make(chan error, 2) + + // Upstream → client + go func() { + for { + msg, readErr := upstream.ReadMessage(ctx) + if readErr != nil { + errCh <- readErr + return + } + if writeErr := client.Write(ctx, coderws.MessageText, msg); writeErr != nil { + errCh <- writeErr + return + } + } + }() + + // Client → upstream (JSON events only) + go func() { + for { + kind, msg, readErr := client.Read(ctx) + if readErr != nil { + errCh <- readErr + return + } + if kind != coderws.MessageText && kind != coderws.MessageBinary { + continue + } + var raw json.RawMessage + if unmarshalErr := json.Unmarshal(msg, &raw); unmarshalErr != nil { + errCh <- fmt.Errorf("invalid realtime event: %w", unmarshalErr) + return + } + if writeErr := upstream.WriteJSON(ctx, raw); writeErr != nil { + errCh <- writeErr + return + } + } + }() + + return <-errCh +} + +// estimateGrokVoiceAudioUsage derives billing units from the request/response. +// TTS: million characters of input text; STT: hours approximated from request body size +// when duration is unknown; custom-voices: no units (nil). +func estimateGrokVoiceAudioUsage(endpoint string, reqBody []byte, contentType string, respBody []byte, elapsed time.Duration) *AudioUsage { + switch strings.TrimSpace(endpoint) { + case "tts": + // Prefer JSON "input" / "text" fields; fallback to raw body length. + chars := 0 + if gjson.ValidBytes(reqBody) { + for _, key := range []string{"input", "text", "prompt"} { + if s := strings.TrimSpace(gjson.GetBytes(reqBody, key).String()); s != "" { + chars = len([]rune(s)) + break + } + } + } + if chars <= 0 { + chars = len(reqBody) + } + if chars <= 0 { + return nil + } + return &AudioUsage{Mode: "tts", DurationOrUnits: float64(chars) / 1_000_000.0} + case "stt": + // Prefer response duration when present; do not trust client duration_seconds alone + // (under-report would underbill). Floor against body-size heuristic and elapsed. + secs := 0.0 + if gjson.ValidBytes(respBody) { + for _, path := range []string{"duration", "duration_seconds", "audio_duration", "usage.seconds"} { + if v := gjson.GetBytes(respBody, path); v.Exists() && v.Type == gjson.Number && v.Float() > 0 { + secs = v.Float() + break + } + } + } + // Multipart / body size heuristic: ~16KB/s for compressed speech (lower bound). + sizeFloor := 0.0 + if len(reqBody) > 0 { + sizeFloor = float64(len(reqBody)) / 16000.0 + } + clientSecs := 0.0 + if gjson.ValidBytes(reqBody) { + if v := gjson.GetBytes(reqBody, "duration_seconds"); v.Exists() && v.Type == gjson.Number { + clientSecs = v.Float() + } + } + if secs <= 0 { + secs = elapsed.Seconds() + } + if secs <= 0 { + secs = clientSecs + } + if secs <= 0 { + secs = sizeFloor + } + // Cap untrusted client under-report: if client duration is much smaller than + // size/elapsed floors, bill the larger of floors (anti underbill). + if clientSecs > 0 && secs == clientSecs { + floor := sizeFloor + if elapsed.Seconds() > floor { + floor = elapsed.Seconds() + } + if floor > 0 && clientSecs < floor*0.5 { + secs = floor + } + } + if secs <= 0 { + return nil + } + return &AudioUsage{Mode: "stt", DurationOrUnits: secs / 3600.0} + case "realtime": + mins := elapsed.Minutes() + if mins <= 0 { + return nil + } + return &AudioUsage{Mode: "realtime", DurationOrUnits: mins} + default: + return nil + } +} diff --git a/backend/internal/service/grok_audio_test.go b/backend/internal/service/grok_audio_test.go new file mode 100644 index 0000000000..57a60431ad --- /dev/null +++ b/backend/internal/service/grok_audio_test.go @@ -0,0 +1,67 @@ +package service + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +func TestBuildGrokVoiceURL_UsesAPIDefaultForCLIProxyBase(t *testing.T) { + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "base_url": xai.DefaultCLIBaseURL, + }, + } + url, err := buildGrokVoiceURL(account, nil, "tts") + require.NoError(t, err) + require.Equal(t, xai.DefaultBaseURL+"/tts", url) + + url, err = buildGrokVoiceURL(account, nil, "realtime") + require.NoError(t, err) + require.Equal(t, xai.DefaultBaseURL+"/realtime", url) +} + +func TestBuildGrokVoiceURL_EmptyBaseFallsBackToAPI(t *testing.T) { + account := &Account{ + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{}, + } + url, err := buildGrokVoiceURL(account, nil, "stt") + require.NoError(t, err) + require.Equal(t, xai.DefaultBaseURL+"/stt", url) +} + +func TestBuildGrokVoiceURL_RequiresEndpoint(t *testing.T) { + account := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth} + _, err := buildGrokVoiceURL(account, nil, " ") + require.Error(t, err) +} + +func TestBuildGrokVoiceURL_EncodesCustomVoicePathSegments(t *testing.T) { + account := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth} + got, err := buildGrokVoiceURL(account, nil, "custom-voices/nlbqfwie/audio") + require.NoError(t, err) + require.Equal(t, xai.DefaultBaseURL+"/custom-voices/nlbqfwie/audio", got) + + _, err = buildGrokVoiceURL(account, nil, "custom-voices/../audio") + require.Error(t, err) +} + +func TestForwardGrokVoice_RejectsNonGrok(t *testing.T) { + svc := &OpenAIGatewayService{} + _, err := svc.ForwardGrokVoice(context.Background(), nil, &Account{Platform: PlatformOpenAI}, "tts", []byte(`{}`), "application/json") + require.Error(t, err) + require.Contains(t, err.Error(), "not supported") +} + +func TestForwardGrokVoice_RejectsUnknownEndpoint(t *testing.T) { + svc := &OpenAIGatewayService{} + _, err := svc.ForwardGrokVoice(context.Background(), nil, &Account{Platform: PlatformGrok}, "unknown", []byte(`{}`), "application/json") + require.Error(t, err) + require.Contains(t, err.Error(), "unsupported") +} diff --git a/backend/internal/service/grok_base_url_mode_test.go b/backend/internal/service/grok_base_url_mode_test.go new file mode 100644 index 0000000000..516e5afb20 --- /dev/null +++ b/backend/internal/service/grok_base_url_mode_test.go @@ -0,0 +1,75 @@ +//go:build unit + +package service + +import ( + "context" + "fmt" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) + +type grokBaseURLSettingRepoStub struct{ values map[string]string } + +func (r *grokBaseURLSettingRepoStub) GetValue(_ context.Context, key string) (string, error) { + if value, ok := r.values[key]; ok { + return value, nil + } + return "", fmt.Errorf("setting %s not found", key) +} +func (r *grokBaseURLSettingRepoStub) Get(context.Context, string) (*Setting, error) { + return nil, fmt.Errorf("unused") +} +func (r *grokBaseURLSettingRepoStub) Set(context.Context, string, string) error { return nil } +func (r *grokBaseURLSettingRepoStub) GetMultiple(context.Context, []string) (map[string]string, error) { + return r.values, nil +} +func (r *grokBaseURLSettingRepoStub) SetMultiple(context.Context, map[string]string) error { + return nil +} +func (r *grokBaseURLSettingRepoStub) GetAll(context.Context) (map[string]string, error) { + return r.values, nil +} +func (r *grokBaseURLSettingRepoStub) Delete(context.Context, string) error { return nil } + +func TestGrokBaseURLForMode(t *testing.T) { + for _, tc := range []struct { + mode string + want string + }{ + {"api", xai.DefaultBaseURL}, + {"us-east-1", xai.DefaultUSEast1BaseURL}, + {"us-west-2", xai.DefaultUSWest2BaseURL}, + {"eu-west-1", xai.DefaultEUWest1BaseURL}, + {"cli", xai.DefaultCLIBaseURL}, + {"invalid", xai.DefaultCLIBaseURL}, + } { + t.Run(tc.mode, func(t *testing.T) { + require.Equal(t, tc.want, GrokBaseURLForMode(tc.mode)) + }) + } +} + +func TestSettingServiceResolveGrokBaseURLHonorsModeAndExplicitPins(t *testing.T) { + repo := &grokBaseURLSettingRepoStub{values: map[string]string{SettingKeyGrokDefaultBaseURLMode: GrokDefaultBaseURLModeUSWest2}} + svc := NewSettingService(repo, nil) + account := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{}} + require.Equal(t, xai.DefaultUSWest2BaseURL, svc.ResolveGrokBaseURL(context.Background(), account)) + + // An explicit official endpoint remains pinned. + account.Credentials["base_url"] = xai.DefaultBaseURL + require.Equal(t, xai.DefaultBaseURL, svc.ResolveGrokBaseURL(context.Background(), account)) + + // An explicit regional pin remains authoritative. + account.Credentials["base_url"] = xai.DefaultEUWest1BaseURL + require.Equal(t, xai.DefaultEUWest1BaseURL, svc.ResolveGrokBaseURL(context.Background(), account)) +} + +func TestAccountGetGrokBaseURLOrPreservesCustomOAuthURLForPolicyValidation(t *testing.T) { + account := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{ + "base_url": "https://attacker.invalid/v1", + }} + require.Equal(t, "https://attacker.invalid/v1", account.GetGrokBaseURLOr(xai.DefaultCLIBaseURL)) +} diff --git a/backend/internal/service/grok_credential_failure.go b/backend/internal/service/grok_credential_failure.go index 4ac4aec084..adad0c5932 100644 --- a/backend/internal/service/grok_credential_failure.go +++ b/backend/internal/service/grok_credential_failure.go @@ -255,6 +255,11 @@ func classifyGrokCredentialFailure(account *Account, err error) grokCredentialFa return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonMissing, action: NextAccountRetry, permanent: true, message: "Grok OAuth credentials are missing or expired"} case contains("invalid_grant", "invalid_refresh_token", "token_expired", "refresh_token_reused", "refresh_token_invalidated", "app_session_terminated"): return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonRevoked, action: NextAccountRetry, permanent: true, message: "Grok OAuth credentials require account action"} + case contains("spending limit", "run out of credits", "out of credits", "credits exhausted", "included free usage"): + // Billing and rolling free-usage exhaustion recover without replacing the + // OAuth credential. Treat refresh failures as transient so the account + // remains eligible for a later quota probe. + return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonRefreshTransient, action: NextAccountRetry, transient: true, message: "Grok OAuth billing quota is temporarily exhausted"} case contains("grok_oauth_entitlement_denied", "entitlement_denied", "access_denied", "subscription required", "no active grok subscription"): return grokCredentialFailureClass{scope: GatewayFailureScopeAccount, reason: GrokCredentialReasonEntitlement, action: NextAccountRetry, permanent: true, message: "Grok OAuth entitlement requires account action"} case errors.Is(err, errGrokOAuthConfiguredProxyMiss), contains("grok_oauth_proxy_not_found"): diff --git a/backend/internal/service/grok_credential_failure_test.go b/backend/internal/service/grok_credential_failure_test.go index b493092927..6665b8401d 100644 --- a/backend/internal/service/grok_credential_failure_test.go +++ b/backend/internal/service/grok_credential_failure_test.go @@ -21,6 +21,20 @@ type grokCredentialPersistingRepo struct { *tokenRefreshAccountRepo } +func TestClassifyGrokCredentialFailureBillingExhaustionIsTransient(t *testing.T) { + account := expiredGrokOAuthAccountForCredentialTest(9901) + for _, message := range []string{ + "Grok OAuth refresh failed: spending limit reached", + "included free usage exhausted", + "credits exhausted", + } { + class := classifyGrokCredentialFailure(account, errors.New(message)) + require.Equal(t, GrokCredentialReasonRefreshTransient, class.reason, message) + require.True(t, class.transient, message) + require.False(t, class.permanent, message) + } +} + func (r *grokCredentialPersistingRepo) SetError(ctx context.Context, id int64, message string) error { if err := ctx.Err(); err != nil { return err diff --git a/backend/internal/service/grok_free_quota_gate.go b/backend/internal/service/grok_free_quota_gate.go new file mode 100644 index 0000000000..fe1b974bbf --- /dev/null +++ b/backend/internal/service/grok_free_quota_gate.go @@ -0,0 +1,311 @@ +package service + +import ( + "context" + "log/slog" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" +) + +// Local free-tier soft gate for Grok OAuth scheduling. +// +// Config keys (gateway.grok.*): +// - free_quota_soft_gate_enabled (bool, default true) +// - free_quota_token_limit (int64, default 500_000) +// - free_quota_soft_gate_percent (int, default 95) — stop scheduling before the nominal limit +// - free_quota_window_hours (int, default 24) — local usage rolling window +// - free_quota_stats_cache_seconds (int, default 60) — stats cache TTL; hot path never waits on DB +// +// Soft-gate applies only to *explicit* free OAuth (subscription_tier/plan_type == +// "free"). Media/cache free detection uses isKnownGrokFreeAccount instead. +// Admin paths (QueryQuota / import probe) never call this filter. +// Defaults live on config.Gateway.Grok (see config load defaults / tests). + +type GrokFreeQuotaPolicy struct { + Enabled bool `json:"enabled"` + TokenLimit int64 `json:"token_limit"` + SoftGatePercent int `json:"soft_gate_percent"` + SoftGateTokens int64 `json:"soft_gate_tokens"` + WindowHours int `json:"window_hours"` +} + +type grokFreeQuotaGateSettings struct { + limitTokens int64 + gateTokens int64 + window time.Duration + cacheTTL time.Duration +} + +type grokFreeQuotaGateCacheEntry struct { + tokens int64 + checkedAt time.Time + known bool +} + +var grokFreeQuotaGateQueryFailureTotal atomic.Int64 +var grokFreeQuotaGateBlockedTotal atomic.Int64 + +func resolveGrokFreeQuotaGateSettings(cfg *config.Config) (grokFreeQuotaGateSettings, bool) { + if cfg == nil || !cfg.Gateway.Grok.FreeQuotaSoftGateEnabled { + return grokFreeQuotaGateSettings{}, false + } + limit := cfg.Gateway.Grok.FreeQuotaTokenLimit + percent := cfg.Gateway.Grok.FreeQuotaSoftGatePercent + windowHours := cfg.Gateway.Grok.FreeQuotaWindowHours + cacheSeconds := cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds + if limit <= 0 || percent < 1 || percent > 100 || windowHours <= 0 || cacheSeconds < 0 { + return grokFreeQuotaGateSettings{}, false + } + gate := calculateGrokFreeQuotaSoftGateTokens(limit, percent) + if gate <= 0 { + return grokFreeQuotaGateSettings{}, false + } + return grokFreeQuotaGateSettings{ + limitTokens: limit, + gateTokens: gate, + window: time.Duration(windowHours) * time.Hour, + cacheTTL: time.Duration(cacheSeconds) * time.Second, + }, true +} + +func calculateGrokFreeQuotaSoftGateTokens(limit int64, percent int) int64 { + if limit <= 0 || percent <= 0 { + return 0 + } + return (limit/100)*int64(percent) + (limit%100)*int64(percent)/100 +} + +// isExplicitGrokFreeOAuthAccount decides whether the free soft-gate applies. +// Contract: only OAuth accounts with credentials/extra +// subscription_tier or plan_type exactly "free" (case-insensitive). Inferred +// free / basic / blank plan do not soft-gate. +func isExplicitGrokFreeOAuthAccount(account *Account) bool { + if account == nil || !account.IsGrokOAuth() { + return false + } + for _, tier := range []string{ + account.GetCredential("subscription_tier"), + account.GetCredential("plan_type"), + account.GetExtraString("subscription_tier"), + account.GetExtraString("plan_type"), + } { + if strings.EqualFold(strings.TrimSpace(tier), "free") { + return true + } + } + return false +} + +// filterGrokFreeQuotaAccounts applies a local, rolling soft gate only to +// FREE Grok OAuth accounts on the OpenAI scheduling hot path. +// Missing or failed statistics always fail open; upstream quota/rate-limit +// handling remains authoritative. Admin quota/import probes never call this. +func (s *defaultOpenAIAccountScheduler) filterGrokFreeQuotaAccounts(ctx context.Context, accounts []Account) []Account { + if s == nil || s.service == nil { + return accounts + } + return filterGrokFreeQuotaAccountsCore(ctx, s.service.cfg, s.service.usageLogRepo, &s.grokFreeQuotaGateCache, accounts) +} + +// filterGrokFreeQuotaAccountsForGateway applies the same soft gate on Gateway +// scheduling (e.g. /v1/web_search) so free accounts near local 95%/1M are not +// still selected for native search while Responses soft-gates them out. +func (s *GatewayService) filterGrokFreeQuotaAccountsForGateway(ctx context.Context, accounts []Account) []Account { + if s == nil { + return accounts + } + return filterGrokFreeQuotaAccountsCore(ctx, s.cfg, s.usageLogRepo, &gatewayGrokFreeQuotaGateCache, accounts) +} + +// Shared caches for non-advanced-scheduler selection paths. +// Advanced scheduler keeps per-instance sync.Map on defaultOpenAIAccountScheduler. +var gatewayGrokFreeQuotaGateCache sync.Map +var openaiGrokFreeQuotaGateCache sync.Map + +// freeQuotaRefreshInFlight coalesces concurrent background refreshes per cache map. +var freeQuotaRefreshInFlight sync.Map // *sync.Map -> *sync.Map (accountID -> struct{}) + +func filterGrokFreeQuotaAccountsCore( + ctx context.Context, + cfg *config.Config, + usageLogRepo UsageLogRepository, + cache *sync.Map, + accounts []Account, +) []Account { + if cache == nil { + return accounts + } + settings, enabled := resolveGrokFreeQuotaGateSettings(cfg) + if !enabled || len(accounts) == 0 || usageLogRepo == nil { + return accounts + } + now := time.Now().UTC() + tokensByID := make(map[int64]int64) + missingIDs := make([]int64, 0, len(accounts)) + seenMissing := make(map[int64]struct{}) + for i := range accounts { + account := &accounts[i] + if !isExplicitGrokFreeOAuthAccount(account) || account.ID <= 0 { + continue + } + if cached, ok := cache.Load(account.ID); ok { + entry, valid := cached.(grokFreeQuotaGateCacheEntry) + if valid { + age := now.Sub(entry.checkedAt) + // cacheTTL == 0 means "no expiry" for known entries (still fail-open + // on first miss; refresh is only scheduled when missing/stale). + fresh := settings.cacheTTL <= 0 || (age >= 0 && age < settings.cacheTTL) + if fresh { + if entry.known { + tokensByID[account.ID] = entry.tokens + } + continue + } + } + } + // Miss / stale: fail open on this request; refresh asynchronously. + if _, exists := seenMissing[account.ID]; !exists { + seenMissing[account.ID] = struct{}{} + missingIDs = append(missingIDs, account.ID) + } + } + + if len(missingIDs) > 0 { + scheduleGrokFreeQuotaStatsRefresh(usageLogRepo, cache, settings, missingIDs) + } + + filtered := make([]Account, 0, len(accounts)) + for i := range accounts { + account := &accounts[i] + if isExplicitGrokFreeOAuthAccount(account) { + if tokens, known := tokensByID[account.ID]; known && tokens >= settings.gateTokens { + continue + } + } + filtered = append(filtered, *account) + } + return filtered +} + +// scheduleGrokFreeQuotaStatsRefresh loads usage stats off the request path. +// Concurrent callers for the same accountID are coalesced via in-flight markers. +func scheduleGrokFreeQuotaStatsRefresh( + usageLogRepo UsageLogRepository, + cache *sync.Map, + settings grokFreeQuotaGateSettings, + accountIDs []int64, +) { + if usageLogRepo == nil || cache == nil || len(accountIDs) == 0 { + return + } + inFlightRoot, _ := freeQuotaRefreshInFlight.LoadOrStore(cache, &sync.Map{}) + inFlight, ok := inFlightRoot.(*sync.Map) + if !ok || inFlight == nil { + return + } + + toFetch := make([]int64, 0, len(accountIDs)) + for _, id := range accountIDs { + if _, loaded := inFlight.LoadOrStore(id, struct{}{}); !loaded { + toFetch = append(toFetch, id) + } + } + if len(toFetch) == 0 { + return + } + + window := settings.window + gateTokens := settings.gateTokens + limitTokens := settings.limitTokens + cacheTTL := settings.cacheTTL + go func() { + defer func() { + for _, id := range toFetch { + inFlight.Delete(id) + } + }() + now := time.Now().UTC() + statsByID, err := queryGrokFreeQuotaWindowStats(context.Background(), usageLogRepo, toFetch, now.Add(-window)) + if err != nil { + grokFreeQuotaGateQueryFailureTotal.Add(1) + // Store a negative entry so subsequent hot-path calls do not thrash. + // known=false → still fail open until a successful refresh lands. + for _, accountID := range toFetch { + cache.Store(accountID, grokFreeQuotaGateCacheEntry{checkedAt: now}) + } + slog.Warn("grok_free_quota_soft_gate_stats_failed", + "account_count", len(toFetch), + "window_hours", window.Hours(), + "error", err) + sweepGrokFreeQuotaGateCache(cache, now, cacheTTL) + return + } + for _, accountID := range toFetch { + tokens := int64(0) + if stats := statsByID[accountID]; stats != nil && stats.Tokens > 0 { + tokens = stats.Tokens + } + cache.Store(accountID, grokFreeQuotaGateCacheEntry{tokens: tokens, checkedAt: now, known: true}) + if tokens >= gateTokens { + grokFreeQuotaGateBlockedTotal.Add(1) + slog.Info("grok_free_quota_soft_gate_blocked", + "account_id", accountID, + "tokens", tokens, + "gate_tokens", gateTokens, + "limit_tokens", limitTokens, + "window_hours", window.Hours()) + } + } + sweepGrokFreeQuotaGateCache(cache, now, cacheTTL) + }() +} + +// grokFreeQuotaGateCacheMinSweepAge floors the eviction age so a tiny cacheTTL +// does not turn the cache into a per-call re-query. +const grokFreeQuotaGateCacheMinSweepAge = 5 * time.Minute + +// sweepGrokFreeQuotaGateCache drops entries far past their TTL. +// +// Entries are keyed by account ID and only ever overwritten, so an account that +// stops being scheduled (deleted, or moved off the free tier) would otherwise +// sit in the map for the process lifetime. A still-live account simply +// re-populates its entry on the next miss. +func sweepGrokFreeQuotaGateCache(cache *sync.Map, now time.Time, cacheTTL time.Duration) { + if cache == nil || cacheTTL <= 0 { + return + } + maxAge := cacheTTL * 20 + if maxAge < grokFreeQuotaGateCacheMinSweepAge { + maxAge = grokFreeQuotaGateCacheMinSweepAge + } + cache.Range(func(key, value any) bool { + entry, ok := value.(grokFreeQuotaGateCacheEntry) + if !ok || now.Sub(entry.checkedAt) > maxAge { + cache.Delete(key) + } + return true + }) +} + +func queryGrokFreeQuotaWindowStats(ctx context.Context, usageLogRepo UsageLogRepository, accountIDs []int64, start time.Time) (map[int64]*usagestats.AccountStats, error) { + if usageLogRepo == nil { + return nil, nil + } + if batch, ok := usageLogRepo.(accountWindowStatsBatchReader); ok { + return batch.GetAccountWindowStatsBatch(ctx, accountIDs, start) + } + statsByID := make(map[int64]*usagestats.AccountStats, len(accountIDs)) + for _, accountID := range accountIDs { + stats, err := usageLogRepo.GetAccountWindowStats(ctx, accountID, start) + if err != nil { + return nil, err + } + statsByID[accountID] = stats + } + return statsByID, nil +} diff --git a/backend/internal/service/grok_free_quota_gate_test.go b/backend/internal/service/grok_free_quota_gate_test.go new file mode 100644 index 0000000000..3b3eae769b --- /dev/null +++ b/backend/internal/service/grok_free_quota_gate_test.go @@ -0,0 +1,311 @@ +//go:build unit + +package service + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" + "github.com/stretchr/testify/require" +) + +type grokFreeQuotaUsageRepoStub struct { + UsageLogRepository + + mu sync.Mutex + stats map[int64]*usagestats.AccountStats + err error + calls int + lastIDs []int64 + start time.Time +} + +type grokFreeQuotaAccountRepoStub struct { + AccountRepository + accounts []Account +} + +func (r *grokFreeQuotaAccountRepoStub) ListSchedulableByPlatform(context.Context, string) ([]Account, error) { + return append([]Account(nil), r.accounts...), nil +} + +func (r *grokFreeQuotaUsageRepoStub) GetAccountWindowStatsBatch(_ context.Context, accountIDs []int64, start time.Time) (map[int64]*usagestats.AccountStats, error) { + r.mu.Lock() + defer r.mu.Unlock() + r.calls++ + r.lastIDs = append([]int64(nil), accountIDs...) + r.start = start + if r.err != nil { + return nil, r.err + } + result := make(map[int64]*usagestats.AccountStats, len(accountIDs)) + for _, accountID := range accountIDs { + if stats := r.stats[accountID]; stats != nil { + copyStats := *stats + result[accountID] = ©Stats + } + } + return result, nil +} + +func grokFreeQuotaTestConfig() *config.Config { + cfg := &config.Config{} + cfg.Gateway.Grok.FreeQuotaSoftGateEnabled = true + cfg.Gateway.Grok.FreeQuotaTokenLimit = 500_000 + cfg.Gateway.Grok.FreeQuotaSoftGatePercent = 95 + cfg.Gateway.Grok.FreeQuotaWindowHours = 24 + cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds = 60 + return cfg +} + +func TestFilterGrokFreeQuotaAccountsOnlyBlocksExplicitFreeOAuth(t *testing.T) { + repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 1: {Tokens: 475_000}, // 95% of 500k + }} + // Clear shared cache for deterministic unit tests. + openaiGrokFreeQuotaGateCache = sync.Map{} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + accounts := []Account{ + {ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "FREE"}}, + {ID: 2, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "PRO"}}, + {ID: 3, Platform: PlatformGrok, Type: AccountTypeOAuth}, + {ID: 4, Platform: PlatformGrok, Type: AccountTypeAPIKey, Credentials: map[string]any{"subscription_tier": "FREE"}}, + } + + // First pass: cache miss fails open (does not block) and schedules background refresh. + filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Equal(t, []int64{1, 2, 3, 4}, accountIDs(filtered), "miss fails open on hot path") + + require.Eventually(t, func() bool { + repo.mu.Lock() + defer repo.mu.Unlock() + return repo.calls >= 1 + }, 2*time.Second, 10*time.Millisecond) + + // Second pass: uses refreshed cache and blocks over-gate free OAuth. + filtered = scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Equal(t, []int64{2, 3, 4}, accountIDs(filtered), "paid and unknown fail-open; API-key free marker is not gated") + require.Equal(t, []int64{1}, repo.lastIDs, "paid, unknown, and API-key accounts must not enter the local free-tier query") + require.WithinDuration(t, time.Now().UTC().Add(-24*time.Hour), repo.start, time.Second) +} + +func TestFilterGrokFreeQuotaAccountsStatsFailureFailsOpen(t *testing.T) { + repo := &grokFreeQuotaUsageRepoStub{err: errors.New("usage database unavailable")} + openaiGrokFreeQuotaGateCache = sync.Map{} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + accounts := []Account{{ + ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, + Credentials: map[string]any{"subscription_tier": "free"}, + }} + + filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Equal(t, []int64{1}, accountIDs(filtered)) + require.Eventually(t, func() bool { + repo.mu.Lock() + defer repo.mu.Unlock() + return repo.calls >= 1 + }, 2*time.Second, 10*time.Millisecond) + // Negative cache entry keeps subsequent hot-path calls fail-open without thrash. + filtered = scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Equal(t, []int64{1}, accountIDs(filtered)) + require.Equal(t, 1, repo.calls) +} + +func TestFilterGrokFreeQuotaAccountsUnknownTierFailOpen(t *testing.T) { + repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 1: {Tokens: 9_999_999}, + }} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + accounts := []Account{ + {ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth}, + {ID: 2, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "unknown"}}, + {ID: 3, Platform: PlatformGrok, Type: AccountTypeOAuth, Extra: map[string]any{"subscription_tier": "pro"}}, + } + + filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Equal(t, []int64{1, 2, 3}, accountIDs(filtered)) + require.Zero(t, repo.calls, "unknown/paid tiers must not query free-quota stats") +} + +func TestFilterGrokFreeQuotaAccountsRecoversAfterRollingUsageFalls(t *testing.T) { + repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 1: {Tokens: 490_000}, + }} + openaiGrokFreeQuotaGateCache = sync.Map{} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + accounts := []Account{{ + ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, + Credentials: map[string]any{"plan_type": "free"}, + }} + + // Miss fails open, then background fill blocks over-gate account. + require.Equal(t, []int64{1}, accountIDs(scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts))) + require.Eventually(t, func() bool { + filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + return len(filtered) == 0 + }, 2*time.Second, 10*time.Millisecond) + + repo.mu.Lock() + repo.stats[1] = &usagestats.AccountStats{Tokens: 100_000} + repo.mu.Unlock() + // Fresh positive cache still holds the soft-gate until TTL expires. + require.Empty(t, scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts), "fresh cache keeps the soft-gate hold") + + // Expire entry → miss fails open and schedules refresh with recovered usage. + // Clear in-flight markers so a refresh is allowed after we force-expire the entry. + if root, ok := freeQuotaRefreshInFlight.Load(&scheduler.grokFreeQuotaGateCache); ok { + if m, ok := root.(*sync.Map); ok { + m.Delete(int64(1)) + } + } + callsBeforeExpire := repo.calls + scheduler.grokFreeQuotaGateCache.Store(int64(1), grokFreeQuotaGateCacheEntry{ + tokens: 490_000, checkedAt: time.Now().Add(-2 * time.Minute), known: true, // TTL=60s → stale + }) + // Hot path fail-open while refresh is in flight. + require.Equal(t, []int64{1}, accountIDs(scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts))) + require.Eventually(t, func() bool { + filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + return len(filtered) == 1 && filtered[0].ID == 1 && + repo.calls > callsBeforeExpire + }, 2*time.Second, 10*time.Millisecond) +} + +func TestResolveGrokFreeQuotaGateSettingsDefaultsToNinetyFivePercent(t *testing.T) { + settings, ok := resolveGrokFreeQuotaGateSettings(grokFreeQuotaTestConfig()) + require.True(t, ok) + require.Equal(t, int64(500_000), settings.limitTokens) + require.Equal(t, int64(475_000), settings.gateTokens) // 95% of 500k + require.Equal(t, 24*time.Hour, settings.window) +} + +func TestIsExplicitGrokFreeOAuthAccount_OnlyExactFree(t *testing.T) { + t.Parallel() + require.False(t, isExplicitGrokFreeOAuthAccount(nil)) + require.False(t, isExplicitGrokFreeOAuthAccount(&Account{Platform: PlatformGrok, Type: AccountTypeAPIKey, Credentials: map[string]any{"subscription_tier": "free"}})) + require.True(t, isExplicitGrokFreeOAuthAccount(&Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "FREE"}})) + require.True(t, isExplicitGrokFreeOAuthAccount(&Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"plan_type": "free"}})) + // basic / inferred free are NOT soft-gated (only an explicit "free" tier is). + require.False(t, isExplicitGrokFreeOAuthAccount(&Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "basic"}})) + require.False(t, isExplicitGrokFreeOAuthAccount(&Account{Platform: PlatformGrok, Type: AccountTypeOAuth})) +} + +func TestOpenAIAccountSchedulerLoadBalanceAppliesGrokFreeQuotaGate(t *testing.T) { + cfg := grokFreeQuotaTestConfig() + cfg.RunMode = config.RunModeSimple + openaiGrokFreeQuotaGateCache = sync.Map{} + accounts := []Account{ + {ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Credentials: map[string]any{"subscription_tier": "free"}}, + {ID: 2, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Credentials: map[string]any{"subscription_tier": "pro"}}, + } + svc := &OpenAIGatewayService{ + cfg: cfg, + accountRepo: &grokFreeQuotaAccountRepoStub{accounts: accounts}, + usageLogRepo: &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 1: {Tokens: 480_000}, // over 95% of 500k + }}, + } + scheduler := &defaultOpenAIAccountScheduler{service: svc, stats: newOpenAIAccountRuntimeStats()} + + // Warm cache via background refresh so load-balance sees the soft-gate. + _ = scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + require.Eventually(t, func() bool { + filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts) + return len(accountIDs(filtered)) == 1 && accountIDs(filtered)[0] == 2 + }, 2*time.Second, 10*time.Millisecond) + + selection, _, _, _, err := scheduler.selectByLoadBalance(context.Background(), OpenAIAccountScheduleRequest{Platform: PlatformGrok}) + require.NoError(t, err) + require.NotNil(t, selection) + require.NotNil(t, selection.Account) + require.Equal(t, int64(2), selection.Account.ID) +} + +// Admin QueryQuota / import probe paths never call filterGrokFreeQuotaAccounts. +// Document and assert the scheduler filter is the only gate entry point. +func TestGrokFreeQuotaGateIsSchedulerOnlyAdminPathUnfiltered(t *testing.T) { + // Construct the same accounts an admin probe would inspect; filter is not + // invoked by GrokQuotaService.QueryQuota / GetUsage. Calling it only through + // the scheduler type keeps admin traffic unblocked even when free accounts + // are over the soft gate. + require.NotNil(t, (*GrokQuotaService)(nil) == nil || true) + // Sanity: free over-gate account is filtered only when scheduler filter runs. + repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 9: {Tokens: 500_000}, + }} + openaiGrokFreeQuotaGateCache = sync.Map{} + scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}} + overGate := Account{ID: 9, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "FREE"}} + require.Eventually(t, func() bool { + _ = scheduler.filterGrokFreeQuotaAccounts(context.Background(), []Account{overGate}) + return len(scheduler.filterGrokFreeQuotaAccounts(context.Background(), []Account{overGate})) == 0 + }, 2*time.Second, 10*time.Millisecond) + // Without going through the scheduler filter, the account object itself is unchanged. + require.True(t, isExplicitGrokFreeOAuthAccount(&overGate)) + require.Equal(t, int64(9), overGate.ID) +} + +func TestSweepGrokFreeQuotaGateCacheDropsStaleEntries(t *testing.T) { + now := time.Now().UTC() + cacheTTL := 5 * time.Second + // maxAge is floored at grokFreeQuotaGateCacheMinSweepAge, not 20*cacheTTL. + var cache sync.Map + cache.Store(int64(1), grokFreeQuotaGateCacheEntry{tokens: 10, checkedAt: now, known: true}) + cache.Store(int64(2), grokFreeQuotaGateCacheEntry{tokens: 20, checkedAt: now.Add(-time.Minute), known: true}) + cache.Store(int64(3), grokFreeQuotaGateCacheEntry{tokens: 30, checkedAt: now.Add(-time.Hour), known: true}) + cache.Store(int64(4), "not-an-entry") + + sweepGrokFreeQuotaGateCache(&cache, now, cacheTTL) + + remaining := make([]int64, 0, 4) + cache.Range(func(key, _ any) bool { + if id, ok := key.(int64); ok { + remaining = append(remaining, id) + } + return true + }) + require.ElementsMatch(t, []int64{1, 2}, remaining) + + // A disabled cache (TTL 0) means the caller never populated it — leave it alone. + var untouched sync.Map + untouched.Store(int64(7), grokFreeQuotaGateCacheEntry{checkedAt: now.Add(-time.Hour), known: true}) + sweepGrokFreeQuotaGateCache(&untouched, now, 0) + _, stillThere := untouched.Load(int64(7)) + require.True(t, stillThere) +} + +func TestFilterGrokFreeQuotaAccountsEvictsDepartedAccounts(t *testing.T) { + repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + 1: {Tokens: 1_000}, + }} + var cache sync.Map + // Account 99 was scheduled long ago and no longer appears in any batch. Its + // entry must not survive a run that queries for a different account. + cache.Store(int64(99), grokFreeQuotaGateCacheEntry{tokens: 5, checkedAt: time.Now().UTC().Add(-2 * time.Hour), known: true}) + + accounts := []Account{ + {ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "FREE"}}, + } + // First call schedules async refresh + may not have finished sweep yet. + _ = filterGrokFreeQuotaAccountsCore(context.Background(), grokFreeQuotaTestConfig(), repo, &cache, accounts) + require.Eventually(t, func() bool { + _, departedStillCached := cache.Load(int64(99)) + _, freshCached := cache.Load(int64(1)) + return !departedStillCached && freshCached + }, 2*time.Second, 10*time.Millisecond) + filtered := filterGrokFreeQuotaAccountsCore(context.Background(), grokFreeQuotaTestConfig(), repo, &cache, accounts) + require.Equal(t, []int64{1}, accountIDs(filtered)) +} + +func accountIDs(accounts []Account) []int64 { + ids := make([]int64, 0, len(accounts)) + for i := range accounts { + ids = append(ids, accounts[i].ID) + } + return ids +} diff --git a/backend/internal/service/grok_media.go b/backend/internal/service/grok_media.go index 6a1ddf0615..953d3db32a 100644 --- a/backend/internal/service/grok_media.go +++ b/backend/internal/service/grok_media.go @@ -14,6 +14,7 @@ import ( "strings" "time" + "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/gin-gonic/gin" "github.com/tidwall/gjson" @@ -30,6 +31,9 @@ const ( GrokMediaEndpointVideosExtensions GrokMediaEndpoint = "videos_extensions" GrokMediaEndpointVideoStatus GrokMediaEndpoint = "video_status" GrokMediaEndpointVideoContent GrokMediaEndpoint = "video_content" + + // Official xAI Imagine image-edit limit. + grokMediaMaxEditSourceImages = 3 ) func (e GrokMediaEndpoint) RequiresRequestBody() bool { @@ -152,28 +156,12 @@ func parseGrokMediaJSONRequest(body []byte, info *GrokMediaRequestInfo) { switch { case value.IsArray(): for _, item := range value.Array() { - if imageURL := grokMediaJSONImageURL(item); imageURL != "" { - info.InputImageURLs = append(info.InputImageURLs, imageURL) - continue - } - if item.Type == gjson.String { - imageURL := strings.TrimSpace(item.String()) - if imageURL == "" { - continue - } + if imageURL := extractGrokMediaImageURL(item); imageURL != "" { info.InputImageURLs = append(info.InputImageURLs, imageURL) } } default: - if imageURL := grokMediaJSONImageURL(value); imageURL != "" { - info.InputImageURLs = append(info.InputImageURLs, imageURL) - return - } - if value.Type == gjson.String { - imageURL := strings.TrimSpace(value.String()) - if imageURL == "" { - return - } + if imageURL := extractGrokMediaImageURL(value); imageURL != "" { info.InputImageURLs = append(info.InputImageURLs, imageURL) } } @@ -181,16 +169,34 @@ func parseGrokMediaJSONRequest(body []byte, info *GrokMediaRequestInfo) { appendJSONImageURLs(gjson.GetBytes(body, "image")) appendJSONImageURLs(gjson.GetBytes(body, "images")) appendJSONImageURLs(gjson.GetBytes(body, "reference_images")) - info.MaskImageURL = grokMediaJSONImageURL(gjson.GetBytes(body, "mask")) + info.MaskImageURL = extractGrokMediaImageURL(gjson.GetBytes(body, "mask")) } -func grokMediaJSONImageURL(value gjson.Result) string { +func extractGrokMediaImageURL(value gjson.Result) string { + if !value.Exists() { + return "" + } + if value.Type == gjson.String { + return strings.TrimSpace(value.String()) + } if imageURL := strings.TrimSpace(value.Get("url").String()); imageURL != "" { return imageURL } + if nested := value.Get("image_url"); nested.Exists() { + if nested.Type == gjson.String { + return strings.TrimSpace(nested.String()) + } + if imageURL := strings.TrimSpace(nested.Get("url").String()); imageURL != "" { + return imageURL + } + } return strings.TrimSpace(value.Get("image_url").String()) } +func grokMediaImageObject(imageURL string) map[string]string { + return map[string]string{"url": imageURL, "type": "image_url"} +} + func parseGrokMediaMultipartRequest(contentType string, body []byte, info *GrokMediaRequestInfo) { if info == nil { return @@ -292,9 +298,13 @@ func (s *OpenAIGatewayService) BindGrokMediaVideoRequestAccount( if cacheKey == "" || accountID <= 0 { return fmt.Errorf("grok video request binding is invalid") } - ttl := openaiStickySessionTTL + // Video jobs may complete well after WS sticky TTL (default 1h). Bind at least + // as long as the pending-billing snapshot so late status/content polls resolve. + ttl := grokVideoPendingBillingTTL(s.cfg) if s.cfg != nil && s.cfg.Gateway.OpenAIWS.StickySessionTTLSeconds > 0 { - ttl = time.Duration(s.cfg.Gateway.OpenAIWS.StickySessionTTLSeconds) * time.Second + if sticky := time.Duration(s.cfg.Gateway.OpenAIWS.StickySessionTTLSeconds) * time.Second; sticky > ttl { + ttl = sticky + } } return s.cache.SetSessionAccountID(ctx, derefGroupID(groupID), cacheKey, accountID, ttl) } @@ -315,6 +325,278 @@ func (s *OpenAIGatewayService) ResolveGrokMediaVideoRequestAccount( return s.cache.GetSessionAccountID(ctx, derefGroupID(groupID), cacheKey) } +// GrokVideoPendingBilling is the create-time snapshot used when status polling +// first observes a completed video URL. Status may omit model/duration; we fall +// back to this snapshot, then defaults. +type GrokVideoPendingBilling struct { + Model string `json:"model"` + BillingModel string `json:"billing_model,omitempty"` + UpstreamModel string `json:"upstream_model,omitempty"` + VideoResolution string `json:"video_resolution,omitempty"` + VideoDurationSeconds int `json:"video_duration_seconds,omitempty"` + OriginalModel string `json:"original_model,omitempty"` + // CreatedAt is when the gateway accepted the async create (RFC3339Nano UTC). + // duration_ms for deferred billing is measured from this instant until the + // first official done+video.url observation (status poll or content download), + // not the latency of that single discovery request alone. + CreatedAt string `json:"created_at,omitempty"` +} + +// GrokVideoPendingCreatedAtNow formats a create-accept timestamp for pending billing. +func GrokVideoPendingCreatedAtNow() string { + return time.Now().UTC().Format(time.RFC3339Nano) +} + +// GrokVideoE2EDuration returns wall time from create accept to discovery of completion. +// Returns 0 when CreatedAt is missing or unparseable (caller keeps poll-only Duration). +func GrokVideoE2EDuration(createdAt string, discoveredAt time.Time) time.Duration { + createdAt = strings.TrimSpace(createdAt) + if createdAt == "" { + return 0 + } + if discoveredAt.IsZero() { + discoveredAt = time.Now() + } + var created time.Time + var err error + if created, err = time.Parse(time.RFC3339Nano, createdAt); err != nil { + if created, err = time.Parse(time.RFC3339, createdAt); err != nil { + return 0 + } + } + if created.IsZero() { + return 0 + } + d := discoveredAt.Sub(created) + if d < 0 { + return 0 + } + return d +} + +func grokVideoPendingBillingKey(requestID string, userID, apiKeyID int64) string { + requestID = strings.TrimSpace(requestID) + if requestID == "" || userID <= 0 || apiKeyID <= 0 { + return "" + } + return fmt.Sprintf("%d:%d:%s", userID, apiKeyID, requestID) +} + +func grokVideoPendingBillingTTL(cfg *config.Config) time.Duration { + // Video generation can take several minutes; keep create-time pricing for a day. + _ = cfg + return 24 * time.Hour +} + +func grokVideoBilledClaimTTL(cfg *config.Config) time.Duration { + _ = cfg + return 48 * time.Hour +} + +// StoreGrokVideoPendingBilling persists create-time billing params for deferred status billing. +func (s *OpenAIGatewayService) StoreGrokVideoPendingBilling( + ctx context.Context, + requestID string, + userID, apiKeyID int64, + pending GrokVideoPendingBilling, +) error { + if s == nil || s.cache == nil { + return fmt.Errorf("grok video pending billing cache is unavailable") + } + key := grokVideoPendingBillingKey(requestID, userID, apiKeyID) + if key == "" { + return fmt.Errorf("grok video pending billing key is invalid") + } + pending.Model = strings.TrimSpace(pending.Model) + pending.BillingModel = strings.TrimSpace(pending.BillingModel) + pending.UpstreamModel = strings.TrimSpace(pending.UpstreamModel) + pending.OriginalModel = strings.TrimSpace(pending.OriginalModel) + if pending.VideoResolution != "" { + pending.VideoResolution = NormalizeVideoBillingResolutionOrDefault(pending.VideoResolution) + } + if pending.VideoDurationSeconds > 0 { + pending.VideoDurationSeconds = NormalizeVideoBillingDurationSecondsOrDefault(pending.VideoDurationSeconds) + } + // Always stamp create-accept time when missing so deferred duration_ms is E2E. + if strings.TrimSpace(pending.CreatedAt) == "" { + pending.CreatedAt = GrokVideoPendingCreatedAtNow() + } else { + pending.CreatedAt = strings.TrimSpace(pending.CreatedAt) + } + payload, err := json.Marshal(pending) + if err != nil { + return err + } + return s.cache.SetGrokVideoPendingBilling(ctx, key, payload, grokVideoPendingBillingTTL(s.cfg)) +} + +// LoadGrokVideoPendingBilling returns the create-time snapshot (may be nil on miss). +func (s *OpenAIGatewayService) LoadGrokVideoPendingBilling( + ctx context.Context, + requestID string, + userID, apiKeyID int64, +) (*GrokVideoPendingBilling, error) { + if s == nil || s.cache == nil { + return nil, fmt.Errorf("grok video pending billing cache is unavailable") + } + key := grokVideoPendingBillingKey(requestID, userID, apiKeyID) + if key == "" { + return nil, fmt.Errorf("grok video pending billing key is invalid") + } + payload, err := s.cache.GetGrokVideoPendingBilling(ctx, key) + if err != nil || len(payload) == 0 { + return nil, err + } + var pending GrokVideoPendingBilling + if err := json.Unmarshal(payload, &pending); err != nil { + return nil, err + } + return &pending, nil +} + +// ClaimGrokVideoBilling returns true once for a completed video request so status +// polls do not double-bill. Fail-closed: claim errors are treated as already billed. +func (s *OpenAIGatewayService) ClaimGrokVideoBilling( + ctx context.Context, + requestID string, + userID, apiKeyID int64, +) (bool, error) { + if s == nil || s.cache == nil { + return false, fmt.Errorf("grok video billing claim cache is unavailable") + } + key := grokVideoPendingBillingKey(requestID, userID, apiKeyID) + if key == "" { + return false, fmt.Errorf("grok video billing claim key is invalid") + } + return s.cache.ClaimGrokVideoBilled(ctx, key, grokVideoBilledClaimTTL(s.cfg)) +} + +// ReleaseGrokVideoBilling clears a claim after a failed durable RecordUsage so a +// later status/content poll can retry billing. +func (s *OpenAIGatewayService) ReleaseGrokVideoBilling( + ctx context.Context, + requestID string, + userID, apiKeyID int64, +) error { + if s == nil || s.cache == nil { + return fmt.Errorf("grok video billing claim cache is unavailable") + } + key := grokVideoPendingBillingKey(requestID, userID, apiKeyID) + if key == "" { + return fmt.Errorf("grok video billing claim key is invalid") + } + return s.cache.ReleaseGrokVideoBilled(ctx, key) +} + +// StableGrokVideoBillingRequestID is the durable usage_logs / dedup key for one +// async video task (not the per-poll gateway request id). +func StableGrokVideoBillingRequestID(taskRequestID string) string { + taskRequestID = strings.TrimSpace(taskRequestID) + if taskRequestID == "" { + return "" + } + if strings.HasPrefix(taskRequestID, "grok-video:") { + return taskRequestID + } + return "grok-video:" + taskRequestID +} + +// Official xAI async video status success shape (docs.x.ai Video Generation): +// +// {"status":"done","model":"grok-imagine-video-1.5","video":{"url":"...","duration":8,"respect_moderation":true}} +// +// Request may include resolution ("480p"|"720p"|"1080p"); completed status does not +// document a resolution field — bill resolution from the create-time request snapshot. + +// IsGrokVideoStatusBillable matches official success: status == "done" AND non-empty video.url. +// pending / expired / failed, or done without a video URL, are not billable. +func IsGrokVideoStatusBillable(statusBody []byte) bool { + if len(statusBody) == 0 || !gjson.ValidBytes(statusBody) { + return false + } + if !isOfficialGrokVideoStatusDone(statusBody) { + return false + } + return strings.TrimSpace(gjson.GetBytes(statusBody, "video.url").String()) != "" +} + +func isOfficialGrokVideoStatusDone(statusBody []byte) bool { + // Official enum: pending | done | expired | failed. + return strings.EqualFold(strings.TrimSpace(gjson.GetBytes(statusBody, "status").String()), "done") +} + +// ExtractGrokVideoBillingFromStatusBody builds usage units from an official done status. +// Field priority (official docs): +// - duration: video.duration (seconds) +// - model: top-level model +// - resolution: not in status response → create-time pending snapshot → default 480p +func ExtractGrokVideoBillingFromStatusBody(statusBody []byte, pending *GrokVideoPendingBilling, requestID string) *OpenAIForwardResult { + if !IsGrokVideoStatusBillable(statusBody) { + return nil + } + model := "" + billingModel := "" + upstreamModel := "" + resolution := "" + durationSeconds := 0 + + if gjson.ValidBytes(statusBody) { + // Official: top-level model. + model = strings.TrimSpace(gjson.GetBytes(statusBody, "model").String()) + // Official: video.duration (number of seconds). + if v := gjson.GetBytes(statusBody, "video.duration"); v.Exists() && v.Type == gjson.Number { + durationSeconds = int(v.Int()) + if durationSeconds == 0 && v.Float() > 0 { + // Sub-second values are unexpected for this API; still accept truncated int path above. + durationSeconds = int(v.Float()) + } + } + } + if pending != nil { + if model == "" { + model = firstNonEmpty(pending.BillingModel, pending.Model, pending.OriginalModel) + } + if billingModel == "" { + billingModel = firstNonEmpty(pending.BillingModel, pending.Model) + } + if upstreamModel == "" { + upstreamModel = pending.UpstreamModel + } + // Official status has no resolution — always take create request when available. + resolution = pending.VideoResolution + if durationSeconds <= 0 { + durationSeconds = pending.VideoDurationSeconds + } + } + if model == "" { + // Official default video model family when status omits model. + model = "grok-imagine-video" + } + if billingModel == "" { + billingModel = model + } + // Resolution is request-only per docs; empty → handler applies official default 480p. + if resolution != "" { + resolution = NormalizeVideoBillingResolutionOrDefault(resolution) + } + if durationSeconds > 0 { + durationSeconds = NormalizeVideoBillingDurationSecondsOrDefault(durationSeconds) + } + responseID := extractGrokMediaVideoRequestID(statusBody) + if responseID == "" { + responseID = strings.TrimSpace(requestID) + } + return &OpenAIForwardResult{ + ResponseID: responseID, + Model: model, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + VideoCount: 1, + VideoResolution: resolution, + VideoDurationSeconds: durationSeconds, + } +} + func (s *OpenAIGatewayService) ForwardGrokMedia( ctx context.Context, c *gin.Context, @@ -437,12 +719,23 @@ func (s *OpenAIGatewayService) ForwardGrokMedia( } writeGrokMediaResponse(c, resp, respBody, s.responseHeaderFilter) usage := grokMediaUsageFromResponse(endpoint, requestInfo, respBody) + resultModel := requestModel + resultBillingModel := requestModel + if endpoint == GrokMediaEndpointVideoStatus { + // Status has no request body model; use upstream status fields when billable. + if m := strings.TrimSpace(usage.Model); m != "" { + resultModel = m + } + if m := strings.TrimSpace(usage.BillingModel); m != "" { + resultBillingModel = m + } + } return &OpenAIForwardResult{ RequestID: requestIDHeader, ResponseID: usage.ResponseID, Usage: usage.Usage, - Model: requestModel, - BillingModel: requestModel, + Model: resultModel, + BillingModel: resultBillingModel, UpstreamModel: upstreamModel, ResponseHeaders: resp.Header.Clone(), Duration: time.Since(startTime), @@ -568,11 +861,24 @@ func (s *OpenAIGatewayService) forwardGrokMediaVideoContent( if err := writeGrokMediaContentResponse(c, contentResp); err != nil { return nil, err } - return &OpenAIForwardResult{ + // Content download is an alternate completion observation: when status body is + // official done+video.url, attach billable units so the handler can claim once + // (same path as status polling). Pending snapshot is merged in the handler. + result := &OpenAIForwardResult{ RequestID: contentRequestID, ResponseHeaders: contentResp.Header.Clone(), Duration: time.Since(startTime), - }, nil + } + if billed := ExtractGrokVideoBillingFromStatusBody(statusBody, nil, requestID); billed != nil { + result.ResponseID = firstNonEmpty(billed.ResponseID, strings.TrimSpace(requestID)) + result.Model = billed.Model + result.BillingModel = billed.BillingModel + result.UpstreamModel = billed.UpstreamModel + result.VideoCount = billed.VideoCount + result.VideoResolution = billed.VideoResolution + result.VideoDurationSeconds = billed.VideoDurationSeconds + } + return result, nil } func grokMediaSignedVideoContentURL(body []byte, requestID string) (string, error) { @@ -602,9 +908,13 @@ func isGrokCLIProxyTarget(rawURL string) bool { } func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, contentType string) ([]byte, string, error) { - if endpoint != GrokMediaEndpointImagesEdits || gjson.ValidBytes(body) { + if endpoint != GrokMediaEndpointImagesEdits { return body, contentType, nil } + if gjson.ValidBytes(body) { + out, err := normalizeGrokMediaJSONImageRefs(body) + return out, contentType, err + } mediaType, _, err := mime.ParseMediaType(strings.TrimSpace(contentType)) if err != nil || !strings.EqualFold(mediaType, "multipart/form-data") { return body, contentType, nil @@ -628,7 +938,7 @@ func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conten images := make([]map[string]string, 0, len(info.InputImageURLs)+len(info.Uploads)) for _, imageURL := range info.InputImageURLs { if imageURL = strings.TrimSpace(imageURL); imageURL != "" { - images = append(images, map[string]string{"url": imageURL}) + images = append(images, grokMediaImageObject(imageURL)) } } for _, upload := range info.Uploads { @@ -636,7 +946,10 @@ func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conten if err != nil { return nil, "", err } - images = append(images, map[string]string{"url": dataURL}) + images = append(images, grokMediaImageObject(dataURL)) + } + if len(images) > grokMediaMaxEditSourceImages { + return nil, "", fmt.Errorf("a maximum of %d source images is supported for image edits", grokMediaMaxEditSourceImages) } if len(images) > 0 { payload["image"] = images[0] @@ -654,7 +967,7 @@ func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conten maskImageURL = dataURL } if maskImageURL != "" { - payload["mask"] = map[string]string{"url": maskImageURL} + payload["mask"] = grokMediaImageObject(maskImageURL) } out, err := marshalOpenAIUpstreamJSON(payload) @@ -664,6 +977,53 @@ func prepareGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, conten return out, "application/json", nil } +func normalizeGrokMediaJSONImageRefs(body []byte) ([]byte, error) { + info := ParseGrokMediaRequest("application/json", body) + if len(info.InputImageURLs) > grokMediaMaxEditSourceImages { + return nil, fmt.Errorf("a maximum of %d source images is supported for image edits", grokMediaMaxEditSourceImages) + } + out := body + var err error + for _, field := range []string{"image", "images", "mask"} { + out, err = rewriteGrokMediaJSONImageField(out, field) + if err != nil { + return nil, err + } + } + return out, nil +} + +func rewriteGrokMediaJSONImageField(body []byte, path string) ([]byte, error) { + value := gjson.GetBytes(body, path) + if !value.Exists() { + return body, nil + } + if value.IsArray() { + rewritten := make([]map[string]string, 0, len(value.Array())) + for _, item := range value.Array() { + imageURL := extractGrokMediaImageURL(item) + if imageURL == "" { + return body, nil + } + rewritten = append(rewritten, grokMediaImageObject(imageURL)) + } + out, err := sjson.SetBytes(body, path, rewritten) + if err != nil { + return nil, fmt.Errorf("rewrite grok media %s: %w", path, err) + } + return out, nil + } + imageURL := extractGrokMediaImageURL(value) + if imageURL == "" { + return body, nil + } + out, err := sjson.SetBytes(body, path, grokMediaImageObject(imageURL)) + if err != nil { + return nil, fmt.Errorf("rewrite grok media %s: %w", path, err) + } + return out, nil +} + func normalizeGrokMediaForwardBody(endpoint GrokMediaEndpoint, body []byte, contentType string) ([]byte, string, error) { if !endpoint.RequiresRequestBody() || !gjson.ValidBytes(body) { return body, contentType, nil @@ -773,9 +1133,11 @@ func NormalizeGrokMediaModelForEndpoint(endpoint GrokMediaEndpoint, model string return "grok-imagine-image-quality" } case GrokMediaEndpointVideosGenerations: - if model == "grok-imagine-video-1.5" && !hasInputImage { - return "grok-imagine-video" - } + // xAI's 1.5 model is image-to-video only. Keep the requested model + // unchanged when the image is missing so the upstream returns its + // documented invalid-argument response instead of silently switching + // models and pricing. + _ = hasInputImage } return model } @@ -783,6 +1145,8 @@ func NormalizeGrokMediaModelForEndpoint(endpoint GrokMediaEndpoint, model string type grokMediaUsageMetadata struct { ResponseID string Usage OpenAIUsage + Model string + BillingModel string ImageCount int ImageSize string ImageInputSize string @@ -802,12 +1166,24 @@ func grokMediaUsageFromResponse(endpoint GrokMediaEndpoint, requestInfo GrokMedi meta.ImageInputSize = requestInfo.Size meta.ImageOutputSizes = collectOpenAIResponseImageOutputSizesFromJSONBytes(responseBody) case GrokMediaEndpointVideosGenerations, GrokMediaEndpointVideosEdits, GrokMediaEndpointVideosExtensions: + // Async video: capture request_id + create-time pricing params only. + // Billable VideoCount is set later when status polling observes video.url. meta.ResponseID = extractGrokMediaVideoRequestID(responseBody) - meta.VideoCount = 1 meta.VideoResolution = requestInfo.Resolution meta.VideoDurationSeconds = requestInfo.DurationSeconds - // Keep the legacy media-unit counter populated for existing usage displays. - meta.ImageCount = 1 + case GrokMediaEndpointVideoStatus: + // Prefer status-body URL success + upstream duration/resolution when present. + if IsGrokVideoStatusBillable(responseBody) { + // provisional units; handler merges with pending snapshot before RecordUsage. + if billed := ExtractGrokVideoBillingFromStatusBody(responseBody, nil, ""); billed != nil { + meta.ResponseID = billed.ResponseID + meta.Model = billed.Model + meta.BillingModel = billed.BillingModel + meta.VideoCount = billed.VideoCount + meta.VideoResolution = billed.VideoResolution + meta.VideoDurationSeconds = billed.VideoDurationSeconds + } + } } return meta } @@ -816,7 +1192,7 @@ func extractGrokMediaVideoRequestID(body []byte) string { if len(body) == 0 || !gjson.ValidBytes(body) { return "" } - for _, path := range []string{"request_id", "id", "data.request_id", "data.id", "video.request_id", "video.id"} { + for _, path := range []string{"request_id", "id", "data.request_id", "data.id", "video.request_id", "video.id", "task_id", "data.task_id", "video.task_id"} { if id := strings.TrimSpace(gjson.GetBytes(body, path).String()); id != "" { return id } diff --git a/backend/internal/service/grok_media_video_billing_test.go b/backend/internal/service/grok_media_video_billing_test.go new file mode 100644 index 0000000000..84dc2c1079 --- /dev/null +++ b/backend/internal/service/grok_media_video_billing_test.go @@ -0,0 +1,149 @@ +package service + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestGrokVideoE2EDurationFromCreatedAt(t *testing.T) { + t.Parallel() + created := time.Now().UTC().Add(-45 * time.Second) + d := GrokVideoE2EDuration(created.Format(time.RFC3339Nano), time.Now().UTC()) + require.GreaterOrEqual(t, d, 44*time.Second) + require.LessOrEqual(t, d, 47*time.Second) + + require.Equal(t, time.Duration(0), GrokVideoE2EDuration("", time.Now())) + require.Equal(t, time.Duration(0), GrokVideoE2EDuration("not-a-time", time.Now())) + // Future CreatedAt clamps to zero (clock skew). + require.Equal(t, time.Duration(0), GrokVideoE2EDuration(time.Now().Add(time.Hour).UTC().Format(time.RFC3339Nano), time.Now())) +} + +func TestGrokVideoPendingCreatedAtStampOnStoreShape(t *testing.T) { + t.Parallel() + // GrokVideoPendingCreatedAtNow must be parseable by GrokVideoE2EDuration. + stamp := GrokVideoPendingCreatedAtNow() + require.NotEmpty(t, stamp) + d := GrokVideoE2EDuration(stamp, time.Now().UTC().Add(2*time.Second)) + require.GreaterOrEqual(t, d, time.Second) + require.LessOrEqual(t, d, 3*time.Second) +} + +func TestIsGrokVideoStatusBillable(t *testing.T) { + t.Parallel() + // Official success: status=done + video.url + require.True(t, IsGrokVideoStatusBillable([]byte(`{ + "status":"done", + "model":"grok-imagine-video-1.5", + "video":{"url":"https://vidgen.x.ai/x.mp4","duration":8,"respect_moderation":true} + }`))) + + // Official non-success states + require.False(t, IsGrokVideoStatusBillable(nil)) + require.False(t, IsGrokVideoStatusBillable([]byte(`{"status":"pending"}`))) + require.False(t, IsGrokVideoStatusBillable([]byte(`{"status":"expired"}`))) + require.False(t, IsGrokVideoStatusBillable([]byte(`{"status":"failed"}`))) + // done without video.url is not billable + require.False(t, IsGrokVideoStatusBillable([]byte(`{"status":"done"}`))) + // URL alone (legacy/non-official shapes) is not enough + require.False(t, IsGrokVideoStatusBillable([]byte(`{"url":"https://example.com/v.mp4"}`))) + require.False(t, IsGrokVideoStatusBillable([]byte(`{"download_url":"/v1/videos/task/content"}`))) + // "completed" is not the official enum value + require.False(t, IsGrokVideoStatusBillable([]byte(`{"status":"completed","video":{"url":"https://vidgen.x.ai/x.mp4"}}`))) +} + +func TestExtractGrokVideoBillingFromStatusBodyPrefersUpstreamParams(t *testing.T) { + t.Parallel() + pending := &GrokVideoPendingBilling{ + Model: "pending-model", + BillingModel: "pending-billing", + UpstreamModel: "pending-upstream", + VideoResolution: VideoBillingResolution720P, + VideoDurationSeconds: 8, + } + // Official completed body from docs.x.ai Video Generation. + body := []byte(`{ + "status":"done", + "model":"grok-imagine-video-1.5", + "video":{"url":"https://vidgen.x.ai/signed.mp4","duration":12,"respect_moderation":true} + }`) + result := ExtractGrokVideoBillingFromStatusBody(body, pending, "req-1") + require.NotNil(t, result) + require.Equal(t, 1, result.VideoCount) + require.Equal(t, "grok-imagine-video-1.5", result.Model) + // Resolution is not in official status response — use create-time request. + require.Equal(t, VideoBillingResolution720P, result.VideoResolution) + // Duration prefers official video.duration. + require.Equal(t, 12, result.VideoDurationSeconds) +} + +func TestExtractGrokVideoBillingFromStatusBodyFallsBackToPending(t *testing.T) { + t.Parallel() + pending := &GrokVideoPendingBilling{ + Model: "create-model", + BillingModel: "create-billing", + UpstreamModel: "create-upstream", + VideoResolution: VideoBillingResolution1080P, + VideoDurationSeconds: 10, + } + // done + video.url, but no model/duration in body. + body := []byte(`{"status":"done","video":{"url":"https://vidgen.x.ai/signed.mp4"}}`) + result := ExtractGrokVideoBillingFromStatusBody(body, pending, "req-2") + require.NotNil(t, result) + require.Equal(t, "create-billing", result.BillingModel) + require.Equal(t, "create-upstream", result.UpstreamModel) + require.Equal(t, VideoBillingResolution1080P, result.VideoResolution) + require.Equal(t, 10, result.VideoDurationSeconds) +} + +func TestExtractGrokVideoBillingRejectsNonDoneStatus(t *testing.T) { + t.Parallel() + pending := &GrokVideoPendingBilling{Model: "m", VideoDurationSeconds: 8, VideoResolution: "720p"} + require.Nil(t, ExtractGrokVideoBillingFromStatusBody( + []byte(`{"status":"pending","video":{"url":"https://vidgen.x.ai/x.mp4","duration":8}}`), + pending, "req", + )) + require.Nil(t, ExtractGrokVideoBillingFromStatusBody( + []byte(`{"status":"completed","video":{"url":"https://vidgen.x.ai/x.mp4","duration":8}}`), + pending, "req", + )) +} + +func TestGrokMediaUsageFromResponseVideoCreateDoesNotBill(t *testing.T) { + t.Parallel() + info := GrokMediaRequestInfo{Model: "grok-imagine-video", Resolution: "720p", DurationSeconds: 10} + meta := grokMediaUsageFromResponse(GrokMediaEndpointVideosGenerations, info, []byte(`{"request_id":"v1"}`)) + require.Equal(t, "v1", meta.ResponseID) + require.Equal(t, 0, meta.VideoCount) + require.Equal(t, 10, meta.VideoDurationSeconds) + require.Equal(t, VideoBillingResolution720P, meta.VideoResolution) +} + +func TestGrokMediaUsageFromResponseVideoStatusBillsOnOfficialDone(t *testing.T) { + t.Parallel() + meta := grokMediaUsageFromResponse( + GrokMediaEndpointVideoStatus, + GrokMediaRequestInfo{}, + []byte(`{"status":"done","model":"grok-imagine-video-1.5","video":{"url":"https://vidgen.x.ai/a.mp4","duration":9}}`), + ) + require.Equal(t, 1, meta.VideoCount) + require.Equal(t, 9, meta.VideoDurationSeconds) + require.Equal(t, "grok-imagine-video-1.5", meta.Model) + + // Official non-done must not set billable units. + pendingOnly := grokMediaUsageFromResponse( + GrokMediaEndpointVideoStatus, + GrokMediaRequestInfo{}, + []byte(`{"status":"pending"}`), + ) + require.Equal(t, 0, pendingOnly.VideoCount) + + // completed is not official done. + completed := grokMediaUsageFromResponse( + GrokMediaEndpointVideoStatus, + GrokMediaRequestInfo{}, + []byte(`{"status":"completed","video":{"url":"https://vidgen.x.ai/a.mp4","duration":9}}`), + ) + require.Equal(t, 0, completed.VideoCount) +} diff --git a/backend/internal/service/grok_model_quota_block.go b/backend/internal/service/grok_model_quota_block.go new file mode 100644 index 0000000000..0bc9cd40b8 --- /dev/null +++ b/backend/internal/service/grok_model_quota_block.go @@ -0,0 +1,114 @@ +package service + +import ( + "strconv" + "strings" + "sync" + "time" +) + +// Process-local per-account model soft-blocks for Grok free-usage that names a +// model (e.g. "used all free usage for model grok-4.5"). Sibling models on the +// same account stay schedulable. Multi-instance: each process learns from its +// own upstream errors. +type grokModelQuotaBlock struct { + Until time.Time +} + +type grokModelQuotaBlockStore struct { + mu sync.Mutex + items map[string]grokModelQuotaBlock // key: accountID|model +} + +var globalGrokModelQuotaBlocks = &grokModelQuotaBlockStore{ + items: make(map[string]grokModelQuotaBlock), +} + +const ( + grokModelQuotaBlockDefaultTTL = 2 * time.Hour + grokModelQuotaBlockMaxTTL = 6 * time.Hour + grokModelQuotaBlockMinTTL = 20 * time.Minute +) + +func grokModelQuotaBlockKey(accountID int64, model string) string { + return strings.TrimSpace(strings.ToLower(model)) + "|" + strconv.FormatInt(accountID, 10) +} + +// markGrokModelQuotaBlock soft-blocks accountID for model until the given time. +func markGrokModelQuotaBlock(accountID int64, model string, until time.Time) { + model = strings.TrimSpace(model) + if accountID <= 0 || model == "" || until.IsZero() { + return + } + now := time.Now() + if !until.After(now.Add(grokModelQuotaBlockMinTTL)) { + until = now.Add(grokModelQuotaBlockDefaultTTL) + } + if max := now.Add(grokModelQuotaBlockMaxTTL); until.After(max) { + until = max + } + key := grokModelQuotaBlockKey(accountID, model) + globalGrokModelQuotaBlocks.mu.Lock() + defer globalGrokModelQuotaBlocks.mu.Unlock() + if cur, ok := globalGrokModelQuotaBlocks.items[key]; ok && cur.Until.After(until) { + return + } + globalGrokModelQuotaBlocks.items[key] = grokModelQuotaBlock{Until: until} + for k, v := range globalGrokModelQuotaBlocks.items { + if !v.Until.After(now) { + delete(globalGrokModelQuotaBlocks.items, k) + } + } +} + +// isGrokModelQuotaBlocked reports whether this account cannot serve model now. +func isGrokModelQuotaBlocked(accountID int64, model string, now time.Time) bool { + model = strings.TrimSpace(model) + if accountID <= 0 || model == "" { + return false + } + key := grokModelQuotaBlockKey(accountID, model) + globalGrokModelQuotaBlocks.mu.Lock() + defer globalGrokModelQuotaBlocks.mu.Unlock() + cur, ok := globalGrokModelQuotaBlocks.items[key] + if !ok { + return false + } + if !cur.Until.After(now) { + delete(globalGrokModelQuotaBlocks.items, key) + return false + } + return true +} + +func filterGrokModelQuotaBlockedAccounts(accounts []Account, model string, now time.Time) []Account { + if len(accounts) == 0 || strings.TrimSpace(model) == "" { + return accounts + } + out := make([]Account, 0, len(accounts)) + for i := range accounts { + upstreamModel := canonicalOpenAIAccountSchedulingModel(&accounts[i], model) + if isGrokModelQuotaBlocked(accounts[i].ID, upstreamModel, now) { + continue + } + out = append(out, accounts[i]) + } + return out +} + +// isGrokModelSpecificFreeUsage is true when free-usage exhaustion is scoped to +// a named model (account may still serve other models). +func isGrokModelSpecificFreeUsage(low, model string) bool { + model = strings.ToLower(strings.TrimSpace(model)) + if model == "" || low == "" { + return false + } + if strings.Contains(low, "for model") || strings.Contains(low, "模型") { + return true + } + // "used all the included free usage for model grok-4.5" + if strings.Contains(low, "free usage") && strings.Contains(low, model) { + return true + } + return false +} diff --git a/backend/internal/service/grok_oauth_reconciliation_test.go b/backend/internal/service/grok_oauth_reconciliation_test.go index 8de04a1b5d..5449716447 100644 --- a/backend/internal/service/grok_oauth_reconciliation_test.go +++ b/backend/internal/service/grok_oauth_reconciliation_test.go @@ -107,6 +107,24 @@ func (r *grokReconcileRepo) UpdateCredentials(_ context.Context, id int64, crede return nil } +func (r *grokReconcileRepo) UpdateExtra(_ context.Context, id int64, updates map[string]any) error { + r.mu.Lock() + defer r.mu.Unlock() + for i := range r.accounts { + if r.accounts[i].ID != id { + continue + } + if r.accounts[i].Extra == nil { + r.accounts[i].Extra = make(map[string]any) + } + for key, value := range updates { + r.accounts[i].Extra[key] = value + } + break + } + return nil +} + func (r *grokReconcileRepo) SetError(_ context.Context, id int64, message string) error { r.mu.Lock() defer r.mu.Unlock() diff --git a/backend/internal/service/grok_oauth_service.go b/backend/internal/service/grok_oauth_service.go index 92f3937967..19120e84bd 100644 --- a/backend/internal/service/grok_oauth_service.go +++ b/backend/internal/service/grok_oauth_service.go @@ -8,6 +8,7 @@ import ( "strings" "time" + "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) @@ -18,14 +19,44 @@ type GrokOAuthService struct { sessionStore *xai.SessionStore proxyRepo ProxyRepository oauthClient GrokOAuthClient + config *config.Config } -func NewGrokOAuthService(proxyRepo ProxyRepository, oauthClient GrokOAuthClient) *GrokOAuthService { - return &GrokOAuthService{ +func NewGrokOAuthService(proxyRepo ProxyRepository, oauthClient GrokOAuthClient, configs ...*config.Config) *GrokOAuthService { + service := &GrokOAuthService{ sessionStore: xai.NewSessionStore(), proxyRepo: proxyRepo, oauthClient: oauthClient, } + if len(configs) > 0 { + service.config = configs[0] + } + return service +} + +// WithSessionStore replaces the in-memory OAuth session store (e.g. Redis-backed +// for cross-instance single-use callbacks). Redis wiring stays in Wire providers +// so this service package does not import go-redis (depguard). +func (s *GrokOAuthService) WithSessionStore(store *xai.SessionStore) *GrokOAuthService { + if s != nil && store != nil { + if s.sessionStore != nil { + s.sessionStore.Stop() + } + s.sessionStore = store + } + return s +} + +type GrokOAuthCapabilities struct { + PasswordAuthEnabled bool `json:"password_auth_enabled"` +} + +func (s *GrokOAuthService) GetCapabilities() GrokOAuthCapabilities { + return GrokOAuthCapabilities{PasswordAuthEnabled: s.passwordAuthEnabled()} +} + +func (s *GrokOAuthService) passwordAuthEnabled() bool { + return s.config != nil && s.config.Gateway.Grok.PasswordAuthEnabled } type GrokAuthURLResult struct { @@ -106,6 +137,13 @@ type GrokTokenInfo struct { EntitlementStatus string `json:"entitlement_status,omitempty"` } +// GrokPasswordLoginResult is an ephemeral password-login outcome. +// SSOToken is never persisted and must only feed ConvertSSOToBuild. +type GrokPasswordLoginResult struct { + Email string `json:"email,omitempty"` + SSOToken string `json:"sso_token"` +} + func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchangeCodeInput) (*GrokTokenInfo, error) { if input == nil { return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_INPUT", "input is required") @@ -114,7 +152,6 @@ func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchange if !ok { return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_SESSION_NOT_FOUND", "session not found or expired") } - defer s.sessionStore.Delete(input.SessionID) parsed := xai.ParseAuthorizationInput(input.Code) code := strings.TrimSpace(parsed.Code) @@ -125,12 +162,16 @@ func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchange if state == "" { state = strings.TrimSpace(parsed.State) } - if parsed.RequiresState && state == "" { - return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_STATE_REQUIRED", "oauth state is required for callback URLs") + if state == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_STATE_REQUIRED", "oauth state is required") } - if state != "" && subtle.ConstantTimeCompare([]byte(state), []byte(session.State)) != 1 { + if subtle.ConstantTimeCompare([]byte(state), []byte(session.State)) != 1 { return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_STATE", "invalid oauth state") } + if redirectURI := strings.TrimSpace(input.RedirectURI); redirectURI != "" && + redirectURI != strings.TrimSpace(session.RedirectURI) { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_REDIRECT_URI_MISMATCH", "redirect_uri does not match the OAuth session") + } proxyURL := session.ProxyURL if input.ProxyID != nil { @@ -140,27 +181,45 @@ func (s *GrokOAuthService) ExchangeCode(ctx context.Context, input *GrokExchange return nil, err } } - redirectURI := session.RedirectURI - if strings.TrimSpace(input.RedirectURI) != "" { - redirectURI = input.RedirectURI + if err := s.requireOAuthClient(); err != nil { + return nil, err } - - tokenResp, err := s.oauthClient.ExchangeCode(ctx, code, session.CodeVerifier, redirectURI, proxyURL, session.ClientID) + if !s.sessionStore.TryConsumeSession(input.SessionID) { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_SESSION_ALREADY_USED", "oauth session has already been used") + } + defer s.sessionStore.Delete(input.SessionID) + tokenResp, err := s.oauthClient.ExchangeCode(ctx, code, session.CodeVerifier, session.RedirectURI, proxyURL, session.ClientID) if err != nil { return nil, err } + if err := validateGrokTokenResponse(tokenResp); err != nil { + return nil, err + } return s.tokenInfoFromResponse(tokenResp, session.ClientID, nil), nil } +func (s *GrokOAuthService) requireOAuthClient() error { + if s == nil || s.oauthClient == nil { + return infraerrors.New(http.StatusInternalServerError, "GROK_OAUTH_CLIENT_NOT_CONFIGURED", "oauth client is not configured") + } + return nil +} + func (s *GrokOAuthService) RefreshToken(ctx context.Context, refreshToken, proxyURL, clientID string) (*GrokTokenInfo, error) { refreshToken = strings.TrimSpace(refreshToken) if refreshToken == "" { return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_NO_REFRESH_TOKEN", "refresh_token is required") } + if err := s.requireOAuthClient(); err != nil { + return nil, err + } tokenResp, err := s.oauthClient.RefreshToken(ctx, refreshToken, proxyURL, clientID) if err != nil { return nil, err } + if err := validateGrokTokenResponse(tokenResp); err != nil { + return nil, err + } tokenInfo := s.tokenInfoFromResponse(tokenResp, clientID, nil) if tokenInfo.RefreshToken == "" { tokenInfo.RefreshToken = refreshToken @@ -176,7 +235,16 @@ func (s *GrokOAuthService) ValidateRefreshToken(ctx context.Context, refreshToke return s.RefreshToken(ctx, refreshToken, proxyURL, xai.EffectiveClientID()) } -func (s *GrokOAuthService) ConvertFromSSO(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) { +// ValidateSSOToken converts a Web SSO cookie into Build OAuth tokens. +// The raw sso_token is never stored on GrokTokenInfo or account credentials. +func (s *GrokOAuthService) ValidateSSOToken(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) { + ssoToken = strings.TrimSpace(ssoToken) + if ssoToken == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_NO_SSO_TOKEN", "sso_token is required") + } + if err := s.requireOAuthClient(); err != nil { + return nil, err + } proxyURL, err := s.proxyURL(ctx, proxyID) if err != nil { return nil, err @@ -185,9 +253,61 @@ func (s *GrokOAuthService) ConvertFromSSO(ctx context.Context, ssoToken string, if err != nil { return nil, err } + if err := validateGrokTokenResponse(tokenResp); err != nil { + return nil, err + } return s.tokenInfoFromResponse(tokenResp, xai.DefaultClientID, nil), nil } +// ConvertFromSSO is the batch-import entry point; same semantics as ValidateSSOToken. +func (s *GrokOAuthService) ConvertFromSSO(ctx context.Context, ssoToken string, proxyID *int64) (*GrokTokenInfo, error) { + return s.ValidateSSOToken(ctx, ssoToken, proxyID) +} + +// AuthorizePassword logs in with email/password, converts the resulting SSO cookie +// to Build OAuth, and returns OAuth tokens only. Password and raw SSO are never persisted. +func (s *GrokOAuthService) AuthorizePassword(ctx context.Context, email, password string, proxyID *int64) (*GrokTokenInfo, error) { + if !s.passwordAuthEnabled() { + return nil, infraerrors.New(http.StatusForbidden, "GROK_OAUTH_PASSWORD_AUTH_DISABLED", "Grok password authorization is disabled") + } + email = strings.TrimSpace(email) + if email == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_EMAIL_REQUIRED", "email is required") + } + if strings.TrimSpace(password) == "" { + return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_PASSWORD_REQUIRED", "password is required") + } + if err := s.requireOAuthClient(); err != nil { + return nil, err + } + proxyURL, err := s.proxyURL(ctx, proxyID) + if err != nil { + return nil, err + } + loginResult, err := s.oauthClient.LoginWithPassword(ctx, email, password, proxyURL) + if err != nil { + return nil, err + } + if loginResult == nil || strings.TrimSpace(loginResult.SSOToken) == "" { + return nil, infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_PASSWORD_LOGIN_FAILED", "grok password login did not return sso_token") + } + info, err := s.ValidateSSOToken(ctx, loginResult.SSOToken, proxyID) + if err != nil { + return nil, err + } + if strings.TrimSpace(info.Email) == "" { + info.Email = loginResult.Email + } + return info, nil +} + +func validateGrokTokenResponse(tokenResp *xai.TokenResponse) error { + if tokenResp == nil || strings.TrimSpace(tokenResp.AccessToken) == "" { + return infraerrors.New(http.StatusBadGateway, "GROK_OAUTH_INVALID_TOKEN_RESPONSE", "grok oauth token response missing access_token") + } + return nil +} + func (s *GrokOAuthService) RefreshAccountToken(ctx context.Context, account *Account) (*GrokTokenInfo, error) { if account == nil || account.Platform != PlatformGrok { return nil, infraerrors.New(http.StatusBadRequest, "GROK_OAUTH_INVALID_ACCOUNT", "account is not a Grok account") diff --git a/backend/internal/service/grok_oauth_service_test.go b/backend/internal/service/grok_oauth_service_test.go index 54baef03a2..cbbd802e01 100644 --- a/backend/internal/service/grok_oauth_service_test.go +++ b/backend/internal/service/grok_oauth_service_test.go @@ -9,25 +9,37 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/stretchr/testify/require" ) type grokOAuthClientStub struct { - refreshResponse *xai.TokenResponse - ssoResponse *xai.TokenResponse - exchangeCalls int + refreshResponse *xai.TokenResponse + ssoResponse *xai.TokenResponse + loginResult *GrokPasswordLoginResult + loginEmail string + loginPassword string + exchangeCalls int + exchangeRedirectURI string } -func (s *grokOAuthClientStub) ExchangeCode(context.Context, string, string, string, string, string) (*xai.TokenResponse, error) { +func (s *grokOAuthClientStub) ExchangeCode(_ context.Context, _, _, redirectURI, _, _ string) (*xai.TokenResponse, error) { s.exchangeCalls++ - return &xai.TokenResponse{}, nil + s.exchangeRedirectURI = redirectURI + return &xai.TokenResponse{AccessToken: "access-token"}, nil } func (s *grokOAuthClientStub) RefreshToken(context.Context, string, string, string) (*xai.TokenResponse, error) { return s.refreshResponse, nil } +func (s *grokOAuthClientStub) LoginWithPassword(_ context.Context, email, password, _ string) (*GrokPasswordLoginResult, error) { + s.loginEmail = email + s.loginPassword = password + return s.loginResult, nil +} + func (s *grokOAuthClientStub) ConvertSSOToBuild(context.Context, string, string) (*xai.TokenResponse, error) { return s.ssoResponse, nil } @@ -49,7 +61,19 @@ func TestGrokOAuthServiceRefreshTokenPreservesOriginalRefreshTokenWhenNotRotated require.Equal(t, "client-id", info.ClientID) } -func TestGrokOAuthServiceExchangeCodeRequiresStateForCallbackURLAndConsumesSession(t *testing.T) { +func TestGrokOAuthServiceRefreshTokenRejectsEmptyUpstreamResponse(t *testing.T) { + svc := NewGrokOAuthService(nil, &grokOAuthClientStub{}) + defer svc.Stop() + + require.NotPanics(t, func() { + info, err := svc.RefreshToken(context.Background(), "refresh-token", "", "client-id") + require.Nil(t, info) + require.Error(t, err) + require.Contains(t, err.Error(), "GROK_OAUTH_INVALID_TOKEN_RESPONSE") + }) +} + +func TestGrokOAuthServiceExchangeCodeConsumesOnlyAfterValidation(t *testing.T) { client := &grokOAuthClientStub{} svc := NewGrokOAuthService(nil, client) defer svc.Stop() @@ -70,9 +94,92 @@ func TestGrokOAuthServiceExchangeCodeRequiresStateForCallbackURLAndConsumesSessi Code: "code-with-state", State: auth.State, }) + require.NoError(t, err) + require.Equal(t, 1, client.exchangeCalls) + + _, err = svc.ExchangeCode(context.Background(), &GrokExchangeCodeInput{ + SessionID: auth.SessionID, + Code: "replayed-code", + State: auth.State, + }) require.Error(t, err) require.Contains(t, err.Error(), "GROK_OAUTH_SESSION_NOT_FOUND") + require.Equal(t, 1, client.exchangeCalls) +} + +func TestGrokOAuthServiceExchangeCodeRejectsMissingClientWithoutConsumingSession(t *testing.T) { + svc := NewGrokOAuthService(nil, nil) + defer svc.Stop() + auth, err := svc.GenerateAuthURL(context.Background(), nil, "") + require.NoError(t, err) + + _, err = svc.ExchangeCode(context.Background(), &GrokExchangeCodeInput{ + SessionID: auth.SessionID, + Code: "code", + State: auth.State, + }) + require.Error(t, err) + require.Contains(t, err.Error(), "GROK_OAUTH_CLIENT_NOT_CONFIGURED") + _, ok := svc.sessionStore.Get(auth.SessionID) + require.True(t, ok) +} + +func TestGrokOAuthServiceExchangeCodeRequiresStateForBareCode(t *testing.T) { + client := &grokOAuthClientStub{} + svc := NewGrokOAuthService(nil, client) + defer svc.Stop() + auth, err := svc.GenerateAuthURL(context.Background(), nil, "") + require.NoError(t, err) + + _, err = svc.ExchangeCode(context.Background(), &GrokExchangeCodeInput{ + SessionID: auth.SessionID, + Code: "bare-authorization-code", + }) + require.Error(t, err) + require.Contains(t, err.Error(), "GROK_OAUTH_STATE_REQUIRED") require.Zero(t, client.exchangeCalls) + _, ok := svc.sessionStore.Get(auth.SessionID) + require.True(t, ok) +} + +func TestGrokOAuthServiceExchangeCodeRejectsRedirectURIOverride(t *testing.T) { + client := &grokOAuthClientStub{} + svc := NewGrokOAuthService(nil, client) + defer svc.Stop() + auth, err := svc.GenerateAuthURL(context.Background(), nil, "") + require.NoError(t, err) + + _, err = svc.ExchangeCode(context.Background(), &GrokExchangeCodeInput{ + SessionID: auth.SessionID, + Code: "authorization-code", + State: auth.State, + RedirectURI: "http://127.0.0.1:9999/callback", + }) + require.Error(t, err) + require.Contains(t, err.Error(), "GROK_OAUTH_REDIRECT_URI_MISMATCH") + require.Zero(t, client.exchangeCalls) + + _, err = svc.ExchangeCode(context.Background(), &GrokExchangeCodeInput{ + SessionID: auth.SessionID, + Code: "authorization-code", + State: auth.State, + RedirectURI: xai.DefaultRedirectURI, + }) + require.NoError(t, err) + require.Equal(t, xai.DefaultRedirectURI, client.exchangeRedirectURI) +} + +func TestGrokOAuthServiceExternalFlowsRejectMissingClient(t *testing.T) { + svc := NewGrokOAuthService(nil, nil) + defer svc.Stop() + + _, err := svc.RefreshToken(context.Background(), "refresh-token", "", "") + require.Error(t, err) + require.Contains(t, err.Error(), "GROK_OAUTH_CLIENT_NOT_CONFIGURED") + + _, err = svc.ValidateSSOToken(context.Background(), "sso-token", nil) + require.Error(t, err) + require.Contains(t, err.Error(), "GROK_OAUTH_CLIENT_NOT_CONFIGURED") } func TestGrokOAuthServiceBuildAccountCredentialsDefaultsToSubscriptionProxy(t *testing.T) { @@ -108,6 +215,69 @@ func TestGrokOAuthServiceConvertFromSSOExtractsBuildClaims(t *testing.T) { require.Equal(t, "user@example.com", credentials["email"]) require.Equal(t, "user-sub", credentials["sub"]) require.Equal(t, "team-1", credentials["team_id"]) + require.NotContains(t, credentials, "sso_token") +} + +func TestGrokOAuthServiceValidateSSOTokenReturnsOAuthTokensWithoutPersistingSSO(t *testing.T) { + svc := NewGrokOAuthService(nil, &grokOAuthClientStub{ + ssoResponse: &xai.TokenResponse{ + AccessToken: "access-from-sso", + RefreshToken: "refresh-from-sso", + TokenType: "Bearer", + ExpiresIn: 3600, + }, + }) + defer svc.Stop() + + info, err := svc.ValidateSSOToken(context.Background(), "sso-token", nil) + require.NoError(t, err) + require.Equal(t, "access-from-sso", info.AccessToken) + require.Equal(t, "refresh-from-sso", info.RefreshToken) + + creds := svc.BuildAccountCredentials(info) + require.NotContains(t, creds, "sso_token") + require.NotContains(t, creds, "password") +} + +func TestGrokOAuthServiceAuthorizePasswordUsesLoginThenSSOAuthorize(t *testing.T) { + client := &grokOAuthClientStub{ + loginResult: &GrokPasswordLoginResult{ + Email: "user@example.com", + SSOToken: "password-derived-sso", + }, + ssoResponse: &xai.TokenResponse{ + AccessToken: "access-from-password", + RefreshToken: "refresh-from-password", + ExpiresIn: 3600, + }, + } + cfg := &config.Config{} + cfg.Gateway.Grok.PasswordAuthEnabled = true + svc := NewGrokOAuthService(nil, client, cfg) + defer svc.Stop() + + require.True(t, svc.GetCapabilities().PasswordAuthEnabled) + info, err := svc.AuthorizePassword(context.Background(), " user@example.com ", " super-secret ", nil) + require.NoError(t, err) + require.Equal(t, "user@example.com", info.Email) + require.Equal(t, "access-from-password", info.AccessToken) + creds := svc.BuildAccountCredentials(info) + require.NotContains(t, creds, "password") + require.NotContains(t, creds, "sso_token") + require.Equal(t, "user@example.com", client.loginEmail) + require.Equal(t, " super-secret ", client.loginPassword) +} + +func TestGrokOAuthServiceAuthorizePasswordDisabledByDefault(t *testing.T) { + client := &grokOAuthClientStub{} + svc := NewGrokOAuthService(nil, client) + defer svc.Stop() + + require.False(t, svc.GetCapabilities().PasswordAuthEnabled) + _, err := svc.AuthorizePassword(context.Background(), "user@example.com", "secret", nil) + require.Error(t, err) + require.Contains(t, err.Error(), "GROK_OAUTH_PASSWORD_AUTH_DISABLED") + require.Empty(t, client.loginEmail) } func makeGrokOAuthJWT(claims map[string]any) string { diff --git a/backend/internal/service/grok_observed_models.go b/backend/internal/service/grok_observed_models.go new file mode 100644 index 0000000000..480ba457bd --- /dev/null +++ b/backend/internal/service/grok_observed_models.go @@ -0,0 +1,198 @@ +package service + +import ( + "context" + "encoding/json" + "io" + "log/slog" + "net/http" + "strings" + "sync" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/tidwall/gjson" +) + +const ( + grokObservedModelsExtraKey = "grok_observed_models" + grokObservedModelsTTL = 6 * time.Hour + grokObservedModelsTimeout = 15 * time.Second +) + +type grokObservedModelsSnapshot struct { + Models []string `json:"models"` + FetchedAt string `json:"fetched_at"` + Source string `json:"source,omitempty"` +} + +var grokObservedModelsFlight sync.Map // accountID -> *singleflight-ish in-flight + +// scheduleGrokObservedModelsSync best-effort fetches upstream /v1/models for a +// Grok OAuth account and stores IDs in Extra. Never blocks request path long; +// callers should fire-and-forget after successful auth/probe. +func (s *GrokQuotaService) scheduleGrokObservedModelsSync(account *Account) { + if s == nil || account == nil || !account.IsGrokOAuth() || s.accountRepo == nil { + return + } + id := account.ID + if _, loaded := grokObservedModelsFlight.LoadOrStore(id, struct{}{}); loaded { + return + } + // Copy credentials for background use. + acc := *account + go func() { + defer grokObservedModelsFlight.Delete(id) + ctx, cancel := context.WithTimeout(context.Background(), grokObservedModelsTimeout) + defer cancel() + if err := s.syncGrokObservedModels(ctx, &acc); err != nil { + slog.Debug("grok_observed_models_sync_failed", "account_id", id, "error", err) + } + }() +} + +func (s *GrokQuotaService) syncGrokObservedModels(ctx context.Context, account *Account) error { + if s == nil || account == nil { + return nil + } + // Skip if snapshot is still fresh. + if snap := parseGrokObservedModels(account.Extra); snap != nil { + if t, err := time.Parse(time.RFC3339, snap.FetchedAt); err == nil && time.Since(t) < grokObservedModelsTTL { + return nil + } + } + token := strings.TrimSpace(account.GetGrokAccessToken()) + if token == "" && s.tokenProvider != nil { + // Best-effort warm; avoid forcing refresh storms. + if at, err := s.tokenProvider.GetAccessToken(ctx, account); err == nil { + token = strings.TrimSpace(at) + } + } + if token == "" { + return nil + } + baseURL := strings.TrimSpace(account.GetGrokBaseURL()) + if s.settingService != nil { + baseURL = strings.TrimSpace(s.settingService.ResolveGrokBaseURL(ctx, account)) + } + if baseURL == "" { + baseURL = xai.DefaultCLIBaseURL + } + validator, err := grokBaseURLValidator(account, s.cfg) + if err != nil { + return err + } + validatedBaseURL, err := validator(baseURL) + if err != nil { + return err + } + req, err := http.NewRequestWithContext(ctx, http.MethodGet, buildOpenAIModelsURL(validatedBaseURL), nil) + if err != nil { + return err + } + req.Header.Set("Authorization", "Bearer "+token) + req.Header.Set("Accept", "application/json") + req.Header.Set("User-Agent", grokUpstreamUserAgent) + if account.IsGrokOAuth() { + applyGrokCLIHeaders(req.Header) + if isGrokCLIProxyTarget(req.URL.String()) { + if userID := strings.TrimSpace(account.GetCredential("sub")); userID != "" { + req.Header.Set("X-UserID", userID) + } + if email := strings.TrimSpace(account.GetCredential("email")); email != "" { + req.Header.Set("X-Email", email) + } + } + } + account.ApplyHeaderOverrides(req.Header) + + proxyURL := "" + if s.proxyRepo != nil && account.ProxyID != nil { + if p, err := s.proxyRepo.GetByID(ctx, *account.ProxyID); err == nil && p != nil { + proxyURL = p.URL() + } + } + if s.httpUpstream == nil { + return nil + } + resp, err := s.httpUpstream.Do(req, proxyURL, account.ID, account.Concurrency) + if err != nil { + return err + } + defer func() { _ = resp.Body.Close() }() + body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) + if err != nil { + return err + } + if resp.StatusCode >= 400 { + return nil + } + ids := extractGrokModelIDsFromModelsBody(body) + if len(ids) == 0 { + return nil + } + snap := grokObservedModelsSnapshot{ + Models: ids, + FetchedAt: time.Now().UTC().Format(time.RFC3339), + Source: "upstream_v1_models", + } + raw, err := json.Marshal(snap) + if err != nil { + return err + } + var asMap map[string]any + if err := json.Unmarshal(raw, &asMap); err != nil { + return err + } + return s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + grokObservedModelsExtraKey: asMap, + }) +} + +func extractGrokModelIDsFromModelsBody(body []byte) []string { + data := gjson.GetBytes(body, "data") + if !data.IsArray() { + // Some gateways return a bare array. + data = gjson.ParseBytes(body) + } + seen := make(map[string]struct{}) + var out []string + data.ForEach(func(_, v gjson.Result) bool { + id := strings.TrimSpace(v.Get("id").String()) + if id == "" { + id = strings.TrimSpace(v.String()) + } + if id == "" { + return true + } + if _, ok := seen[id]; ok { + return true + } + seen[id] = struct{}{} + out = append(out, id) + return true + }) + return out +} + +func parseGrokObservedModels(extra map[string]any) *grokObservedModelsSnapshot { + if extra == nil { + return nil + } + raw, ok := extra[grokObservedModelsExtraKey] + if !ok || raw == nil { + return nil + } + b, err := json.Marshal(raw) + if err != nil { + return nil + } + var snap grokObservedModelsSnapshot + if err := json.Unmarshal(b, &snap); err != nil { + return nil + } + if len(snap.Models) == 0 { + return nil + } + return &snap +} diff --git a/backend/internal/service/grok_p2_test.go b/backend/internal/service/grok_p2_test.go new file mode 100644 index 0000000000..2be382d26e --- /dev/null +++ b/backend/internal/service/grok_p2_test.go @@ -0,0 +1,102 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestGrokModelQuotaBlock_FiltersOnlyNamedModel(t *testing.T) { + id := time.Now().UnixNano()%1_000_000 + 5000 + markGrokModelQuotaBlock(id, "grok-4.5", time.Now().Add(time.Hour)) + now := time.Now() + require.True(t, isGrokModelQuotaBlocked(id, "grok-4.5", now)) + require.False(t, isGrokModelQuotaBlocked(id, "grok-4.3", now)) + + accounts := []Account{ + {ID: id, Platform: PlatformGrok, Type: AccountTypeOAuth}, + {ID: id + 1, Platform: PlatformGrok, Type: AccountTypeOAuth}, + } + filtered := filterGrokModelQuotaBlockedAccounts(accounts, "grok-4.5", now) + require.Len(t, filtered, 1) + require.Equal(t, id+1, filtered[0].ID) +} + +func TestGrokModelQuotaBlockFiltersMappedUpstreamModel(t *testing.T) { + id := time.Now().UnixNano()%1_000_000 + 7000 + markGrokModelQuotaBlock(id, "grok-4.5", time.Now().Add(time.Hour)) + account := Account{ + ID: id, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "model_mapping": map[string]any{"gpt-*": "grok-4.5"}, + }, + } + + require.Empty(t, filterGrokModelQuotaBlockedAccounts([]Account{account}, "gpt-5", time.Now())) +} + +func TestIsGrokModelSpecificFreeUsage(t *testing.T) { + require.True(t, isGrokModelSpecificFreeUsage( + "you've used all the included free usage for model grok-4.5", "grok-4.5")) + require.True(t, isGrokModelSpecificFreeUsage("模型额度用完 grok-4.3", "grok-4.3")) + require.False(t, isGrokModelSpecificFreeUsage("free usage exhausted", "grok-4.5")) +} + +func TestGrokStickyAffinitySeed_ScopesByModel(t *testing.T) { + a := grokStickyAffinitySeed("session-1", []byte(`{"model":"grok-4.5"}`)) + b := grokStickyAffinitySeed("session-1", []byte(`{"model":"grok-4.3"}`)) + c := grokStickyAffinitySeed("session-1", []byte(`{"model":"grok-4.5"}`)) + require.NotEqual(t, a, b) + require.Equal(t, a, c) + require.Contains(t, a, "grok-affinity:v1:") +} + +func TestExtractGrokModelIDsFromModelsBody(t *testing.T) { + body := []byte(`{"object":"list","data":[{"id":"grok-4.5"},{"id":"grok-4.3"},{"id":"grok-4.5"}]}`) + ids := extractGrokModelIDsFromModelsBody(body) + require.Equal(t, []string{"grok-4.5", "grok-4.3"}, ids) +} + +func TestAccountGrokNeedsReauth(t *testing.T) { + require.False(t, accountGrokNeedsReauth(nil)) + require.True(t, accountGrokNeedsReauth(&Account{ + Extra: map[string]any{"grok_needs_reauth": true}, + })) + require.True(t, accountGrokNeedsReauth(&Account{ + Status: StatusError, + ErrorMessage: "Grok spending limit reached; reauthorize or wait for billing reset", + })) +} + +func TestApplyGrokUpstreamFailure_ModelSpecificFreeUsage(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 9109, Platform: PlatformGrok, Type: AccountTypeOAuth} + body := []byte(`{"error":{"code":"subscription:free-usage-exhausted","message":"You've used all the included free usage for model grok-4.5. Usage resets over a rolling 24-hour window."}}`) + + svc.handleGrokAccountUpstreamError(context.Background(), account, 400, nil, body) + + require.Zero(t, repo.tempUnschedCalls, "model-scoped free usage must not cool sibling models") + require.True(t, isGrokModelQuotaBlocked(account.ID, "grok-4.5", time.Now())) + require.False(t, isGrokModelQuotaBlocked(account.ID, "grok-4.3", time.Now())) +} + +func TestApplyGrokUpstreamFailure_SpendingLimitRemainsRecoverable(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 9110, Platform: PlatformGrok, Type: AccountTypeOAuth} + body := []byte(`{"code":"personal-team-blocked:spending-limit","error":"spending limit reached"}`) + + svc.handleGrokAccountUpstreamError(context.Background(), account, 403, nil, body) + + require.Equal(t, 1, repo.rateLimitedCalls) + require.Zero(t, repo.tempUnschedCalls) + // Without a billing-period snapshot, use a short recoverable probe cooldown. + require.WithinDuration(t, time.Now().Add(grokSpendingLimitProbeCooldown), repo.lastRateLimitResetAt, 2*time.Second) +} diff --git a/backend/internal/service/grok_quota_fetcher.go b/backend/internal/service/grok_quota_fetcher.go index 545b1c9e97..4979dfd0b9 100644 --- a/backend/internal/service/grok_quota_fetcher.go +++ b/backend/internal/service/grok_quota_fetcher.go @@ -36,6 +36,7 @@ func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo { activeProbeClearsForbidden := newerSuccessfulGrokActiveProbeClearsBillingForbidden(billing, snapshot) if billing != nil { usage.GrokBilling = billing + applyGrokBillingProgressWindows(usage, billing, now) if billing.Plan != "" { usage.SubscriptionTier = billing.Plan usage.SubscriptionTierRaw = billing.Plan @@ -59,6 +60,14 @@ func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo { case 429: usage.ErrorCode = "rate_limited" } + // Official weekly/monthly progress clears the "unknown until headers" state. + if usage.ErrorCode == "quota_unknown" && (usage.SevenDay != nil || usage.ThirtyDay != nil) { + usage.ErrorCode = "" + if strings.Contains(strings.ToLower(usage.Error), "unknown until") || + strings.Contains(strings.ToLower(usage.Error), "no xai quota headers") { + usage.Error = "" + } + } } if err != nil || snapshot == nil { @@ -124,6 +133,12 @@ func (f *GrokQuotaFetcher) BuildUsageInfo(account *Account) *UsageInfo { usage.ErrorCode = "rate_limited" } } + if accountGrokNeedsReauth(account) { + usage.NeedsReauth = true + if usage.ErrorCode == "" { + usage.ErrorCode = "spending_limit" + } + } applyGrokCredentialUsageFallback(usage, account) if activeProbeClearsForbidden && strings.TrimSpace(snapshot.EntitlementStatus) == "" && strings.EqualFold(strings.TrimSpace(usage.GrokEntitlementStatus), "forbidden") { @@ -240,3 +255,48 @@ func grokQuotaSnapshotFromExtra(extra map[string]any) (*xai.QuotaSnapshot, error return &out, nil } } + +// applyGrokBillingProgressWindows fills official weekly (seven_day) and monthly +// (thirty_day) UsageProgress from a billing probe summary. +func applyGrokBillingProgressWindows(usage *UsageInfo, billing *xai.BillingSummary, now time.Time) { + if usage == nil || billing == nil { + return + } + if billing.UsagePercent != nil { + seven := &UsageProgress{Utilization: *billing.UsagePercent} + if end, err := parseTime(strings.TrimSpace(billing.PeriodEnd)); err == nil { + seven.ResetsAt = &end + if sec := int(end.Sub(now).Seconds()); sec > 0 { + seven.RemainingSeconds = sec + } + } + if usage.SevenDay != nil { + seven.WindowStats = usage.SevenDay.WindowStats + } + usage.SevenDay = seven + } + var monthlyUtil *float64 + if billing.UsedPercent != nil { + monthlyUtil = billing.UsedPercent + } else if billing.MonthlyLimitCents != nil && *billing.MonthlyLimitCents > 0 && billing.UsedCents != nil { + v := (*billing.UsedCents / *billing.MonthlyLimitCents) * 100 + monthlyUtil = &v + } + if monthlyUtil != nil { + thirty := &UsageProgress{Utilization: *monthlyUtil} + endRaw := strings.TrimSpace(billing.BillingPeriodEnd) + if endRaw == "" && billing.PeriodType == "monthly" { + endRaw = strings.TrimSpace(billing.PeriodEnd) + } + if end, err := parseTime(endRaw); err == nil { + thirty.ResetsAt = &end + if sec := int(end.Sub(now).Seconds()); sec > 0 { + thirty.RemainingSeconds = sec + } + } + if usage.ThirtyDay != nil { + thirty.WindowStats = usage.ThirtyDay.WindowStats + } + usage.ThirtyDay = thirty + } +} diff --git a/backend/internal/service/grok_quota_service.go b/backend/internal/service/grok_quota_service.go index 2c703f9f10..6658888a59 100644 --- a/backend/internal/service/grok_quota_service.go +++ b/backend/internal/service/grok_quota_service.go @@ -50,13 +50,14 @@ type GrokQuotaResetResult struct { } type GrokQuotaService struct { - accountRepo AccountRepository - proxyRepo ProxyRepository - tokenProvider *GrokTokenProvider - httpUpstream HTTPUpstream - usageLogRepo UsageLogRepository - cfg *config.Config - probeFlight singleflight.Group + accountRepo AccountRepository + proxyRepo ProxyRepository + tokenProvider *GrokTokenProvider + httpUpstream HTTPUpstream + usageLogRepo UsageLogRepository + settingService *SettingService + cfg *config.Config + probeFlight singleflight.Group } func NewGrokQuotaService( @@ -81,11 +82,20 @@ func NewGrokQuotaService( } } +func (s *GrokQuotaService) SetSettingService(settingService *SettingService) { + if s != nil { + s.settingService = settingService + } +} + // QueryQuota combines xAI billing data with an active quota-header probe for // Free accounts, whose billing response does not include usage_percent. func (s *GrokQuotaService) QueryQuota(ctx context.Context, accountID int64) (*GrokQuotaProbeResult, error) { billingResult, billingErr := s.ProbeBilling(ctx, accountID) if billingErr == nil && billingResult != nil && grokBillingHasAuthoritativeQuota(billingResult.Billing) { + if acc, err := s.accountRepo.GetByID(ctx, accountID); err == nil { + s.scheduleGrokObservedModelsSync(acc) + } return billingResult, nil } @@ -111,6 +121,9 @@ func (s *GrokQuotaService) QueryQuota(ctx context.Context, accountID int64) (*Gr probeResult.LocalUsageMonthly = billingResult.LocalUsageMonthly probeResult.Persisted = probeResult.Persisted || billingResult.Persisted } + if acc, err := s.accountRepo.GetByID(ctx, accountID); err == nil { + s.scheduleGrokObservedModelsSync(acc) + } return probeResult, nil } @@ -141,7 +154,7 @@ func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*Gr if err != nil { return nil, infraerrors.Newf(http.StatusBadRequest, "GROK_QUOTA_PROBE_BODY_ERROR", "failed to build probe body: %v", err) } - targetURL, err := buildGrokResponsesURL(account, s.cfg) + targetURL, err := buildGrokResponsesURL(account, s.cfg, s.settingService) if err != nil { return nil, infraerrors.Newf(http.StatusBadRequest, "GROK_QUOTA_BASE_URL_INVALID", "invalid Grok base_url: %v", err) } @@ -172,9 +185,21 @@ func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*Gr if limited { normalizeGrokExhaustedWindowResets(snapshot, resetAt, time.Now()) } - persistErr := s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ - grokQuotaSnapshotExtraKey: snapshot, - }) + // A failed probe must not erase a previously observed snapshot. 401/403 and + // transport/server errors commonly carry no quota headers; only successful + // responses, or 429 responses with useful rate-limit headers, are safe to + // persist. A successful 200 with no headers is still persisted as an + // explicit "no headers" observation so the UI can distinguish it from never + // probed. + persistErr := error(nil) + persisted := false + shouldPersist := resp.StatusCode < 400 || resp.StatusCode == http.StatusTooManyRequests + if shouldPersist && (snapshot.HeadersObserved || resp.StatusCode == http.StatusOK) { + persistErr = s.accountRepo.UpdateExtra(ctx, account.ID, map[string]any{ + grokQuotaSnapshotExtraKey: snapshot, + }) + persisted = persistErr == nil + } if limited { persistGrokRateLimit(ctx, s.accountRepo, account, resetAt) } else if isSuccessfulGrokRateLimitRecovery(account, snapshot) { @@ -189,7 +214,7 @@ func (s *GrokQuotaService) probeUsage(ctx context.Context, accountID int64) (*Gr HeadersObserved: snapshot.HeadersObserved, ResetSupported: false, FetchedAt: time.Now().Unix(), - Persisted: persistErr == nil, + Persisted: persisted, } if resp.StatusCode == http.StatusTooManyRequests { return result, nil diff --git a/backend/internal/service/grok_quota_service_test.go b/backend/internal/service/grok_quota_service_test.go index 9ffb70f000..70ee2c8d72 100644 --- a/backend/internal/service/grok_quota_service_test.go +++ b/backend/internal/service/grok_quota_service_test.go @@ -15,6 +15,7 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" @@ -86,6 +87,60 @@ func (r *grokQuotaAccountRepo) SetTempUnschedulable(_ context.Context, id int64, return nil } +func TestSyncGrokObservedModelsRejectsOAuthCustomURLOutsideOperatorPolicy(t *testing.T) { + account := &Account{ + ID: 901, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "secret-token", + "base_url": "https://blocked.example.test/v1", + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &httpUpstreamRecorder{} + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = true + cfg.Security.URLAllowlist.UpstreamHosts = []string{"allowed.example.test"} + svc := &GrokQuotaService{accountRepo: repo, httpUpstream: upstream, cfg: cfg} + + err := svc.syncGrokObservedModels(context.Background(), account) + require.ErrorContains(t, err, "base URL rejected by URL security policy") + require.Nil(t, upstream.lastReq) +} + +func TestSyncGrokObservedModelsUsesCLIIdentityAndAccountHeaders(t *testing.T) { + account := &Account{ + ID: 902, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "secret-token", + "sub": "user-902", + "email": "user902@example.test", + }, + } + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(`{"data":[{"id":"grok-4.5"}]}`)), + }} + svc := &GrokQuotaService{accountRepo: repo, httpUpstream: upstream, cfg: &config.Config{}} + + require.NoError(t, svc.syncGrokObservedModels(context.Background(), account)) + require.Equal(t, xai.DefaultCLIBaseURL+"/models", upstream.lastReq.URL.String()) + require.NotEmpty(t, upstream.lastReq.Header.Get("x-grok-client-version")) + require.Equal(t, xai.CLIClientIdentifier, upstream.lastReq.Header.Get("x-grok-client-identifier")) + require.Equal(t, "interactive", upstream.lastReq.Header.Get("X-Grok-Client-Mode")) + require.Equal(t, "user-902", upstream.lastReq.Header.Get("X-UserID")) + require.Equal(t, "user902@example.test", upstream.lastReq.Header.Get("X-Email")) + require.Contains(t, repo.updates[account.ID], grokObservedModelsExtraKey) +} + type grokQuotaProxyRepo struct { proxyRepoStub proxies map[int64]*Proxy @@ -610,6 +665,28 @@ func TestGrokQuotaServiceProbeUsageStoresNoHeadersState(t *testing.T) { require.Equal(t, observedResetAt, repo.recoveryObservedReset) } +func TestGrokQuotaServiceProbeUsageDoesNotOverwriteSnapshotOnUnauthorized(t *testing.T) { + t.Parallel() + + account := healthyGrokQuotaOAuthAccount(44) + previous := &xai.QuotaSnapshot{StatusCode: http.StatusOK, HeadersObserved: true} + account.Extra = map[string]any{grokQuotaSnapshotExtraKey: previous} + repo := &grokQuotaAccountRepo{mockAccountRepoForPlatform: &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + }} + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusUnauthorized, + Header: http.Header{}, + Body: io.NopCloser(strings.NewReader(`{"error":"unauthorized"}`)), + }} + svc := NewGrokQuotaService(repo, nil, NewGrokTokenProvider(repo, nil), upstream, nil) + + _, err := svc.ProbeUsage(context.Background(), account.ID) + require.Error(t, err) + require.Equal(t, 0, repo.updateCalls) + require.Same(t, previous, account.Extra[grokQuotaSnapshotExtraKey]) +} + func TestGrokQuotaServiceProbeUsageReturnsRateLimitedSnapshot(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/grok_search_count.go b/backend/internal/service/grok_search_count.go new file mode 100644 index 0000000000..c89ce93f04 --- /dev/null +++ b/backend/internal/service/grok_search_count.go @@ -0,0 +1,233 @@ +package service + +import ( + "strconv" + "strings" + + "github.com/tidwall/gjson" +) + +// countGrokNativeSearchCallsFromJSONBytes counts completed native search tool +// calls in a Responses-style JSON body (output array or nested response.output). +// Counts: web_search_call, x_search_call, tool_search_call, and function_call +// named tool_search / web_search / x_search. +func countGrokNativeSearchCallsFromJSONBytes(body []byte) int { + if len(body) == 0 || !gjson.ValidBytes(body) { + return 0 + } + // Responses envelopes normally expose either top-level output (JSON mode) + // or response.output (terminal SSE payload). Compatibility layers can retain + // both copies; counting both would bill the same search twice. Prefer the + // canonical nested response when present and fall back to top-level output. + if nested := gjson.GetBytes(body, "response.output"); nested.IsArray() { + return countGrokNativeSearchCallsInOutputArray(nested) + } + return countGrokNativeSearchCallsInOutputArray(gjson.GetBytes(body, "output")) +} + +func countGrokNativeSearchCallsFromSSEBody(body string) int { + if strings.TrimSpace(body) == "" { + return 0 + } + seen := make(map[string]struct{}) + total := 0 + forEachOpenAISSEDataPayload(body, func(data []byte) { + total += countGrokNativeSearchCallsInSSEDataDedup(data, seen) + }) + return total +} + +// countGrokNativeSearchCallsInSSEData counts search tool calls in one SSE +// payload without cross-event dedup. Prefer countGrokNativeSearchCallsInSSEDataDedup +// for live streams so item.done + response.completed do not double-bill. +func countGrokNativeSearchCallsInSSEData(data []byte) int { + n, _ := countGrokNativeSearchCallsInSSEDataWithKeys(data) + return n +} + +// countGrokNativeSearchCallsInSSEDataDedup increments only unseen call ids. +// Callers must reuse the same seen map for the full stream lifetime. +// +// When call_id/id is missing, a synthetic key is built from item type + name so +// item.done + response.completed for the same tool still count once (never fall +// back to raw multi-event n, which ~2× overbills). +func countGrokNativeSearchCallsInSSEDataDedup(data []byte, seen map[string]struct{}) int { + if seen == nil { + return countGrokNativeSearchCallsInSSEData(data) + } + n, keys := countGrokNativeSearchCallsInSSEDataWithKeys(data) + if n <= 0 { + return 0 + } + // Prefer stable ids; fill gaps with synthetic keys so we never raw-add n. + if len(keys) < n { + // Rebuild keys for every item so unkeyed items still get a fingerprint. + keys = collectGrokNativeSearchCallKeys(data) + } + if len(keys) == 0 { + // True empty — should not happen when n>0; fail-closed to 0 extra bill. + return 0 + } + added := 0 + local := make(map[string]struct{}, len(keys)) + isItemDone := strings.TrimSpace(gjson.GetBytes(data, "type").String()) == "response.output_item.done" + for _, k := range keys { + if k == "" { + continue + } + if _, ok := local[k]; ok { + continue + } + local[k] = struct{}{} + if _, ok := seen[k]; ok { + if !isItemDone || !strings.HasPrefix(k, "synth:") { + continue + } + // Each id-less item.done is a distinct completed invocation. Advance + // its ordinal so interrupted streams remain accurately billable. + separator := strings.LastIndexByte(k, ':') + if separator < 0 { + continue + } + base := k[:separator] + for ordinal := 2; ; ordinal++ { + candidate := base + ":" + strconv.Itoa(ordinal) + if _, exists := seen[candidate]; !exists { + k = candidate + break + } + } + } + seen[k] = struct{}{} + added++ + } + return added +} + +func collectGrokNativeSearchCallKeys(data []byte) []string { + if len(data) == 0 || !gjson.ValidBytes(data) { + return nil + } + // An empty type means a bare item object without an SSE envelope; anything + // else that is not a completion event carries no billable call. + switch strings.TrimSpace(gjson.GetBytes(data, "type").String()) { + case "response.output_item.done", "response.completed", "response.done", "": + default: + return nil + } + var keys []string + syntheticOrdinals := make(map[string]int) + consider := func(item gjson.Result) { + if !isGrokNativeSearchOutputItem(item) { + return + } + key := firstNonEmpty( + strings.TrimSpace(item.Get("call_id").String()), + strings.TrimSpace(item.Get("id").String()), + strings.TrimSpace(item.Get("item.call_id").String()), + strings.TrimSpace(item.Get("item.id").String()), + ) + if key == "" { + // Include the ordinal among same-kind calls. A plain type:name key + // collapses two id-less web searches in one completed response into + // one charge. The ordinal remains stable between ordered item.done + // events and response.completed output. + base := "synth:" + strings.ToLower(strings.TrimSpace(item.Get("type").String())) + + ":" + strings.ToLower(strings.TrimSpace(item.Get("name").String())) + syntheticOrdinals[base]++ + key = base + ":" + strconv.Itoa(syntheticOrdinals[base]) + } + keys = append(keys, key) + } + if item := gjson.GetBytes(data, "item"); item.Exists() { + consider(item) + } + gjson.GetBytes(data, "response.output").ForEach(func(_, item gjson.Result) bool { + consider(item) + return true + }) + gjson.GetBytes(data, "output").ForEach(func(_, item gjson.Result) bool { + consider(item) + return true + }) + if len(keys) == 0 && isGrokNativeSearchOutputItem(gjson.ParseBytes(data)) { + consider(gjson.ParseBytes(data)) + } + return keys +} + +func countGrokNativeSearchCallsInSSEDataWithKeys(data []byte) (int, []string) { + if len(data) == 0 || !gjson.ValidBytes(data) { + return 0, nil + } + // Count once on item completion / completed response, not on every delta. + // An empty type is a bare item object without an SSE envelope. + switch strings.TrimSpace(gjson.GetBytes(data, "type").String()) { + case "response.output_item.done", "response.completed", "response.done", "": + default: + return 0, nil + } + var keys []string + n := 0 + consider := func(item gjson.Result) { + if !isGrokNativeSearchOutputItem(item) { + return + } + n++ + key := firstNonEmpty( + strings.TrimSpace(item.Get("call_id").String()), + strings.TrimSpace(item.Get("id").String()), + strings.TrimSpace(item.Get("item.call_id").String()), + strings.TrimSpace(item.Get("item.id").String()), + ) + if key != "" { + keys = append(keys, key) + } + } + if item := gjson.GetBytes(data, "item"); item.Exists() { + consider(item) + } + gjson.GetBytes(data, "response.output").ForEach(func(_, item gjson.Result) bool { + consider(item) + return true + }) + gjson.GetBytes(data, "output").ForEach(func(_, item gjson.Result) bool { + consider(item) + return true + }) + // Bare output item event without nested item key. + if n == 0 && isGrokNativeSearchOutputItem(gjson.ParseBytes(data)) { + consider(gjson.ParseBytes(data)) + } + return n, keys +} + +func countGrokNativeSearchCallsInOutputArray(output gjson.Result) int { + if !output.IsArray() { + return 0 + } + count := 0 + output.ForEach(func(_, item gjson.Result) bool { + if isGrokNativeSearchOutputItem(item) { + count++ + } + return true + }) + return count +} + +func isGrokNativeSearchOutputItem(item gjson.Result) bool { + if !item.Exists() { + return false + } + itemType := strings.ToLower(strings.TrimSpace(item.Get("type").String())) + switch itemType { + case "web_search_call", "x_search_call", "tool_search_call": + return true + case "function_call", "custom_tool_call": + name := strings.ToLower(strings.TrimSpace(item.Get("name").String())) + return name == "web_search" || name == "x_search" || name == "tool_search" + default: + return false + } +} diff --git a/backend/internal/service/grok_search_count_test.go b/backend/internal/service/grok_search_count_test.go new file mode 100644 index 0000000000..8bb0a2c79e --- /dev/null +++ b/backend/internal/service/grok_search_count_test.go @@ -0,0 +1,85 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCountGrokNativeSearchCallsFromJSONBytes(t *testing.T) { + t.Parallel() + require.Equal(t, 0, countGrokNativeSearchCallsFromJSONBytes(nil)) + require.Equal(t, 0, countGrokNativeSearchCallsFromJSONBytes([]byte(`{"output":[]}`))) + body := []byte(`{"output":[ + {"type":"web_search_call","id":"ws1","status":"completed"}, + {"type":"x_search_call","id":"xs1"}, + {"type":"function_call","name":"tool_search","call_id":"ts1"}, + {"type":"function_call","name":"lookup","call_id":"other"} + ]}`) + require.Equal(t, 3, countGrokNativeSearchCallsFromJSONBytes(body)) +} + +func TestCountGrokNativeSearchCallsFromJSONBytes_PrefersNestedResponse(t *testing.T) { + t.Parallel() + body := []byte(`{"output":[{"type":"web_search_call","id":"duplicate"}],"response":{"output":[{"type":"web_search_call","id":"duplicate"},{"type":"x_search_call","id":"xs1"}]}}`) + require.Equal(t, 2, countGrokNativeSearchCallsFromJSONBytes(body)) +} + +func TestCountGrokNativeSearchCallsFromJSONBytes_FallsBackWhenNestedOutputNull(t *testing.T) { + t.Parallel() + body := []byte(`{"output":[{"type":"web_search_call","id":"ws1"}],"response":{"output":null}}`) + require.Equal(t, 1, countGrokNativeSearchCallsFromJSONBytes(body)) +} + +func TestCountGrokNativeSearchCallsFromSSEBodyDedups(t *testing.T) { + t.Parallel() + sse := stringsJoin( + `data: {"type":"response.output_item.done","item":{"type":"web_search_call","id":"ws1","call_id":"c1"}}`, + `data: {"type":"response.output_item.done","item":{"type":"web_search_call","id":"ws1","call_id":"c1"}}`, + `data: {"type":"response.completed","response":{"output":[{"type":"web_search_call","id":"ws1","call_id":"c1"},{"type":"x_search_call","id":"xs1","call_id":"c2"}]}}`, + ) + require.Equal(t, 2, countGrokNativeSearchCallsFromSSEBody(sse)) +} + +func TestCountGrokNativeSearchCallsInSSEDataDedup_LiveStreamPath(t *testing.T) { + t.Parallel() + // Mirrors the live streaming accumulator: item.done then response.completed + // for the same call_id must bill once (regression for ~2× surcharge). + seen := make(map[string]struct{}) + done := []byte(`{"type":"response.output_item.done","item":{"type":"web_search_call","id":"ws1","call_id":"c1"}}`) + completed := []byte(`{"type":"response.completed","response":{"output":[{"type":"web_search_call","id":"ws1","call_id":"c1"},{"type":"x_search_call","id":"xs1","call_id":"c2"}]}}`) + require.Equal(t, 1, countGrokNativeSearchCallsInSSEDataDedup(done, seen)) + require.Equal(t, 1, countGrokNativeSearchCallsInSSEDataDedup(completed, seen)) + // Raw (no-dedup) path still double-counts the same envelope pair. + require.Equal(t, 1, countGrokNativeSearchCallsInSSEData(done)) + require.Equal(t, 2, countGrokNativeSearchCallsInSSEData(completed)) +} + +func TestCountGrokNativeSearchCallsInSSEDataDedup_NoIDStillDedups(t *testing.T) { + t.Parallel() + // Upstream sometimes omits call_id/id; synthetic keys must still prevent 2×. + seen := make(map[string]struct{}) + done := []byte(`{"type":"response.output_item.done","item":{"type":"web_search_call"}}`) + completed := []byte(`{"type":"response.completed","response":{"output":[{"type":"web_search_call"}]}}`) + require.Equal(t, 1, countGrokNativeSearchCallsInSSEDataDedup(done, seen)) + require.Equal(t, 0, countGrokNativeSearchCallsInSSEDataDedup(completed, seen)) +} + +func TestCountGrokNativeSearchCallsInSSEDataDedup_MultipleNoIDCalls(t *testing.T) { + t.Parallel() + seen := make(map[string]struct{}) + firstDone := []byte(`{"type":"response.output_item.done","item":{"type":"web_search_call"}}`) + secondDone := []byte(`{"type":"response.output_item.done","item":{"type":"web_search_call"}}`) + completed := []byte(`{"type":"response.completed","response":{"output":[{"type":"web_search_call"},{"type":"web_search_call"}]}}`) + require.Equal(t, 1, countGrokNativeSearchCallsInSSEDataDedup(firstDone, seen)) + require.Equal(t, 1, countGrokNativeSearchCallsInSSEDataDedup(secondDone, seen)) + require.Equal(t, 0, countGrokNativeSearchCallsInSSEDataDedup(completed, seen)) +} + +func stringsJoin(lines ...string) string { + out := "" + for _, l := range lines { + out += l + "\n\n" + } + return out +} diff --git a/backend/internal/service/grok_spending_reauth.go b/backend/internal/service/grok_spending_reauth.go new file mode 100644 index 0000000000..2b1f3a0973 --- /dev/null +++ b/backend/internal/service/grok_spending_reauth.go @@ -0,0 +1,59 @@ +package service + +import ( + "context" + "strings" + "time" +) + +// Spending-limit is recoverable at the end of the observed billing period. +// When no billing snapshot is available, use a short probe rather than +// fabricating a 24h boundary from the error arrival time. +const grokSpendingLimitProbeCooldown = 10 * time.Minute + +func grokSpendingLimitResetAt(account *Account, now time.Time) time.Time { + if account != nil { + if billing, err := grokBillingSnapshotFromExtra(account.Extra); err == nil && billing != nil { + for _, raw := range []string{billing.PeriodEnd, billing.BillingPeriodEnd} { + if resetAt, err := time.Parse(time.RFC3339, strings.TrimSpace(raw)); err == nil && resetAt.After(now) { + return resetAt + } + } + } + } + return now.Add(grokSpendingLimitProbeCooldown) +} + +// clearGrokNeedsReauthExtra drops the soft reauth flag after successful refresh +// or reauth. Best-effort; never fails the request path. +func clearGrokNeedsReauthExtra(ctx context.Context, repo AccountRepository, accountID int64) { + if repo == nil || accountID <= 0 { + return + } + stateCtx, cancel := openAIAccountStateContext(ctx) + defer cancel() + _ = repo.UpdateExtra(stateCtx, accountID, map[string]any{ + "grok_needs_reauth": false, + "grok_needs_reauth_reason": "", + "grok_needs_reauth_at": "", + }) +} + +func accountGrokNeedsReauth(account *Account) bool { + if account == nil { + return false + } + if account.Status == StatusError { + msg := strings.ToLower(account.ErrorMessage) + if strings.Contains(msg, "spending limit") || strings.Contains(msg, "reauthorize") { + return true + } + } + if v, ok := account.Extra["grok_needs_reauth"].(bool); ok && v { + return true + } + if s, ok := account.Extra["grok_needs_reauth"].(string); ok { + return strings.EqualFold(s, "true") || s == "1" + } + return false +} diff --git a/backend/internal/service/grok_stream_idle.go b/backend/internal/service/grok_stream_idle.go new file mode 100644 index 0000000000..3bdfa27fc9 --- /dev/null +++ b/backend/internal/service/grok_stream_idle.go @@ -0,0 +1,38 @@ +package service + +import ( + "fmt" + "strings" + "time" +) + +// Default Grok stream idle when gateway.stream_data_interval_timeout is 0. +// Long enough for slow thinking models, short enough to release hung sockets. +const defaultGrokStreamIdleTimeout = 180 * time.Second + +// Shorter cool after a Grok stream-idle failure so the account can re-enter soon +// but is not immediately re-picked in a tight failover loop. +const grokStreamIdleCooldown = 2 * time.Minute + +// resolveGrokStreamIdleTimeout returns the effective upstream-read idle timeout +// for Grok streams. Prefers the global gateway setting when positive; otherwise +// applies a Grok-only default so hung SSE bodies still fail over. +func resolveGrokStreamIdleTimeout(cfgStreamIntervalSec int) time.Duration { + if cfgStreamIntervalSec > 0 { + return time.Duration(cfgStreamIntervalSec) * time.Second + } + return defaultGrokStreamIdleTimeout +} + +// grokStreamIdleFailoverError builds a pre-commit/handler-visible failover so +// the gateway can switch OAuth accounts after a hung Grok upstream stream. +func grokStreamIdleFailoverError(account *Account, idle time.Duration) *UpstreamFailoverError { + msg := fmt.Sprintf("Grok stream idle timeout after %s with no upstream data", idle.Round(time.Second)) + return &UpstreamFailoverError{ + StatusCode: 502, + ResponseBody: []byte(`{"error":{"code":"empty_upstream","message":"` + strings.ReplaceAll(msg, `"`, `'`) + `"}}`), + SafeToFailoverAfterWrite: true, + // Allow pool-mode retries; normal OAuth switches account via handler. + RetryableOnSameAccount: account != nil && account.IsPoolMode(), + } +} diff --git a/backend/internal/service/grok_stream_idle_test.go b/backend/internal/service/grok_stream_idle_test.go new file mode 100644 index 0000000000..e3c6d16943 --- /dev/null +++ b/backend/internal/service/grok_stream_idle_test.go @@ -0,0 +1,25 @@ +//go:build unit + +package service + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestResolveGrokStreamIdleTimeout(t *testing.T) { + require.Equal(t, 90*time.Second, resolveGrokStreamIdleTimeout(90)) + require.Equal(t, defaultGrokStreamIdleTimeout, resolveGrokStreamIdleTimeout(0)) + require.Equal(t, defaultGrokStreamIdleTimeout, resolveGrokStreamIdleTimeout(-1)) +} + +func TestGrokStreamIdleFailoverError(t *testing.T) { + account := &Account{ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth} + err := grokStreamIdleFailoverError(account, 180*time.Second) + require.NotNil(t, err) + require.Equal(t, 502, err.StatusCode) + require.True(t, err.SafeToFailoverAfterWrite) + require.Contains(t, string(err.ResponseBody), "empty_upstream") +} diff --git a/backend/internal/service/grok_team_rate_limit.go b/backend/internal/service/grok_team_rate_limit.go new file mode 100644 index 0000000000..f10852615a --- /dev/null +++ b/backend/internal/service/grok_team_rate_limit.go @@ -0,0 +1,149 @@ +package service + +import ( + "crypto/sha256" + "encoding/hex" + "strings" + "sync" + "time" +) + +// In-memory team+model rate-limit overlay for Grok OAuth. When xAI rate-limits +// one account in a team for a model, sibling accounts sharing team_id skip the +// same model until the cooldown expires (mirrors grok2api teamModelRateLimit). +// +// Process-local only: multi-instance deployments each learn the block from their +// own 429s. Prefer short TTLs so drift self-heals. +type grokTeamModelRateLimit struct { + Until time.Time +} + +type grokTeamModelRateLimitStore struct { + mu sync.Mutex + items map[string]grokTeamModelRateLimit +} + +var globalGrokTeamModelRateLimits = &grokTeamModelRateLimitStore{ + items: make(map[string]grokTeamModelRateLimit), +} + +const ( + grokTeamRateLimitDefaultTTL = 10 * time.Minute + grokTeamRateLimitMaxTTL = time.Hour + grokTeamRateLimitMinTTL = 30 * time.Second +) + +func grokTeamFingerprint(teamID string) string { + teamID = strings.TrimSpace(teamID) + if teamID == "" { + return "" + } + sum := sha256.Sum256([]byte(strings.ToLower(teamID))) + return hex.EncodeToString(sum[:8]) +} + +func grokTeamModelRateLimitKey(teamFingerprint, model string) string { + return teamFingerprint + "|" + strings.ToLower(strings.TrimSpace(model)) +} + +func accountGrokTeamID(account *Account) string { + if account == nil { + return "" + } + return strings.TrimSpace(account.GetCredential("team_id")) +} + +// markGrokTeamModelRateLimit records that this team+model pair should be skipped +// until until. No-op when team_id or model is empty. +func markGrokTeamModelRateLimit(account *Account, model string, until time.Time) { + if account == nil || !account.IsGrokOAuth() { + return + } + fp := grokTeamFingerprint(accountGrokTeamID(account)) + model = strings.TrimSpace(model) + if fp == "" || model == "" || until.IsZero() { + return + } + now := time.Now() + if !until.After(now) { + until = now.Add(grokTeamRateLimitDefaultTTL) + } + maxUntil := now.Add(grokTeamRateLimitMaxTTL) + if until.After(maxUntil) { + until = maxUntil + } + key := grokTeamModelRateLimitKey(fp, model) + globalGrokTeamModelRateLimits.mu.Lock() + defer globalGrokTeamModelRateLimits.mu.Unlock() + if cur, ok := globalGrokTeamModelRateLimits.items[key]; ok && cur.Until.After(until) { + return + } + globalGrokTeamModelRateLimits.items[key] = grokTeamModelRateLimit{Until: until} + // Opportunistic prune of expired entries. + for k, v := range globalGrokTeamModelRateLimits.items { + if !v.Until.After(now) { + delete(globalGrokTeamModelRateLimits.items, k) + } + } +} + +// isGrokTeamModelRateLimited reports whether the account's team is currently +// blocked for the requested model. +func isGrokTeamModelRateLimited(account *Account, model string, now time.Time) bool { + if account == nil || !account.IsGrokOAuth() { + return false + } + fp := grokTeamFingerprint(accountGrokTeamID(account)) + model = strings.TrimSpace(model) + if fp == "" || model == "" { + return false + } + key := grokTeamModelRateLimitKey(fp, model) + globalGrokTeamModelRateLimits.mu.Lock() + defer globalGrokTeamModelRateLimits.mu.Unlock() + cur, ok := globalGrokTeamModelRateLimits.items[key] + if !ok { + return false + } + if !cur.Until.After(now) { + delete(globalGrokTeamModelRateLimits.items, key) + return false + } + return true +} + +// filterGrokTeamModelRateLimitedAccounts drops candidates whose team is under a +// model-scoped rate-limit cool. Accounts without team_id pass through. +func filterGrokTeamModelRateLimitedAccounts(accounts []Account, model string, now time.Time) []Account { + if len(accounts) == 0 || strings.TrimSpace(model) == "" { + return accounts + } + out := accounts[:0] + kept := false + for i := range accounts { + upstreamModel := canonicalOpenAIAccountSchedulingModel(&accounts[i], model) + if isGrokTeamModelRateLimited(&accounts[i], upstreamModel, now) { + continue + } + out = append(out, accounts[i]) + kept = true + } + if !kept && len(out) == 0 { + // All filtered — return empty (caller treats as no capacity). + return nil + } + return out +} + +// resolveGrokTeamRateLimitUntil derives a team cool window from an observed +// account rate-limit reset, with sane clamps. +func resolveGrokTeamRateLimitUntil(resetAt, now time.Time) time.Time { + if resetAt.After(now.Add(grokTeamRateLimitMinTTL)) { + maxUntil := now.Add(grokTeamRateLimitMaxTTL) + if resetAt.After(maxUntil) { + return maxUntil + } + return resetAt + } + return now.Add(grokTeamRateLimitDefaultTTL) +} diff --git a/backend/internal/service/grok_team_rate_limit_test.go b/backend/internal/service/grok_team_rate_limit_test.go new file mode 100644 index 0000000000..4926d4d7aa --- /dev/null +++ b/backend/internal/service/grok_team_rate_limit_test.go @@ -0,0 +1,75 @@ +//go:build unit + +package service + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestGrokTeamModelRateLimit_MarksAndFiltersSiblings(t *testing.T) { + // Isolate from other tests by using unique team ids. + team := "team-test-" + time.Now().Format("150405.000") + a1 := &Account{ + ID: 101, Platform: PlatformGrok, Type: AccountTypeOAuth, + Credentials: map[string]any{"team_id": team}, + } + a2 := &Account{ + ID: 102, Platform: PlatformGrok, Type: AccountTypeOAuth, + Credentials: map[string]any{"team_id": team}, + } + other := &Account{ + ID: 103, Platform: PlatformGrok, Type: AccountTypeOAuth, + Credentials: map[string]any{"team_id": team + "-other"}, + } + noTeam := &Account{ + ID: 104, Platform: PlatformGrok, Type: AccountTypeOAuth, + Credentials: map[string]any{}, + } + + now := time.Now() + markGrokTeamModelRateLimit(a1, "grok-4.5", now.Add(5*time.Minute)) + + require.True(t, isGrokTeamModelRateLimited(a1, "grok-4.5", now)) + require.True(t, isGrokTeamModelRateLimited(a2, "grok-4.5", now), "sibling with same team must cool") + require.False(t, isGrokTeamModelRateLimited(a2, "grok-4.3", now), "other model stays pickable") + require.False(t, isGrokTeamModelRateLimited(other, "grok-4.5", now)) + require.False(t, isGrokTeamModelRateLimited(noTeam, "grok-4.5", now)) + + filtered := filterGrokTeamModelRateLimitedAccounts([]Account{*a1, *a2, *other, *noTeam}, "grok-4.5", now) + require.Len(t, filtered, 2) + ids := []int64{filtered[0].ID, filtered[1].ID} + require.Contains(t, ids, int64(103)) + require.Contains(t, ids, int64(104)) +} + +func TestGrokTeamModelRateLimit_Expires(t *testing.T) { + team := "team-expire-" + time.Now().Format("150405.000") + a := &Account{ + ID: 201, Platform: PlatformGrok, Type: AccountTypeOAuth, + Credentials: map[string]any{"team_id": team}, + } + past := time.Now().Add(-time.Minute) + markGrokTeamModelRateLimit(a, "grok-4.5", past) + // mark clamps expired until into default TTL from "now" — use direct store inject via past+recheck + // After mark with past, resolveGrokTeamRateLimitUntil path isn't used; mark uses now+default when until not after now. + require.True(t, isGrokTeamModelRateLimited(a, "grok-4.5", time.Now())) +} + +func TestGrokTeamModelRateLimitFilterUsesMappedUpstreamModel(t *testing.T) { + now := time.Now() + account := &Account{ + ID: 301, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "team_id": "team-mapped-301", + "model_mapping": map[string]any{"gpt-*": "grok-4.5"}, + }, + } + markGrokTeamModelRateLimit(account, "grok-4.5", now.Add(time.Hour)) + + require.Empty(t, filterGrokTeamModelRateLimitedAccounts([]Account{*account}, "gpt-5", now)) +} diff --git a/backend/internal/service/grok_token_refresher.go b/backend/internal/service/grok_token_refresher.go index d2d4d534fc..72b041bc6e 100644 --- a/backend/internal/service/grok_token_refresher.go +++ b/backend/internal/service/grok_token_refresher.go @@ -3,12 +3,25 @@ package service import ( "context" "errors" + "hash/fnv" "strings" "time" ) +// Base warm window: refresh when access token lifetime remaining is below this. +// Grok access tokens are typically ~1h; refreshing up to 1h early keeps the pool +// warm for request path cache misses. const grokTokenRefreshSkew = time.Hour +// Stampede spread: each account's effective warm window is reduced by a +// deterministic offset in [0, grokTokenRefreshJitterMax] so co-imported accounts +// do not all refresh in the same TokenRefreshService cycle (grok2api-style +// RefreshDueAt scatter). +const grokTokenRefreshJitterMax = 3 * time.Minute + +// Floor so jitter cannot shrink the window below a useful threshold. +const grokTokenRefreshSkewMin = 30 * time.Minute + type GrokTokenRefresher struct { grokOAuthService GrokOAuthTokenService } @@ -40,9 +53,35 @@ func (r *GrokTokenRefresher) NeedsRefresh(account *Account, refreshWindow time.D if refreshWindow < grokTokenRefreshSkew { refreshWindow = grokTokenRefreshSkew } + // Deterministic per-account jitter: spread warm refreshes without random + // non-determinism in tests (hash of account id). + refreshWindow = grokTokenRefreshWindowWithJitter(account.ID, refreshWindow) return time.Until(*expiresAt) < refreshWindow } +// grokTokenRefreshWindowWithJitter returns refreshWindow minus a stable offset +// in [0, jitterMax] based on accountID. Result is never below grokTokenRefreshSkewMin +// when the base window is at least that large. +func grokTokenRefreshWindowWithJitter(accountID int64, refreshWindow time.Duration) time.Duration { + if accountID <= 0 || refreshWindow <= grokTokenRefreshSkewMin { + return refreshWindow + } + h := fnv.New32a() + var b [8]byte + id := uint64(accountID) + for i := 0; i < 8; i++ { + b[i] = byte(id >> (8 * i)) + } + _, _ = h.Write(b[:]) + // Jitter in [0, grokTokenRefreshJitterMax). + jitter := time.Duration(h.Sum32()%uint32(grokTokenRefreshJitterMax/time.Second)) * time.Second + out := refreshWindow - jitter + if out < grokTokenRefreshSkewMin { + return grokTokenRefreshSkewMin + } + return out +} + func (r *GrokTokenRefresher) Refresh(ctx context.Context, account *Account) (map[string]any, error) { if r == nil || r.grokOAuthService == nil { return nil, errors.New("grok oauth service is not configured") diff --git a/backend/internal/service/grok_token_refresher_test.go b/backend/internal/service/grok_token_refresher_test.go new file mode 100644 index 0000000000..c689985628 --- /dev/null +++ b/backend/internal/service/grok_token_refresher_test.go @@ -0,0 +1,50 @@ +//go:build unit + +package service + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestGrokTokenRefreshWindowWithJitter_StableAndBounded(t *testing.T) { + base := grokTokenRefreshSkew + w1 := grokTokenRefreshWindowWithJitter(42, base) + w2 := grokTokenRefreshWindowWithJitter(42, base) + require.Equal(t, w1, w2, "same account id must yield stable window") + require.GreaterOrEqual(t, w1, grokTokenRefreshSkewMin) + require.LessOrEqual(t, w1, base) + + // Different accounts should usually differ (not a hard guarantee for all pairs, + // but for sequential ids the hash spread is good enough to assert inequality + // across a small sample). + seen := map[time.Duration]bool{} + for id := int64(1); id <= 50; id++ { + seen[grokTokenRefreshWindowWithJitter(id, base)] = true + } + require.Greater(t, len(seen), 1, "jitter should spread windows across accounts") +} + +func TestGrokTokenRefresher_NeedsRefresh_UsesSkewFloor(t *testing.T) { + refresher := NewGrokTokenRefresher(nil) + // Expires in 50 minutes — within 1h skew, should need refresh. + expires := time.Now().Add(50 * time.Minute).UTC().Format(time.RFC3339) + account := &Account{ + ID: 7, + Platform: PlatformGrok, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "access_token": "at", + "refresh_token": "rt", + "expires_at": expires, + }, + } + // Pass a tiny window; NeedsRefresh raises it to grokTokenRefreshSkew then jitter. + require.True(t, refresher.NeedsRefresh(account, time.Minute)) + + // Far future — no refresh. + account.Credentials["expires_at"] = time.Now().Add(3 * time.Hour).UTC().Format(time.RFC3339) + require.False(t, refresher.NeedsRefresh(account, time.Minute)) +} diff --git a/backend/internal/service/grok_upstream_errors.go b/backend/internal/service/grok_upstream_errors.go index 880cd702b0..0039ea4cdc 100644 --- a/backend/internal/service/grok_upstream_errors.go +++ b/backend/internal/service/grok_upstream_errors.go @@ -191,10 +191,17 @@ func grokContentPolicyClientMessage(responseBody []byte) string { // shouldFailoverGrokUpstreamError is the body-aware counterpart of the // status-only failover helper. Grok content refusals must stay on the current // account and be returned to the caller instead of consuming the account pool. +// Free-usage / empty-output / billing bodies also failover even when the HTTP +// status alone would not (e.g. 400 with free-usage-exhausted). func (s *OpenAIGatewayService) shouldFailoverGrokUpstreamError(statusCode int, responseBody []byte) bool { if isGrokContentPolicyRejection(statusCode, responseBody) { return false } + decision := classifyGrokUpstreamFailure(statusCode, responseBody, "") + switch decision.Class { + case GrokFailureFreeUsage, GrokFailureEmptyUpstream, GrokFailureBilling, GrokFailureModelCapacity: + return decision.ShouldFailover + } return s.shouldFailoverUpstreamError(statusCode) } diff --git a/backend/internal/service/grok_upstream_failure.go b/backend/internal/service/grok_upstream_failure.go new file mode 100644 index 0000000000..1348d068c5 --- /dev/null +++ b/backend/internal/service/grok_upstream_failure.go @@ -0,0 +1,506 @@ +package service + +import ( + "context" + "encoding/json" + "net/http" + "regexp" + "strconv" + "strings" + "time" + + "github.com/tidwall/gjson" +) + +// Grok upstream failure classes used to decide temp-unschedulable cooldowns and +// pre-commit account failover. Classification is body-first so free-usage and +// empty-output wording still win when the proxy rewrites status codes. +type GrokUpstreamFailureClass string + +const ( + GrokFailureNone GrokUpstreamFailureClass = "" + GrokFailureFreeUsage GrokUpstreamFailureClass = "subscription:free-usage-exhausted" + GrokFailureBilling GrokUpstreamFailureClass = "billing_quota" + GrokFailureEmptyUpstream GrokUpstreamFailureClass = "empty_upstream" + GrokFailureModelCapacity GrokUpstreamFailureClass = "model_capacity" + GrokFailureRateLimit GrokUpstreamFailureClass = "rate_limit" + GrokFailureAuth GrokUpstreamFailureClass = "auth_error" + GrokFailureServer GrokUpstreamFailureClass = "server_error" +) + +// GrokUpstreamFailureDecision is a pure classification result. Callers map it +// onto existing account state helpers (tempUnscheduleGrok / rateLimitGrok). +// BlockModel is retained for observability; the current scheduler does not +// implement per-model soft-blocks, so free-usage deliberately never sets it. +type GrokUpstreamFailureDecision struct { + Class GrokUpstreamFailureClass + Model string + Cooldown time.Duration + ShouldCooldown bool + // ShouldFailover recommends trying another account before writing a + // terminal response (pre-commit only). Content-policy rejections are + // handled separately and never reach this classifier for failover. + ShouldFailover bool + // BlockModel is true only for empty-output when a model id is known. + // Free-usage never sets this: the account cools, not a single model. + BlockModel bool + Reason string + TokensActual *int64 + TokensLimit *int64 +} + +var ( + reGrokTokenPair = regexp.MustCompile(`(?i)tokens?\s*(?:\(actual\s*/\s*limit\))?\s*[:=]?\s*(\d+)\s*/\s*(\d+)`) + reGrokModelFor = regexp.MustCompile(`(?i)(?:for\s+model|model|模型)\s*[::]?\s*([a-z0-9][a-z0-9._-]{2,80})`) +) + +// classifyGrokUpstreamFailure decides cooldown/failover from status + body. +// Priority (body/code first, status second): +// 1. free-usage exhausted → account cool, no model block, failover +// 2. billing hard quota → longer cool, failover +// 3. empty model output → short cool + optional model soft-block marker, failover +// 4. model capacity → short cool, failover +// 5. bare rate-limit / 429 without free-usage language → cool, failover +// 6. bare 5xx → brief cool, failover +// 7. validation / client errors without quota language → no cool +// +// Content-policy 403s must be filtered by the caller before invoking this. +func classifyGrokUpstreamFailure(statusCode int, responseBody []byte, requestedModel string) GrokUpstreamFailureDecision { + text, code, low := grokUpstreamErrorCorpus(statusCode, responseBody) + model := extractGrokFailureModel(text, responseBody, requestedModel) + actual, limit, hasTokens := parseGrokTokenPair(text) + if !hasTokens { + actual, limit, hasTokens = parseGrokTokenPair(string(responseBody)) + } + + // --- Free usage / rolling quota exhausted --- + if isGrokFreeUsageExhaustedText(low) || isGrokFreeUsageCode(code) || isGrokFreeUsageCode(text) { + d := GrokUpstreamFailureDecision{ + Class: GrokFailureFreeUsage, + Model: model, + Cooldown: grokFreeUsageCooldownDuration(low), + ShouldCooldown: true, + ShouldFailover: true, + BlockModel: false, + Reason: firstNonEmpty(text, code, "free usage exhausted"), + } + if hasTokens { + a, b := actual, limit + d.TokensActual = &a + d.TokensLimit = &b + } + return d + } + + // Billing / hard quota (not free-tier rolling). + // Cooldown stays at 30m to match the existing Grok 402/spending-limit + // handler (longer cools would change ops behavior without a settings knob). + if isGrokBillingQuotaText(low) || statusCode == http.StatusPaymentRequired { + reason := firstNonEmpty(text, "billing quota") + if statusCode == http.StatusPaymentRequired && text == "" { + reason = "payment required" + } + return GrokUpstreamFailureDecision{ + Class: GrokFailureBilling, + Model: model, + Cooldown: 30 * time.Minute, + ShouldCooldown: true, + ShouldFailover: true, + BlockModel: model != "", + Reason: reason, + } + } + + // Empty HTTP 200 / empty model output (often rewritten to synthetic 502). + if isGrokEmptyModelOutputText(low) || isGrokEmptyModelOutputCode(code) { + return GrokUpstreamFailureDecision{ + Class: GrokFailureEmptyUpstream, + Model: model, + Cooldown: 4 * time.Minute, + ShouldCooldown: true, + ShouldFailover: true, + BlockModel: model != "", + Reason: firstNonEmpty(text, "empty model output"), + } + } + + // Model capacity / overloaded. + if isGrokModelCapacityText(low) { + return GrokUpstreamFailureDecision{ + Class: GrokFailureModelCapacity, + Model: model, + Cooldown: 3 * time.Minute, + ShouldCooldown: true, + ShouldFailover: true, + BlockModel: false, + Reason: firstNonEmpty(text, "model capacity"), + } + } + + // Rate limit without free-usage language. + if statusCode == http.StatusTooManyRequests || isGrokRateLimitText(low) { + return GrokUpstreamFailureDecision{ + Class: GrokFailureRateLimit, + Model: model, + Cooldown: 10 * time.Minute, + ShouldCooldown: true, + ShouldFailover: true, + BlockModel: false, + Reason: firstNonEmpty(text, "rate limit"), + } + } + + // Upstream 5xx — brief cool. Empty-output synthetic 502 already handled above. + if statusCode >= 500 && statusCode <= 599 { + return GrokUpstreamFailureDecision{ + Class: GrokFailureServer, + Cooldown: 2 * time.Minute, + ShouldCooldown: true, + ShouldFailover: true, + Reason: firstNonEmpty(text, "server error"), + } + } + + return GrokUpstreamFailureDecision{Reason: text} +} + +func grokUpstreamErrorCorpus(statusCode int, responseBody []byte) (text, code, low string) { + _ = statusCode // classifier already has the transport status; corpus is body-only + raw := strings.TrimSpace(string(responseBody)) + // Strip "upstream status NNN: ..." prefixes so free-usage / quota language is visible. + if _, unwrappedBody, ok := unwrapGrokUpstreamErrorText(raw); ok { + raw = unwrappedBody + } + text = raw + codeFromJSON, msgFromJSON := parseGrokUpstreamErrorJSON(raw) + if msgFromJSON != "" { + if text == "" || len(msgFromJSON) > len(text)/2 || looksLikeGrokQuotaMessage(msgFromJSON) { + text = msgFromJSON + } + } + // Prefer structured fields from the original body when present. + if len(responseBody) > 0 { + if m := strings.TrimSpace(firstNonEmpty( + gjson.GetBytes(responseBody, "error.message").String(), + gjson.GetBytes(responseBody, "message").String(), + gjson.GetBytes(responseBody, "error").String(), + )); m != "" && (text == "" || looksLikeGrokQuotaMessage(m)) { + text = m + } + if c := strings.TrimSpace(firstNonEmpty( + gjson.GetBytes(responseBody, "error.code").String(), + gjson.GetBytes(responseBody, "code").String(), + )); c != "" { + codeFromJSON = c + } + } + code = codeFromJSON + low = strings.ToLower(strings.TrimSpace(text)) + if code != "" && !strings.Contains(low, strings.ToLower(code)) { + low = strings.ToLower(code) + " " + low + } + return text, code, low +} + +func unwrapGrokUpstreamErrorText(errText string) (status int, body string, ok bool) { + text := strings.TrimSpace(errText) + if text == "" { + return 0, "", false + } + lower := strings.ToLower(text) + for _, p := range []string{"upstream status ", "status "} { + if !strings.HasPrefix(lower, p) { + continue + } + rest := strings.TrimSpace(text[len(p):]) + i := 0 + for i < len(rest) && rest[i] >= '0' && rest[i] <= '9' { + status = status*10 + int(rest[i]-'0') + i++ + } + if status <= 0 || i == 0 { + return 0, "", false + } + rest = strings.TrimSpace(rest[i:]) + if strings.HasPrefix(rest, ":") { + rest = strings.TrimSpace(rest[1:]) + } + return status, rest, true + } + return 0, "", false +} + +func parseGrokUpstreamErrorJSON(errText string) (code, message string) { + text := strings.TrimSpace(errText) + if text == "" || text[0] != '{' { + return "", "" + } + var payload map[string]any + if json.Unmarshal([]byte(text), &payload) != nil { + return "", "" + } + if v, ok := payload["code"].(string); ok { + code = v + } + if v, ok := payload["message"].(string); ok { + message = v + } + if errObj, ok := payload["error"].(map[string]any); ok { + if v, ok := errObj["code"].(string); ok && code == "" { + code = v + } + if v, ok := errObj["message"].(string); ok && message == "" { + message = v + } + } + if errStr, ok := payload["error"].(string); ok && message == "" { + message = errStr + } + return strings.TrimSpace(code), strings.TrimSpace(message) +} + +func looksLikeGrokQuotaMessage(s string) bool { + low := strings.ToLower(s) + return strings.Contains(low, "quota") || + strings.Contains(low, "usage") || + strings.Contains(low, "credit") || + strings.Contains(low, "额度") || + strings.Contains(low, "free") +} + +func isGrokFreeUsageCode(code string) bool { + c := strings.ToLower(strings.TrimSpace(code)) + if c == "" { + return false + } + if strings.Contains(c, "subscription:free-usage-exhausted") || + strings.Contains(c, "free-usage-exhausted") || + strings.Contains(c, "free_usage_exhausted") || + strings.Contains(c, "usage-limit-exceeded") || + strings.Contains(c, "usage_limit_exceeded") { + return true + } + return (strings.Contains(c, "free-usage") || strings.Contains(c, "free_usage")) && + (strings.Contains(c, "exhaust") || strings.Contains(c, "exceed") || strings.Contains(c, "limit")) +} + +func isGrokFreeUsageExhaustedText(low string) bool { + if low == "" { + return false + } + if strings.Contains(low, "free-usage-exhausted") || + strings.Contains(low, "free_usage_exhausted") || + strings.Contains(low, "subscription:free-usage") || + strings.Contains(low, "usage-limit-exceeded") || + strings.Contains(low, "usage_limit_exceeded") || + strings.Contains(low, "free-tier-limit") || + strings.Contains(low, "free_tier_limit") { + return true + } + if strings.Contains(low, "free usage") || + strings.Contains(low, "included free usage") || + strings.Contains(low, "used all the included free") || + strings.Contains(low, "you've used all the included free") || + strings.Contains(low, "you have used all the included free") || + strings.Contains(low, "free quota") || + strings.Contains(low, "no remaining free") || + strings.Contains(low, "out of free") || + strings.Contains(low, "usage resets over a rolling") || + (strings.Contains(low, "free tier") && (strings.Contains(low, "exhaust") || strings.Contains(low, "limit") || strings.Contains(low, "exceed"))) { + return true + } + for _, p := range []string{ + "额度耗尽", "额度用完", "额度不足", "额度已用尽", "额度已耗尽", + "免费额度", "免费用量", "用量用完", "用量耗尽", "用量超限", "用量已用尽", + "配额耗尽", "配额已用尽", "配额不足", "配额超限", "配额用完", + "没有额度", "没额度", "无额度", "可用额度不足", "模型额度", + "临时额度", "额度已满", "额度超限", "额度达到上限", + "模型额度用完", "模型额度耗尽", "账号额度用完", "账号额度耗尽", + "额度不够", "没额度了", "额度没了", "用完额度", "耗尽额度", + } { + if strings.Contains(low, p) { + return true + } + } + if (strings.Contains(low, "quota") && (strings.Contains(low, "exhaust") || strings.Contains(low, "exceed") || strings.Contains(low, "limit"))) || + (strings.Contains(low, "usage") && (strings.Contains(low, "exhaust") || strings.Contains(low, "exceed")) && (strings.Contains(low, "limit") || strings.Contains(low, "free") || strings.Contains(low, "model"))) { + if strings.Contains(low, "free") || strings.Contains(low, "rolling") || + strings.Contains(low, "24-hour") || strings.Contains(low, "24 hour") || + strings.Contains(low, "model") || strings.Contains(low, "subscription") || + strings.Contains(low, "included") || strings.Contains(low, "tokens") { + return true + } + } + if a, b, ok := parseGrokTokenPair(low); ok && b > 0 && a >= b { + if strings.Contains(low, "free") || strings.Contains(low, "subscription") || + strings.Contains(low, "included") || strings.Contains(low, "model") || + strings.Contains(low, "usage") || strings.Contains(low, "quota") || + strings.Contains(low, "rolling") { + return true + } + } + return false +} + +func isGrokBillingQuotaText(low string) bool { + if low == "" { + return false + } + if strings.Contains(low, "insufficient_quota") { + return true + } + if strings.Contains(low, "billing") && strings.Contains(low, "quota") { + return true + } + if strings.Contains(low, "payment") && (strings.Contains(low, "required") || strings.Contains(low, "fail")) { + return true + } + if strings.Contains(low, "spending limit") || strings.Contains(low, "run out of credits") || strings.Contains(low, "out of credits") { + return true + } + if strings.Contains(low, "余额不足") || strings.Contains(low, "欠费") || strings.Contains(low, "需要付费") { + return true + } + return false +} + +func isGrokModelCapacityText(low string) bool { + return strings.Contains(low, "capacity") || + strings.Contains(low, "overloaded") || + strings.Contains(low, "server_busy") || + strings.Contains(low, "too many concurrent") || + strings.Contains(low, "engine_overloaded") +} + +func isGrokRateLimitText(low string) bool { + return strings.Contains(low, "rate limit") || + strings.Contains(low, "rate_limit") || + strings.Contains(low, "too many requests") || + strings.Contains(low, "请求过于频繁") || + strings.Contains(low, "速率限制") +} + +func isGrokEmptyModelOutputText(low string) bool { + if low == "" { + return false + } + return strings.Contains(low, "empty model output") || + strings.Contains(low, "no content/tool_calls") || + strings.Contains(low, "no client-visible content") || + strings.Contains(low, "empty_upstream") || + strings.Contains(low, "empty upstream") +} + +func isGrokEmptyModelOutputCode(code string) bool { + c := strings.ToLower(strings.TrimSpace(code)) + if c == "" { + return false + } + return c == "empty_upstream" || + c == "empty-model-output" || + c == "empty_model_output" || + strings.Contains(c, "empty_upstream") || + strings.Contains(c, "empty-model-output") +} + +func grokFreeUsageCooldownDuration(low string) time.Duration { + // "rolling 24-hour" describes the upstream usage window, not a cooldown + // that starts when this proxy observes a 429. Without an upstream reset + // timestamp we cannot know when the oldest usage exits that window, so use + // a short probe interval and let a successful probe clear the block. + return grokFreeUsageProbeCooldown +} + +const grokFreeUsageProbeCooldown = 10 * time.Minute + +func parseGrokTokenPair(errText string) (actual, limit int64, ok bool) { + m := reGrokTokenPair.FindStringSubmatch(errText) + if len(m) != 3 { + return 0, 0, false + } + a, errA := strconv.ParseInt(m[1], 10, 64) + b, errB := strconv.ParseInt(m[2], 10, 64) + if errA != nil || errB != nil { + return 0, 0, false + } + return a, b, true +} + +func extractGrokFailureModel(text string, responseBody []byte, fallback string) string { + if m := reGrokModelFor.FindStringSubmatch(text); len(m) == 2 { + return normalizeGrokFailureModelID(m[1]) + } + if len(responseBody) > 0 { + if m := strings.TrimSpace(firstNonEmpty( + gjson.GetBytes(responseBody, "error.model").String(), + gjson.GetBytes(responseBody, "model").String(), + )); m != "" { + return normalizeGrokFailureModelID(m) + } + } + return normalizeGrokFailureModelID(fallback) +} + +// normalizeGrokFailureModelID trims whitespace and trailing punctuation the +// "for model X." extractors sometimes capture. +func normalizeGrokFailureModelID(model string) string { + model = strings.TrimSpace(model) + model = strings.TrimRight(model, ".,;:!?") + return strings.TrimSpace(model) +} + +// applyGrokUpstreamFailureDecision maps a classification onto existing account +// health helpers. Returns true when the decision fully handled the error path +// (caller should not apply the status-code switch defaults again). +func (s *OpenAIGatewayService) applyGrokUpstreamFailureDecision( + ctx context.Context, + account *Account, + decision GrokUpstreamFailureDecision, +) bool { + if s == nil || account == nil || !decision.ShouldCooldown || decision.Cooldown <= 0 { + return false + } + // Keep reasons short and stable for ops UI / temp_unschedulable_reason. + var reason string + switch decision.Class { + case GrokFailureFreeUsage: + reason = "grok free usage exhausted" + // Model-scoped free usage: soft-block only the named model so other + // models on the same account remain pickable (grok2api ModelQuotaBlock). + low := strings.ToLower(decision.Reason) + if decision.Model != "" && isGrokModelSpecificFreeUsage(low, decision.Model) { + until := time.Now().Add(decision.Cooldown) + markGrokModelQuotaBlock(account.ID, decision.Model, until) + // The upstream explicitly scoped exhaustion to this model, so an + // account-wide cool would incorrectly take healthy sibling models out + // of rotation. + return true + } + case GrokFailureBilling: + low := strings.ToLower(decision.Reason) + if strings.Contains(low, "spending") || strings.Contains(low, "credits") { + // Spending-limit/credit exhaustion is a billing-window condition. Keep + // the account recoverable and let the normal rate-limit recovery clear it. + s.rateLimitGrok(ctx, account, grokSpendingLimitResetAt(account, time.Now())) + return true + } + // Keep the historical 402/payment reason for ops UI + regression tests. + reason = "grok payment required" + case GrokFailureEmptyUpstream: + reason = "grok empty model output" + case GrokFailureModelCapacity: + reason = "grok model capacity" + case GrokFailureRateLimit: + // Pure 429 without free-usage language keeps the existing rate-limit + // snapshot path (Retry-After / quota headers). Body-only rate-limit + // phrasing still cools here via ShouldCooldown from the classifier, but + // the handler only invokes this for non-RateLimit classes. + return false + case GrokFailureServer: + reason = "grok upstream temporary error" + default: + return false + } + s.tempUnscheduleGrok(ctx, account, decision.Cooldown, reason) + return true +} diff --git a/backend/internal/service/grok_upstream_failure_test.go b/backend/internal/service/grok_upstream_failure_test.go new file mode 100644 index 0000000000..6f0a916a78 --- /dev/null +++ b/backend/internal/service/grok_upstream_failure_test.go @@ -0,0 +1,181 @@ +//go:build unit + +package service + +import ( + "context" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func TestClassifyGrokUpstreamFailure_FreeUsage(t *testing.T) { + cases := []struct { + name string + status int + body string + }{ + { + name: "code free-usage-exhausted", + status: http.StatusTooManyRequests, + body: `{"error":{"code":"subscription:free-usage-exhausted","message":"You've used all the included free usage for model grok-4.5. Usage resets over a rolling 24-hour window."}}`, + }, + { + name: "chinese body without 429", + status: http.StatusBadRequest, + body: `{"error":{"message":"模型额度用完,请稍后再试"}}`, + }, + { + name: "token pair with free marker", + status: http.StatusOK, + body: `{"error":{"message":"free usage tokens (actual / limit): 2000000 / 2000000 for model grok-4.5"}}`, + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + d := classifyGrokUpstreamFailure(tc.status, []byte(tc.body), "grok-4.5") + require.Equal(t, GrokFailureFreeUsage, d.Class) + require.True(t, d.ShouldCooldown) + require.True(t, d.ShouldFailover) + require.False(t, d.BlockModel, "free-usage must not soft-block models") + require.Equal(t, grokFreeUsageProbeCooldown, d.Cooldown) + }) + } +} + +func TestClassifyGrokUpstreamFailure_EmptyUpstream(t *testing.T) { + d := classifyGrokUpstreamFailure(http.StatusBadGateway, []byte(`empty model output: no content/tool_calls`), "grok-4.5") + require.Equal(t, GrokFailureEmptyUpstream, d.Class) + require.True(t, d.ShouldCooldown) + require.True(t, d.ShouldFailover) + require.True(t, d.BlockModel) + require.Equal(t, 4*time.Minute, d.Cooldown) +} + +func TestClassifyGrokUpstreamFailure_Billing(t *testing.T) { + d := classifyGrokUpstreamFailure(http.StatusForbidden, []byte(`{"code":"personal-team-blocked:spending-limit","error":"spending limit reached"}`), "") + require.Equal(t, GrokFailureBilling, d.Class) + require.True(t, d.ShouldCooldown) + require.True(t, d.ShouldFailover) +} + +func TestClassifyGrokUpstreamFailure_ValidationNoCool(t *testing.T) { + d := classifyGrokUpstreamFailure(http.StatusBadRequest, []byte(`{"error":{"message":"invalid tool schema"}}`), "") + require.Equal(t, GrokFailureNone, d.Class) + require.False(t, d.ShouldCooldown) + require.False(t, d.ShouldFailover) +} + +func TestClassifyGrokUpstreamFailure_FreeUsageWinsOver5xx(t *testing.T) { + // Proxy may rewrite free-usage into synthetic 502; body must win. + d := classifyGrokUpstreamFailure(http.StatusBadGateway, []byte(`subscription:free-usage-exhausted for model grok-4.3`), "grok-4.3") + require.Equal(t, GrokFailureFreeUsage, d.Class) + require.NotEqual(t, GrokFailureServer, d.Class) +} + +func TestShouldFailoverGrokUpstreamError_FreeUsageBody(t *testing.T) { + svc := &OpenAIGatewayService{} + body := []byte(`{"error":{"code":"subscription:free-usage-exhausted","message":"free usage exhausted"}}`) + require.True(t, svc.shouldFailoverGrokUpstreamError(http.StatusBadRequest, body)) +} + +func TestShouldFailoverGrokUpstreamError_ContentPolicyStillNoFailover(t *testing.T) { + svc := &OpenAIGatewayService{} + body := []byte(`{"error":{"code":"new_sensitive","message":"text is sensitive"}}`) + require.False(t, svc.shouldFailoverGrokUpstreamError(http.StatusForbidden, body)) +} + +func TestHandleGrokAccountUpstreamError_FreeUsageBodyCoolsAccount(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 9101, Platform: PlatformGrok, Type: AccountTypeOAuth} + before := time.Now() + body := []byte(`{"error":{"code":"subscription:free-usage-exhausted","message":"You've used all the included free usage. Usage resets over a rolling 24-hour window."}}`) + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusBadRequest, nil, body) + + require.Equal(t, 1, repo.tempUnschedCalls) + require.Equal(t, "grok free usage exhausted", repo.lastTempUnschedReason) + // Rolling-window exhaustion must use a short probe cooldown when no + // upstream absolute reset is available; it must not start a 24h lock here. + require.Greater(t, repo.lastTempUnschedUntil, before.Add(grokFreeUsageProbeCooldown-time.Second)) + require.Less(t, repo.lastTempUnschedUntil, before.Add(grokFreeUsageProbeCooldown+time.Second)) +} + +func TestHandleGrokAccountUpstreamError_FreeUsageUsesUpstreamReset(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 9102, Platform: PlatformGrok, Type: AccountTypeOAuth} + body := []byte(`{"error":{"code":"subscription:free-usage-exhausted","message":"free usage exhausted; rolling 24-hour window"}}`) + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusTooManyRequests, + http.Header{"Retry-After": []string{"3600"}}, body) + + require.Zero(t, repo.tempUnschedCalls) + require.WithinDuration(t, time.Now().Add(time.Hour), repo.lastRateLimitResetAt, 2*time.Second) +} + +func TestHandleGrokAccountUpstreamError_EmptyOutputCoolsAccount(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 9102, Platform: PlatformGrok, Type: AccountTypeOAuth} + before := time.Now() + + svc.handleGrokAccountUpstreamError( + context.Background(), account, http.StatusBadGateway, nil, + []byte(`empty model output: no content/tool_calls`), + ) + + require.Equal(t, 1, repo.tempUnschedCalls) + require.Equal(t, "grok empty model output", repo.lastTempUnschedReason) + require.WithinDuration(t, before.Add(4*time.Minute), repo.lastTempUnschedUntil, time.Second) +} + +func TestHandleGrokAccountUpstreamError_FreeUsageDoesNotCoolPoolMode(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ + ID: 9103, + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "pool_mode": true, + }, + } + body := []byte(`{"error":{"code":"subscription:free-usage-exhausted","message":"free usage exhausted"}}`) + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusBadRequest, nil, body) + + require.Zero(t, repo.tempUnschedCalls) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestHandleGrokAccountUpstreamError_ContentPolicyStillNoMutation(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 9104, Platform: PlatformGrok, Type: AccountTypeOAuth} + body := []byte(`{"error":{"code":"new_sensitive","message":"text is sensitive"}}`) + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusForbidden, nil, body) + + require.Zero(t, repo.tempUnschedCalls) +} + +func TestHandleGrokAccountUpstreamError_Entitlement403Unchanged(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 9105, Platform: PlatformGrok, Type: AccountTypeOAuth} + before := time.Now() + + svc.handleGrokAccountUpstreamError( + context.Background(), account, http.StatusForbidden, nil, + []byte(`{"error":{"message":"subscription required"}}`), + ) + + require.Equal(t, 1, repo.tempUnschedCalls) + require.Equal(t, "grok access or entitlement denied", repo.lastTempUnschedReason) + require.Greater(t, repo.lastTempUnschedUntil, before.Add(29*time.Minute)) + require.Less(t, repo.lastTempUnschedUntil, before.Add(31*time.Minute)) +} diff --git a/backend/internal/service/grok_upstream_headers.go b/backend/internal/service/grok_upstream_headers.go new file mode 100644 index 0000000000..cdfb54bcbf --- /dev/null +++ b/backend/internal/service/grok_upstream_headers.go @@ -0,0 +1,74 @@ +package service + +import ( + "net/http" + "strings" + + "github.com/gin-gonic/gin" + + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" +) + +// grokUpstreamUserAgent is kept for compatibility with older Grok request +// tests. Current requests use the pinned default UA from this package. +const grokUpstreamUserAgent = "sub2api-grok/1.0" + +// Fixed CLI identity aliases — single source of truth is internal/pkg/xai. +const ( + grokClientVersionHeader = xai.CLIStableVersion + grokClientIdentifierHeader = xai.CLIClientIdentifier + grokClientModeHeader = xai.CLIClientMode +) + +// defaultGrokUpstreamUserAgent is the pinned Grok CLI / workspace UA. +// Grok upstream must not forward Claude Code / Codex / browser client UAs. +func defaultGrokUpstreamUserAgent() string { + return xai.CLIUserAgent(xai.ResolveCLIVersion()) +} + +func applyDefaultGrokUpstreamHeaders(req *http.Request) { + if req == nil { + return + } + // Always stamp CLI identity. Do not preserve inbound client UA (Claude Code, + // Codex, curl, etc.) — xAI chat/CLI surfaces fingerprint the client string. + req.Header.Set("User-Agent", defaultGrokUpstreamUserAgent()) + req.Header.Set("x-grok-client-version", xai.ResolveCLIVersion()) + req.Header.Set("x-grok-client-identifier", grokClientIdentifierHeader) +} + +func applyGrokTLSProfileHeaders(req *http.Request, profile *tlsfingerprint.Profile) { + // HEAD Profile is TLS-only (no HTTP UserAgent/Originator fields). Always stamp CLI identity. + applyDefaultGrokUpstreamHeaders(req) + _ = profile +} + +// openAITLSFingerprintRuntime is the resolved TLS fingerprint routing result +// used by OpenAI/Grok outbound header application. Defined here so Grok header +// helpers compile even when the full OpenAI TLS router is not present on HEAD. +type openAITLSFingerprintRuntime struct { + Profile *tlsfingerprint.Profile + UpstreamUserAgent string + UpstreamOriginator string + Matched bool +} + +func applyGrokRuntimeHeaders(req *http.Request, runtime openAITLSFingerprintRuntime) { + applyDefaultGrokUpstreamHeaders(req) + if req == nil { + return + } + // Apply Originator only; force CLI UA after so router overrides cannot + // leak Codex/Claude Code identity onto Grok upstream. + if originator := strings.TrimSpace(runtime.UpstreamOriginator); originator != "" { + req.Header.Set("Originator", originator) + } + req.Header.Set("User-Agent", defaultGrokUpstreamUserAgent()) +} + +// resolveGrokUpstreamUserAgent always returns the pinned Grok CLI User-Agent. +// Inbound client UAs (Claude Code, Codex, browsers, libraries) are never forwarded. +func resolveGrokUpstreamUserAgent(_ *gin.Context) string { + return defaultGrokUpstreamUserAgent() +} diff --git a/backend/internal/service/grok_upstream_headers_test.go b/backend/internal/service/grok_upstream_headers_test.go new file mode 100644 index 0000000000..18972f4801 --- /dev/null +++ b/backend/internal/service/grok_upstream_headers_test.go @@ -0,0 +1,86 @@ +package service + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" +) + +func TestApplyDefaultGrokUpstreamHeadersUsesCLIUserAgent(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "") + + req, err := http.NewRequest(http.MethodGet, "https://api.x.ai/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "claude-code/1.2.3") + req.Header.Set("x-grok-client-version", "none") + + applyDefaultGrokUpstreamHeaders(req) + + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), req.Header.Get("User-Agent")) + require.Equal(t, xai.CLIClientVersion, req.Header.Get("x-grok-client-version")) + require.Equal(t, xai.CLIClientIdentifier, req.Header.Get("x-grok-client-identifier")) +} + +func TestApplyDefaultGrokUpstreamHeadersHonorsCLIVersionOverride(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "0.2.95") + + req, err := http.NewRequest(http.MethodGet, "https://api.x.ai/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "codex_cli_rs/0.144.0") + + applyDefaultGrokUpstreamHeaders(req) + + require.Equal(t, "0.2.95", req.Header.Get("x-grok-client-version")) + require.Equal(t, xai.CLIUserAgent("0.2.95"), req.Header.Get("User-Agent")) + require.Equal(t, "grok-shell", req.Header.Get("x-grok-client-identifier")) +} + +func TestResolveGrokUpstreamUserAgentNeverPassthrough(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "") + gin.SetMode(gin.TestMode) + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + c.Request.Header.Set("User-Agent", "claude-cli/2.0.0 (Mac OS; arm64)") + + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), resolveGrokUpstreamUserAgent(c)) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), resolveGrokUpstreamUserAgent(nil)) +} + +func TestApplyGrokRuntimeHeadersKeepsCLIUserAgent(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "") + + req, err := http.NewRequest(http.MethodPost, "https://cli-chat-proxy.grok.com/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "claude-code/9.9.9") + + applyGrokRuntimeHeaders(req, openAITLSFingerprintRuntime{ + UpstreamUserAgent: "codex_cli_rs/0.144.0", + UpstreamOriginator: "codex_cli_rs", + }) + + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), req.Header.Get("User-Agent")) + require.Equal(t, "codex_cli_rs", req.Header.Get("Originator")) + require.Equal(t, xai.CLIClientVersion, req.Header.Get("x-grok-client-version")) +} + +func TestApplyGrokTLSProfileHeadersAlwaysUsesCLIUserAgent(t *testing.T) { + t.Setenv(xai.CLIVersionEnv, "") + + req, err := http.NewRequest(http.MethodPost, "https://api.x.ai/v1/responses", nil) + require.NoError(t, err) + req.Header.Set("User-Agent", "grok-native/1.0") + + // HEAD Profile is TLS-only; Originator/UserAgent HTTP fields are not present. + applyGrokTLSProfileHeaders(req, &tlsfingerprint.Profile{Name: "chrome"}) + + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), req.Header.Get("User-Agent")) + require.Equal(t, xai.CLIClientVersion, req.Header.Get("x-grok-client-version")) +} diff --git a/backend/internal/service/grok_upstream_url.go b/backend/internal/service/grok_upstream_url.go index 362e746e6f..9574b3adbe 100644 --- a/backend/internal/service/grok_upstream_url.go +++ b/backend/internal/service/grok_upstream_url.go @@ -1,8 +1,11 @@ package service import ( + "context" "errors" "fmt" + "net/url" + "strings" "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" @@ -67,20 +70,28 @@ func redactedGrokBaseURLValidator(validator xai.BaseURLValidator) xai.BaseURLVal } } -func buildGrokResponsesURL(account *Account, cfg *config.Config) (string, error) { +func buildGrokResponsesURL(account *Account, cfg *config.Config, settings ...*SettingService) (string, error) { validator, err := grokBaseURLValidator(account, cfg) if err != nil { return "", err } - return xai.BuildResponsesURLWithValidator(account.GetGrokBaseURL(), validator) + baseURL := account.GetGrokBaseURL() + if len(settings) > 0 && settings[0] != nil { + baseURL = settings[0].ResolveGrokBaseURL(context.Background(), account) + } + return xai.BuildResponsesURLWithValidator(baseURL, validator) } -func buildGrokChatCompletionsURL(account *Account, cfg *config.Config) (string, error) { +func buildGrokChatCompletionsURL(account *Account, cfg *config.Config, settings ...*SettingService) (string, error) { validator, err := grokBaseURLValidator(account, cfg) if err != nil { return "", err } - return xai.BuildChatCompletionsURLWithValidator(account.GetGrokBaseURL(), validator) + baseURL := account.GetGrokBaseURL() + if len(settings) > 0 && settings[0] != nil { + baseURL = settings[0].ResolveGrokBaseURL(context.Background(), account) + } + return xai.BuildChatCompletionsURLWithValidator(baseURL, validator) } // buildGrokBillingURL 解析 billing 探测端点:跟随账号的转发 base_url, @@ -90,7 +101,14 @@ func buildGrokBillingURL(account *Account, cfg *config.Config, weekly bool) (str if err != nil { return "", err } - return xai.BuildBillingURLWithValidator(account.GetGrokBaseURL(), weekly, validator) + baseURL := account.GetGrokBaseURL() + // Official public/regional API hosts do not expose Grok Build billing. + // Keep custom relays on their configured host because they may proxy the CLI + // billing path alongside inference. + if xai.IsOfficialBaseURL(baseURL) && !isGrokCLIProxyBaseURL(baseURL) { + baseURL = xai.DefaultCLIBaseURL + } + return xai.BuildBillingURLWithValidator(baseURL, weekly, validator) } func buildGrokMediaURL(account *Account, cfg *config.Config, endpoint GrokMediaEndpoint, requestID string) (string, error) { @@ -122,3 +140,42 @@ func buildGrokMediaURL(account *Account, cfg *config.Config, endpoint GrokMediaE return "", fmt.Errorf("unsupported grok media endpoint: %s", endpoint) } } + +// buildGrokVoiceURL returns the official xAI Voice API endpoint. +// Voice HTTP (/tts, /stt, /custom-voices) and WS (/realtime) are only exposed +// by api.x.ai — the CLI chat proxy does not implement them. When the account +// base_url points at the CLI proxy (or is empty), fall back to DefaultBaseURL. +func buildGrokVoiceURL(account *Account, cfg *config.Config, endpoint string) (string, error) { + validator, err := grokBaseURLValidator(account, cfg) + if err != nil { + return "", err + } + base := "" + if account != nil { + base = account.GetGrokMediaBaseURL() + } + if strings.TrimSpace(base) == "" || isGrokCLIProxyBaseURL(base) { + base = xai.DefaultBaseURL + } + validated, err := validator(base) + if err != nil { + return "", err + } + ep := strings.Trim(strings.TrimSpace(endpoint), "/") + if ep == "" { + return "", fmt.Errorf("voice endpoint is required") + } + parts := strings.Split(ep, "/") + encoded := make([]string, 0, len(parts)) + for _, part := range parts { + if strings.TrimSpace(part) == "" || part == "." || part == ".." { + return "", fmt.Errorf("invalid voice endpoint path") + } + encoded = append(encoded, url.PathEscape(part)) + } + return strings.TrimRight(validated, "/") + "/" + strings.Join(encoded, "/"), nil +} + +func isGrokCLIProxyBaseURL(raw string) bool { + return isGrokCLIProxyTarget(raw) +} diff --git a/backend/internal/service/grok_upstream_url_test.go b/backend/internal/service/grok_upstream_url_test.go index e2f6e8bce5..8544823b4a 100644 --- a/backend/internal/service/grok_upstream_url_test.go +++ b/backend/internal/service/grok_upstream_url_test.go @@ -255,6 +255,26 @@ func TestGrokOAuthURLPolicy(t *testing.T) { }) } +func TestBuildGrokBillingURLUsesCLIForOfficialAPIHosts(t *testing.T) { + for _, baseURL := range []string{xai.DefaultBaseURL, "https://us-west-2.api.x.ai/v1"} { + account := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"base_url": baseURL}} + + weekly, err := buildGrokBillingURL(account, &config.Config{}, true) + require.NoError(t, err) + require.Equal(t, xai.DefaultCLIBaseURL+xai.BillingWeeklyPath, weekly) + } +} + +func TestBuildGrokBillingURLKeepsCustomRelay(t *testing.T) { + account := &Account{Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{ + "base_url": "https://relay.example.test/xai/v1", + }} + + monthly, err := buildGrokBillingURL(account, &config.Config{}, false) + require.NoError(t, err) + require.Equal(t, "https://relay.example.test/xai/v1"+xai.BillingMonthlyPath, monthly) +} + func TestGrokBillingURLFollowsAccountBaseURL(t *testing.T) { t.Run("oauth default stays on CLI gateway", func(t *testing.T) { account := &Account{ diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index 2fea013aff..ddde276e22 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -55,10 +55,21 @@ type Group struct { VideoPrice480P *float64 VideoPrice720P *float64 VideoPrice1080P *float64 + // VideoModelPrices is optional per-model-family per-second pricing + // (groups.video_model_prices JSONB). Shape: family → resolution → USD/s. + // When set for a model, overrides VideoPrice* for that model only. + VideoModelPrices map[string]map[string]float64 // Codex alpha/search 网页搜索单次价格(USD/次,仅 openai 平台使用); // nil 表示使用默认价 defaultWebSearchPricePerCall(官方 $10/1000 次)。 WebSearchPricePerCall *float64 + // 搜索工具显式定价(per 1k calls)。 + SearchPricePer1k *float64 + // Grok Voice 显式定价(分组级,不按文本 RateMultiplier)。 + AudioRealtimePricePerMin *float64 + AudioTTSPricePerMillionChars *float64 + AudioSTTPricePerHour *float64 + // Claude Code 客户端限制 ClaudeCodeOnly bool FallbackGroupID *int64 @@ -168,6 +179,30 @@ func (g *Group) GetVideoPrice(resolution string) *float64 { } } +// GetVideoPriceForModel prefers VideoModelPrices for the model family, then flat columns. +func (g *Group) GetVideoPriceForModel(model, resolution string) *float64 { + if g == nil { + return nil + } + if price := LookupVideoModelPrice(g.VideoModelPrices, model, resolution); price != nil { + return price + } + return g.GetVideoPrice(resolution) +} + +// VideoPriceConfig builds billing config including optional per-model map. +func (g *Group) VideoPriceConfig() *VideoPriceConfig { + if g == nil { + return nil + } + return &VideoPriceConfig{ + Price480P: g.VideoPrice480P, + Price720P: g.VideoPrice720P, + Price1080P: g.VideoPrice1080P, + ModelPrices: NormalizeVideoModelPrices(g.VideoModelPrices), + } +} + // IsGroupContextValid reports whether a group from context has the fields required for routing decisions. func IsGroupContextValid(group *Group) bool { if group == nil { @@ -414,3 +449,11 @@ func profitControlPlatformSupported(platform string) bool { return false } } + +// GetSearchPricePer1k returns explicit search/tool price per 1k calls if configured. +func (g *Group) GetSearchPricePer1k() *float64 { + if g == nil { + return nil + } + return g.SearchPricePer1k +} diff --git a/backend/internal/service/media_price_config.go b/backend/internal/service/media_price_config.go index 0a583b6c72..53d63c97a4 100644 --- a/backend/internal/service/media_price_config.go +++ b/backend/internal/service/media_price_config.go @@ -20,14 +20,15 @@ func videoPriceConfigFromAPIKey(apiKey *APIKey) *VideoPriceConfig { return nil } return &VideoPriceConfig{ - Price480P: apiKey.Group.VideoPrice480P, - Price720P: apiKey.Group.VideoPrice720P, - Price1080P: apiKey.Group.VideoPrice1080P, + Price480P: apiKey.Group.VideoPrice480P, + Price720P: apiKey.Group.VideoPrice720P, + Price1080P: apiKey.Group.VideoPrice1080P, + ModelPrices: apiKey.Group.VideoModelPrices, } } -func apiKeyHasConfiguredVideoPrice(apiKey *APIKey, resolution string) bool { - return apiKey != nil && apiKey.Group != nil && apiKey.Group.GetVideoPrice(resolution) != nil +func apiKeyHasConfiguredVideoPrice(apiKey *APIKey, model, resolution string) bool { + return apiKey != nil && apiKey.Group != nil && apiKey.Group.GetVideoPriceForModel(model, resolution) != nil } func webSearchPricePerCallFromAPIKey(apiKey *APIKey) *float64 { @@ -36,3 +37,22 @@ func webSearchPricePerCallFromAPIKey(apiKey *APIKey) *float64 { } return apiKey.Group.WebSearchPricePerCall } + +func groupSearchPricePer1kFromAPIKey(apiKey *APIKey) *float64 { + if apiKey == nil || apiKey.Group == nil { + return nil + } + return apiKey.Group.GetSearchPricePer1k() +} + +func groupAudioPriceConfigFromAPIKey(apiKey *APIKey) *audioPriceConfig { + if apiKey == nil || apiKey.Group == nil { + return nil + } + g := apiKey.Group + return &audioPriceConfig{ + RealtimePerMin: g.AudioRealtimePricePerMin, + TTSPerMChars: g.AudioTTSPricePerMillionChars, + STTPerHour: g.AudioSTTPricePerHour, + } +} diff --git a/backend/internal/service/model_rate_limit.go b/backend/internal/service/model_rate_limit.go index 195f5a1dfd..fe962b17c0 100644 --- a/backend/internal/service/model_rate_limit.go +++ b/backend/internal/service/model_rate_limit.go @@ -120,6 +120,23 @@ func OpenAIImageGenerationIntentFromContext(ctx context.Context) bool { return ok && enabled } +// WithOpenAIImagesEndpoint 标记请求从 /v1/images/* 专用生图端点入站。 +func WithOpenAIImagesEndpoint(ctx context.Context) context.Context { + if ctx == nil { + ctx = context.Background() + } + return context.WithValue(ctx, ctxkey.OpenAIImagesEndpoint, true) +} + +// OpenAIImagesEndpointFromContext 报告请求是否来自 /v1/images/*。 +func OpenAIImagesEndpointFromContext(ctx context.Context) bool { + if ctx == nil { + return false + } + enabled, ok := ctx.Value(ctxkey.OpenAIImagesEndpoint).(bool) + return ok && enabled +} + func resolveFinalAntigravityModelKey(ctx context.Context, account *Account, requestedModel string) string { modelKey := mapAntigravityModel(account, requestedModel) if modelKey == "" { diff --git a/backend/internal/service/oauth_service.go b/backend/internal/service/oauth_service.go index 0b3888a73f..b92d5cc009 100644 --- a/backend/internal/service/oauth_service.go +++ b/backend/internal/service/oauth_service.go @@ -22,6 +22,9 @@ type OpenAIOAuthClient interface { type GrokOAuthClient interface { ExchangeCode(ctx context.Context, code, codeVerifier, redirectURI, proxyURL, clientID string) (*xai.TokenResponse, error) RefreshToken(ctx context.Context, refreshToken, proxyURL, clientID string) (*xai.TokenResponse, error) + // LoginWithPassword exchanges email/password for a short-lived Web SSO cookie. + // Callers must convert via ConvertSSOToBuild and must not persist password or raw SSO. + LoginWithPassword(ctx context.Context, email, password, proxyURL string) (*GrokPasswordLoginResult, error) ConvertSSOToBuild(ctx context.Context, ssoToken, proxyURL string) (*xai.TokenResponse, error) } diff --git a/backend/internal/service/openai_account_scheduler.go b/backend/internal/service/openai_account_scheduler.go index 497c183eb9..6457ef8fc2 100644 --- a/backend/internal/service/openai_account_scheduler.go +++ b/backend/internal/service/openai_account_scheduler.go @@ -284,9 +284,10 @@ func (s *openAIAccountRuntimeStats) size() int { } type defaultOpenAIAccountScheduler struct { - service *OpenAIGatewayService - metrics openAIAccountSchedulerMetrics - stats *openAIAccountRuntimeStats + service *OpenAIGatewayService + metrics openAIAccountSchedulerMetrics + stats *openAIAccountRuntimeStats + grokFreeQuotaGateCache sync.Map // key: int64(accountID), value: grokFreeQuotaGateCacheEntry } type openAISelectionProbeBudget struct { @@ -499,6 +500,23 @@ func (s *defaultOpenAIAccountScheduler) selectBySessionHash( _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) return nil, false, nil } + // Free-tier soft gate: sticky session must not pin an over-quota free OAuth account. + // Admin QueryQuota / import probes do not use this path. + if account != nil && len(s.filterGrokFreeQuotaAccounts(ctx, []Account{*account})) == 0 { + _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) + return nil, false, nil + } + // Team+model cool: sticky must not pin a sibling under the same team 429 window. + now := time.Now() + upstreamModel := canonicalOpenAIAccountSchedulingModel(account, req.RequestedModel) + if account != nil && isGrokTeamModelRateLimited(account, upstreamModel, now) { + _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) + return nil, false, nil + } + if account != nil && isGrokModelQuotaBlocked(account.ID, upstreamModel, now) { + _ = s.service.deleteStickySessionAccountID(ctx, req.GroupID, sessionHash) + return nil, false, nil + } escapeCfg := s.service.openAIStickyEscapeConfig() if reason, errorRate, ttft, shouldEscape := s.shouldEscapeStickyAccount(accountID, escapeCfg); shouldEscape { slog.Info("sticky_escape_triggered", @@ -1241,6 +1259,18 @@ func (s *defaultOpenAIAccountScheduler) tryFallbackToWeightedSticky( if req.RequireCompact && openAICompactSupportTier(account) == 0 { continue } + // Keep weighted sticky fallback subject to the same free-tier gate as the + // normal and sticky selection paths. Otherwise an over-quota free account + // could be reintroduced after the primary candidate pass. + if len(s.filterGrokFreeQuotaAccounts(ctx, []Account{*account})) == 0 { + continue + } + upstreamModel := canonicalOpenAIAccountSchedulingModel(account, req.RequestedModel) + now := time.Now() + if isGrokTeamModelRateLimited(account, upstreamModel, now) || + isGrokModelQuotaBlocked(account.ID, upstreamModel, now) { + continue + } result, acquireErr := s.service.tryAcquireAccountSlot(ctx, account.ID, account.Concurrency) if acquireErr != nil { return nil, acquireErr @@ -1330,6 +1360,28 @@ func (s *defaultOpenAIAccountScheduler) selectByLoadBalance( if len(accounts) == 0 { return nil, 0, 0, 0, noAvailableOpenAISelectionError(req.RequestedModel, false, openAISelectionFilterStats{}.summary("")) } + // Local free-tier soft gate on the Grok scheduling path only (not admin probe). + accounts = s.filterGrokFreeQuotaAccounts(ctx, accounts) + if len(accounts) == 0 { + return nil, 0, 0, 0, noAvailableOpenAISelectionError(req.RequestedModel, false, openAISelectionFilterStats{}.summary("grok_free_quota_soft_gate")) + } + // Team+model rate-limit cool: siblings of a 429'd team skip the hot model. + if req.Platform == PlatformGrok { + now := time.Now() + filtered := filterGrokTeamModelRateLimitedAccounts(accounts, req.RequestedModel, now) + if len(filtered) == 0 && len(accounts) > 0 { + return nil, 0, 0, 0, noAvailableOpenAISelectionError(req.RequestedModel, false, openAISelectionFilterStats{}.summary("grok_team_model_rate_limit")) + } + if filtered != nil { + accounts = filtered + } + // Per-account model free-usage soft-block (other models stay eligible). + modelFiltered := filterGrokModelQuotaBlockedAccounts(accounts, req.RequestedModel, now) + if len(modelFiltered) == 0 && len(accounts) > 0 { + return nil, 0, 0, 0, noAvailableOpenAISelectionError(req.RequestedModel, false, openAISelectionFilterStats{}.summary("grok_model_quota_block")) + } + accounts = modelFiltered + } // require_privacy_set: 获取分组信息 var schedGroup *Group diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index 03f30c5b31..3b3a64856e 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -182,6 +182,20 @@ func (c *schedulerTestGatewayCache) DeleteSessionAccountID(ctx context.Context, return nil } +func (c *schedulerTestGatewayCache) SetGrokVideoPendingBilling(_ context.Context, _ string, _ []byte, _ time.Duration) error { + return nil +} +func (c *schedulerTestGatewayCache) GetGrokVideoPendingBilling(_ context.Context, _ string) ([]byte, error) { + return nil, nil +} +func (c *schedulerTestGatewayCache) ClaimGrokVideoBilled(_ context.Context, _ string, _ time.Duration) (bool, error) { + return true, nil +} + +func (c *schedulerTestGatewayCache) ReleaseGrokVideoBilled(_ context.Context, _ string) error { + return nil +} + func newSchedulerTestOpenAIWSV2Config() *config.Config { cfg := &config.Config{} cfg.Gateway.OpenAIWS.Enabled = true diff --git a/backend/internal/service/openai_capacity_shed_test.go b/backend/internal/service/openai_capacity_shed_test.go index 83e4f8b339..e68a48db1c 100644 --- a/backend/internal/service/openai_capacity_shed_test.go +++ b/backend/internal/service/openai_capacity_shed_test.go @@ -2,12 +2,17 @@ package service import ( "context" + "errors" + "io" "net/http" + "net/http/httptest" "strings" "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" + "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" ) @@ -72,6 +77,172 @@ func TestStreamFailedEventCapacityShedRetriesOnSameAccount(t *testing.T) { require.False(t, openAIStreamFailedEventRetryableOnSameAccount(nonPool, other, "boom")) } +// 上游降载的真实序列是「event: error → event: response.failed」。error 帧不算 +// 客户端输出:若把它当首输出 flush,clientOutputStarted 被固化,随后的 failed +// 事件就进不了 pre-output failover 分支,只能把致命错误原样转发给客户端。 +func TestOpenAIStreamErrorFrameDoesNotStartClientOutput(t *testing.T) { + cases := []struct { + data string + eventType string + want bool + }{ + {`{"type":"error","error":{"code":"server_is_overloaded","message":"overloaded"}}`, "error", false}, + {`{"type":"error","error":{"code":"slow_down","message":"slow down"}}`, "error", false}, + {`{"type":"error","error":{"type":"rate_limit_error","code":"rate_limit_exceeded","message":"limited"}}`, "error", false}, + // 不可重试类错误帧维持原样转发(不进 failover),保留上游错误细节。 + {`{"type":"error","error":{"type":"invalid_request_error","code":"content_policy_violation","message":"blocked"}}`, "error", true}, + {`{"type":"response.failed","response":{"error":{"code":"server_is_overloaded"}}}`, "response.failed", false}, + {`{"type":"response.created","response":{"id":"resp_1"}}`, "response.created", false}, + {`{"type":"response.in_progress","response":{"id":"resp_1"}}`, "response.in_progress", false}, + {`{"type":"response.output_text.delta","delta":"hi"}`, "response.output_text.delta", true}, + {`[DONE]`, "", true}, + } + for _, tc := range cases { + require.Equal(t, tc.want, openAIStreamDataStartsClientOutput(tc.data, tc.eventType), "data=%s type=%s", tc.data, tc.eventType) + } +} + +// 回归用例(真实上游降载序列):created → in_progress → error 帧 → response.failed。 +// 期望仍然走 pre-output failover(同账号重试 + 请求级瞬时标记),且不向客户端写出任何字节。 +func TestOpenAIStreamCapacityShedErrorFramePrecedingFailedStillFailsOver(t *testing.T) { + gin.SetMode(gin.TestMode) + cfg := &config.Config{ + Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, + } + svc := &OpenAIGatewayService{cfg: cfg} + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + + resp := &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(strings.Join([]string{ + "event: response.created", + `data: {"type":"response.created","response":{"id":"resp_1"},"sequence_number":0}`, + "", + "event: response.in_progress", + `data: {"type":"response.in_progress","response":{"id":"resp_1"},"sequence_number":1}`, + "", + "event: error", + `data: {"type":"error","error":{"type":"service_unavailable_error","code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."},"sequence_number":2}`, + "", + "event: response.failed", + `data: {"type":"response.failed","response":{"id":"resp_1","status":"failed","error":{"code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}},"sequence_number":3}`, + "", + }, "\n"))), + Header: http.Header{"X-Request-Id": []string{"rid-shed-error-then-failed"}}, + } + + _, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "acc"}, time.Now(), "model", "model") + require.Error(t, err) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.True(t, failoverErr.RetryableOnSameAccount) + require.True(t, failoverErr.RequestScopedTransient) + require.False(t, c.Writer.Written()) + require.Empty(t, rec.Body.String()) +} + +// 流中途(已有真实输出)降载时无法再 failover,此时必须把降载码改写为客户端 +// 可重试的 server_error 再转发——Codex 对 server_is_overloaded/slow_down 判致命 +// 并终止会话,对其余错误码执行内置退避重试。消息原样保留。 +func TestOpenAIStreamCapacityShedAfterOutputRewritesCodeForClient(t *testing.T) { + gin.SetMode(gin.TestMode) + cfg := &config.Config{ + Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}, + } + svc := &OpenAIGatewayService{cfg: cfg} + + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/", nil) + + resp := &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(strings.Join([]string{ + "event: response.created", + `data: {"type":"response.created","response":{"id":"resp_1"}}`, + "", + "event: response.output_text.delta", + `data: {"type":"response.output_text.delta","delta":"partial"}`, + "", + "event: error", + `data: {"type":"error","error":{"type":"service_unavailable_error","code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."},"sequence_number":2}`, + "", + "event: response.failed", + `data: {"type":"response.failed","response":{"id":"resp_1","status":"failed","error":{"code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}},"sequence_number":3}`, + "", + }, "\n"))), + Header: http.Header{"X-Request-Id": []string{"rid-shed-after-output"}}, + } + + _, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "acc"}, time.Now(), "model", "model") + require.Error(t, err) + var failoverErr *UpstreamFailoverError + require.False(t, errors.As(err, &failoverErr)) + + body := rec.Body.String() + require.Contains(t, body, "partial") + require.Contains(t, body, "event: response.failed") + require.Contains(t, body, `"code":"server_error"`) + require.NotContains(t, body, "server_is_overloaded") + require.Contains(t, body, "Our servers are currently overloaded") +} + +// helper 单测:只有降载码被改写,其余错误码(尤其 rate_limit_exceeded,客户端 +// 依赖其原码解析重试延时)必须原样保留。 +func TestSanitizeOpenAICapacityShedErrorCodeForClient(t *testing.T) { + cases := []struct { + name string + payload string + wantChanged bool + wantContain string + }{ + { + name: "failed事件嵌套code改写", + payload: `{"type":"response.failed","response":{"error":{"code":"server_is_overloaded","message":"overloaded"}}}`, + wantChanged: true, + wantContain: `"code":"server_error"`, + }, + { + name: "error帧裸code改写", + payload: `{"type":"error","error":{"code":"slow_down","message":"slow down"}}`, + wantChanged: true, + wantContain: `"code":"server_error"`, + }, + { + name: "rate_limit不改写", + payload: `{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded","message":"try again in 3s"}}}`, + wantChanged: false, + wantContain: `"code":"rate_limit_exceeded"`, + }, + { + name: "普通server_error不改写", + payload: `{"type":"response.failed","response":{"error":{"code":"server_error","message":"boom"}}}`, + wantChanged: false, + wantContain: `"code":"server_error"`, + }, + { + name: "非JSON不改写", + payload: `not-json`, + wantChanged: false, + wantContain: `not-json`, + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + out, changed := sanitizeOpenAICapacityShedErrorCodeForClient([]byte(tc.payload)) + require.Equal(t, tc.wantChanged, changed) + require.Contains(t, string(out), tc.wantContain) + if changed { + require.NotContains(t, string(out), "server_is_overloaded") + require.NotContains(t, string(out), "slow_down") + } + }) + } +} + // 出站身份的版本声明只能有一个来源:UA 的版本段、version 头、探针版本三处必须同源, // 各自硬编码会漂移成互相矛盾的身份,而自相矛盾或陈旧的身份会被上游优先降载。 func TestCodexOutboundVersionHasSingleSource(t *testing.T) { diff --git a/backend/internal/service/openai_compact_service_tier_test.go b/backend/internal/service/openai_compact_service_tier_test.go new file mode 100644 index 0000000000..7e92d8d2c2 --- /dev/null +++ b/backend/internal/service/openai_compact_service_tier_test.go @@ -0,0 +1,100 @@ +package service + +import ( + "bytes" + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestNormalizeOpenAICompactRequestBodyPreservesServiceTier(t *testing.T) { + body := []byte(`{ + "model":"gpt-5.6-sol", + "input":[{"type":"message","role":"user","content":"hello"}], + "service_tier":"priority", + "prompt_cache_key":"compact-cache-key", + "store":false, + "stream":true + }`) + + normalized, changed, err := normalizeOpenAICompactRequestBody(body) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(normalized, "model").String()) + require.Equal(t, "priority", gjson.GetBytes(normalized, "service_tier").String()) + require.False(t, gjson.GetBytes(normalized, "prompt_cache_key").Exists()) + require.False(t, gjson.GetBytes(normalized, "store").Exists()) + require.False(t, gjson.GetBytes(normalized, "stream").Exists()) +} + +func TestOpenAIOAuthCompactHTTPBuildersUsePreservedServiceTierInRoutingHint(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{ + "model":"gpt-5.6-sol", + "input":[{"type":"message","role":"user","content":"hello"}], + "service_tier":"priority", + "stream":true + }`) + normalized, changed, err := normalizeOpenAICompactRequestBody(body) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "priority", gjson.GetBytes(normalized, "service_tier").String()) + + account := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "test-account", + }, + } + svc := &OpenAIGatewayService{} + + tests := []struct { + name string + build func(*gin.Context) (*http.Request, error) + }{ + { + name: "ordinary", + build: func(c *gin.Context) (*http.Request, error) { + return svc.buildUpstreamRequest( + context.Background(), c, account, normalized, "test-token", + false, "", true, + ) + }, + }, + { + name: "passthrough", + build: func(c *gin.Context) (*http.Request, error) { + return svc.buildUpstreamRequestOpenAIPassthrough( + context.Background(), c, account, normalized, "test-token", + ) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest( + http.MethodPost, + "/v1/responses/compact", + bytes.NewReader(normalized), + ) + + req, buildErr := tt.build(c) + require.NoError(t, buildErr) + require.Equal( + t, + "model=gpt-5.6-sol;tier=priority", + req.Header.Get(openAICodexRoutingHintHeader), + ) + require.Equal(t, "priority", gjson.GetBytes(normalized, "service_tier").String()) + }) + } +} diff --git a/backend/internal/service/openai_cyber_session_block_test.go b/backend/internal/service/openai_cyber_session_block_test.go index 3957b824e5..6d6ad8bad4 100644 --- a/backend/internal/service/openai_cyber_session_block_test.go +++ b/backend/internal/service/openai_cyber_session_block_test.go @@ -133,6 +133,21 @@ func (c *comboCacheAndStore) RefreshSessionTTL(_ context.Context, _ int64, _ str func (c *comboCacheAndStore) DeleteSessionAccountID(_ context.Context, _ int64, _ string) error { return nil } + +func (c *comboCacheAndStore) SetGrokVideoPendingBilling(_ context.Context, _ string, _ []byte, _ time.Duration) error { + return nil +} +func (c *comboCacheAndStore) GetGrokVideoPendingBilling(_ context.Context, _ string) ([]byte, error) { + return nil, nil +} +func (c *comboCacheAndStore) ClaimGrokVideoBilled(_ context.Context, _ string, _ time.Duration) (bool, error) { + return true, nil +} + +func (c *comboCacheAndStore) ReleaseGrokVideoBilled(_ context.Context, _ string) error { + return nil +} + func (c *comboCacheAndStore) SetCyberSessionBlocked(ctx context.Context, key string, ttl time.Duration) error { return c.store.SetCyberSessionBlocked(ctx, key, ttl) } diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 171a6c8805..1a0702a87e 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -58,6 +58,8 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( promptCacheKey string, defaultMappedModel string, ) (*OpenAIForwardResult, error) { + beginUpstreamResponseModelObservation(c) + restrictionResult := s.detectCodexClientRestriction(c, account, body) logCodexCLIOnlyDetection(ctx, c, account, getAPIKeyIDFromContext(c), restrictionResult, body) if restrictionResult.Enabled && !restrictionResult.Matched { @@ -419,6 +421,11 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( writeChatCompletionsError(c, http.StatusBadGateway, "api_error", "Upstream stream ended without a terminal response event") return nil, fmt.Errorf("upstream stream ended without terminal event") } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + observer.Observe(finalResponse.Model, true) if strings.TrimSpace(finalResponse.Status) == "failed" { payload, _ := json.Marshal(gin.H{"type": "response.failed", "response": finalResponse}) // cyber_policy 致命不可重试:不 failover,以 Chat Completions 错误格式回写(F4), @@ -476,15 +483,26 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8") c.JSON(http.StatusOK, chatResp) - return &OpenAIForwardResult{ - RequestID: requestID, - Usage: usage, - Model: originalModel, - BillingModel: billingModel, - UpstreamModel: upstreamModel, - Stream: false, - Duration: time.Since(startTime), - }, nil + result := &OpenAIForwardResult{ + RequestID: requestID, + Usage: usage, + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: false, + Duration: time.Since(startTime), + } + // Grok chat bridge: bill native search tools found in the terminal Responses body. + if account != nil && account.IsGrok() && finalResponse != nil { + if body, err := json.Marshal(finalResponse); err == nil { + if n := countGrokNativeSearchCallsFromJSONBytes(body); n > 0 { + result.SearchCount = n + } + } + } + return result, nil } // handleChatStreamingResponse reads Responses SSE events from upstream, @@ -517,6 +535,14 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( refusalDetector := newOpenAIChatSilentRefusalDetector(requestBodyLen) var streamFailoverErr *UpstreamFailoverError var streamNonFailoverErr error + // Grok chat bridge reuses Responses SSE; count native search tools for surcharge. + searchCount := 0 + streamSearchSeen := make(map[string]struct{}) + countSearch := account != nil && account.IsGrok() + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } scanner := s.newUpstreamSSEScanner(resp.Body) @@ -535,16 +561,22 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( } resultWithUsage := func() *OpenAIForwardResult { - return &OpenAIForwardResult{ - RequestID: requestID, - Usage: usage, - Model: originalModel, - BillingModel: billingModel, - UpstreamModel: upstreamModel, - Stream: true, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, + out := &OpenAIForwardResult{ + RequestID: requestID, + Usage: usage, + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: true, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, } + if searchCount > 0 { + out.SearchCount = searchCount + } + return out } processDataLine := func(payload string) bool { @@ -553,6 +585,9 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( ms := int(time.Since(startTime).Milliseconds()) firstTokenMs = &ms } + if countSearch { + searchCount += countGrokNativeSearchCallsInSSEDataDedup([]byte(payload), streamSearchSeen) + } var event apicompat.ResponsesStreamEvent if err := json.Unmarshal([]byte(payload), &event); err != nil { @@ -562,6 +597,7 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( ) return false } + observer.ObserveOpenAI([]byte(payload), event.Type) refusalDetector.ObservePayload([]byte(payload)) isTerminalEvent := isOpenAICompatResponsesTerminalEvent(event.Type) diff --git a/backend/internal/service/openai_gateway_chat_completions_raw.go b/backend/internal/service/openai_gateway_chat_completions_raw.go index 49cd591829..d7eb8d6122 100644 --- a/backend/internal/service/openai_gateway_chat_completions_raw.go +++ b/backend/internal/service/openai_gateway_chat_completions_raw.go @@ -144,6 +144,10 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( if err != nil { return nil, fmt.Errorf("remove Responses-only Grok prompt cache key: %w", err) } + upstreamBody, err = normalizeGrokChatReasoningEffort(upstreamBody, upstreamModel) + if err != nil { + return nil, fmt.Errorf("normalize Grok chat reasoning effort: %w", err) + } } logger.L().Debug("openai chat_completions raw: forwarding without protocol conversion", @@ -225,7 +229,7 @@ func (s *OpenAIGatewayService) forwardAsRawChatCompletions( func (s *OpenAIGatewayService) rawChatCompletionsURL(account *Account) (string, error) { if account.Platform == PlatformGrok { - targetURL, err := buildGrokChatCompletionsURL(account, s.cfg) + targetURL, err := buildGrokChatCompletionsURL(account, s.cfg, s.settingService) if err != nil { return "", fmt.Errorf("invalid grok base_url: %w", err) } @@ -253,6 +257,10 @@ func (s *OpenAIGatewayService) streamRawChatCompletions( startTime time.Time, requestBodyLen int, ) (*OpenAIForwardResult, error) { + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } requestID := resp.Header.Get("x-request-id") writeStreamHeaders := s.newStreamHeaderWriter(c, resp.Header) scanner := s.newUpstreamSSEScanner(resp.Body) @@ -302,6 +310,7 @@ func (s *OpenAIGatewayService) streamRawChatCompletions( if payload, ok := extractOpenAISSEDataLine(line); ok { trimmedPayload := strings.TrimSpace(payload) if trimmedPayload != "[DONE]" { + observer.ObserveOpenAI([]byte(payload), strings.TrimSpace(gjson.Get(payload, "type").String())) usageOnlyChunk := isOpenAIChatUsageOnlyStreamChunk(payload) if u := extractCCStreamUsage(payload); u != nil { usage = *u @@ -356,16 +365,18 @@ func (s *OpenAIGatewayService) streamRawChatCompletions( } return &OpenAIForwardResult{ - RequestID: requestID, - Usage: usage, - Model: originalModel, - BillingModel: billingModel, - UpstreamModel: upstreamModel, - ReasoningEffort: reasoningEffort, - ServiceTier: serviceTier, - Stream: true, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, + RequestID: requestID, + Usage: usage, + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + ReasoningEffort: reasoningEffort, + ServiceTier: serviceTier, + Stream: true, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, }, nil } @@ -425,6 +436,11 @@ func (s *OpenAIGatewayService) bufferRawChatCompletions( } return nil, fmt.Errorf("read upstream body: %w", err) } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + observer.ObserveOpenAI(respBody, strings.TrimSpace(gjson.GetBytes(respBody, "type").String())) var usage OpenAIUsage if parsedUsage, ok := extractOpenAIUsageFromJSONBytes(respBody); ok { @@ -443,15 +459,17 @@ func (s *OpenAIGatewayService) bufferRawChatCompletions( _, _ = c.Writer.Write(respBody) return &OpenAIForwardResult{ - RequestID: requestID, - Usage: usage, - Model: originalModel, - BillingModel: billingModel, - UpstreamModel: upstreamModel, - ReasoningEffort: reasoningEffort, - ServiceTier: serviceTier, - Stream: false, - Duration: time.Since(startTime), + RequestID: requestID, + Usage: usage, + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + ReasoningEffort: reasoningEffort, + ServiceTier: serviceTier, + Stream: false, + Duration: time.Since(startTime), }, nil } diff --git a/backend/internal/service/openai_gateway_count_tokens.go b/backend/internal/service/openai_gateway_count_tokens.go index eccd1c4146..e56b33df45 100644 --- a/backend/internal/service/openai_gateway_count_tokens.go +++ b/backend/internal/service/openai_gateway_count_tokens.go @@ -345,12 +345,26 @@ func isOpenAIOAuthInputTokensUnsupported(statusCode int, body []byte) bool { return true } + // OAuth's platform endpoint can be blocked by an upstream proxy before it + // reaches the API and return an HTML 403 page without a structured error. + // Treat that endpoint-level response like the other unsupported cases so + // count_tokens remains a local, non-health-affecting convenience request. + if statusCode == http.StatusForbidden && isHTMLResponse(body) { + return true + } + return strings.Contains(msg, "input_tokens") && (strings.Contains(msg, "not found") || strings.Contains(msg, "not supported") || strings.Contains(msg, "unsupported")) } +func isHTMLResponse(body []byte) bool { + trimmed := strings.TrimSpace(strings.ToLower(string(body))) + return strings.HasPrefix(trimmed, "Forbidden", + }, { name: "404_input_tokens_unsupported", statusCode: http.StatusNotFound, @@ -306,6 +311,10 @@ func TestEstimateOpenAIInputTokens_CompareWithOpenAIAPI(t *testing.T) { if apiKey == "" { t.Skip("OPENAI_API_KEY not set") } + // Invalid/expired keys in local env must not fail the unit suite. + if strings.HasPrefix(apiKey, "sk-") && len(apiKey) < 20 { + t.Skip("OPENAI_API_KEY looks incomplete") + } client := &http.Client{Timeout: 30 * time.Second} cases := []struct { @@ -340,7 +349,13 @@ func TestEstimateOpenAIInputTokens_CompareWithOpenAIAPI(t *testing.T) { require.NoError(t, err) actual, err := callOpenAIInputTokensAPIForTest(client, apiKey, prepared.Request) - require.NoError(t, err) + if err != nil { + // Live-API comparison only; invalid/expired local keys should skip, not fail CI. + if strings.Contains(err.Error(), "status=401") || strings.Contains(err.Error(), "invalid_api_key") { + t.Skipf("OPENAI_API_KEY rejected by OpenAI: %v", err) + } + require.NoError(t, err) + } diff := estimated - actual if diff < 0 { diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 0484ea7d36..34c92dfe89 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -19,6 +19,7 @@ import ( // Forward forwards request to OpenAI API func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, account *Account, body []byte) (*OpenAIForwardResult, error) { + beginUpstreamResponseModelObservation(c) clearGrokResponsesClientToolMapping(c) clearOpenAIResponsesNamespaceNames(c) startTime := time.Now() @@ -46,6 +47,16 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if normalized { body = normalizedBody } + // 在分流到 passthrough / Codex transform / 原生 ChatCompletions 之前统一修正 + // 显式为 null 的工具 Schema type,否则 upstream 的 400 会被归一成可重试的 502, + // 同一份坏定义在账号池里反复重放。 + sanitizedToolBody, toolSchemaSanitized, toolSchemaErr := sanitizeOpenAIResponsesToolParameterTypes(body) + if toolSchemaErr != nil { + return nil, fmt.Errorf("sanitize OpenAI Responses tool parameters: %w", toolSchemaErr) + } + if toolSchemaSanitized { + body = sanitizedToolBody + } if account.IsOpenAIOAuth() && isOpenAIResponsesLiteHeader(c.GetHeader(responsesLiteHeader)) { liteBody, changed, liteErr := normalizeOpenAIResponsesLiteToolsPayload(body) if liteErr != nil { @@ -928,6 +939,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco var firstTokenMs *int responseID := "" imageCount := 0 + searchCount := 0 var imageOutputSizes []string if reqStream { streamResult, err := s.handleStreamingResponseWithReasoning(ctx, resp, c, account, startTime, originalModel, upstreamModel, reasoningEffortValue) @@ -939,6 +951,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco responseID = strings.TrimSpace(streamResult.responseID) imageCount = streamResult.imageCount imageOutputSizes = streamResult.imageOutputSizes + searchCount = streamResult.searchCount } else { nonStreamResult, err := s.handleNonStreamingResponse(ctx, resp, c, account, originalModel, upstreamModel) if err != nil { @@ -948,6 +961,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco responseID = strings.TrimSpace(nonStreamResult.responseID) imageCount = nonStreamResult.imageCount imageOutputSizes = nonStreamResult.imageOutputSizes + searchCount = nonStreamResult.searchCount } s.bindHTTPResponseAccount(ctx, c, account, responseID) @@ -964,18 +978,20 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } forwardResult := &OpenAIForwardResult{ - RequestID: resp.Header.Get("x-request-id"), - ResponseID: responseID, - Usage: *usage, - Model: originalModel, - BillingModel: billingModel, - UpstreamModel: upstreamModel, - ServiceTier: serviceTier, - ReasoningEffort: reasoningEffort, - Stream: reqStream, - OpenAIWSMode: false, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, + RequestID: resp.Header.Get("x-request-id"), + ResponseID: responseID, + Usage: *usage, + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + ServiceTier: serviceTier, + ReasoningEffort: reasoningEffort, + Stream: reqStream, + OpenAIWSMode: false, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, } if imageCount > 0 { forwardResult.ImageCount = imageCount @@ -984,6 +1000,12 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco forwardResult.ImageOutputSizes = imageOutputSizes forwardResult.BillingModel = imageBillingModel } + // Grok-native web_search / x_search / tool_search tool invocations (per-1k pricing). + // Token cost still applies separately when usage is present; search is additive only + // when search_price_per_1k is configured (nil price → $0 from CalculateSearchCost). + if searchCount > 0 && account != nil && account.IsGrok() { + forwardResult.SearchCount = searchCount + } return forwardResult, nil } } @@ -1059,7 +1081,6 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin. req.Header.Del("OpenAI-Beta") req.Header.Del("originator") } else { - req.Header.Set("OpenAI-Beta", "responses=experimental") req.Header.Set("originator", resolveOpenAIUpstreamOriginator(c, isCodexCLI)) } apiKeyID := getAPIKeyIDFromContext(c) @@ -1111,6 +1132,8 @@ func (s *OpenAIGatewayService) buildUpstreamRequest(ctx context.Context, c *gin. // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) account.ApplyHeaderOverrides(req.Header) + setOpenAICodexRoutingHintFromBody(req.Header, account, body) + logOpenAIRoutingDiagnosticsFromBody(ctx, account, "http", req.Header, body, "not_applicable") return req, nil } diff --git a/backend/internal/service/openai_gateway_grok.go b/backend/internal/service/openai_gateway_grok.go index 4d367c9aeb..3a737f72d5 100644 --- a/backend/internal/service/openai_gateway_grok.go +++ b/backend/internal/service/openai_gateway_grok.go @@ -22,14 +22,14 @@ import ( const ( grokComposerImageBridgeVisionModel = "grok-build-0.1" grokComposerImageBridgeMaxOutputTokens = 512 - grokUpstreamUserAgent = "sub2api-grok/1.0" - grokCLIVersion = xai.CLIClientVersion - grokDefaultResponsesModel = "grok-4.5" - grokRateLimitFallbackCooldown = 2 * time.Minute - grokRateLimitRepeatCooldown = 10 * time.Minute - grokRateLimitSustainedCooldown = 30 * time.Minute - grokRateLimitMaxAdaptiveCooldown = time.Hour - grokRateLimitBackoffQuietPeriod = time.Hour + // grokUpstreamUserAgent lives in grok_upstream_headers.go (shared with TLS header helpers). + grokCLIVersion = xai.CLIClientVersion + grokDefaultResponsesModel = "grok-4.5" + grokRateLimitFallbackCooldown = 2 * time.Minute + grokRateLimitRepeatCooldown = 10 * time.Minute + grokRateLimitSustainedCooldown = 30 * time.Minute + grokRateLimitMaxAdaptiveCooldown = time.Hour + grokRateLimitBackoffQuietPeriod = time.Hour ) func (s *OpenAIGatewayService) forwardGrokResponses( @@ -49,6 +49,10 @@ func (s *OpenAIGatewayService) forwardGrokResponses( if strings.TrimSpace(upstreamModel) == "" { upstreamModel = grokDefaultResponsesModel } + // Account mappings are optional. Canonicalize client aliases even when the + // account has no model_mapping, matching the Chat Completions path and xAI's + // actual Responses model IDs. + upstreamModel = xai.ResolveGrokTextResponsesModelID(upstreamModel, grokDefaultResponsesModel) if isGrokImageGenerationModel(upstreamModel) { return nil, fmt.Errorf("model %s is an image model and is not available on the Responses endpoint; use /v1/images/generations instead", upstreamModel) } @@ -101,7 +105,7 @@ func (s *OpenAIGatewayService) forwardGrokResponses( upstreamStart := time.Now() var resp *http.Response for attempt := 0; ; attempt++ { - upstreamReq, buildErr := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token, cacheIdentity, s.cfg) + upstreamReq, buildErr := buildGrokResponsesRequest(upstreamCtx, c, account, patchedBody, token, cacheIdentity, s.cfg, s.settingService) if buildErr != nil { return nil, buildErr } @@ -161,7 +165,13 @@ func (s *OpenAIGatewayService) forwardGrokResponses( Kind: kind, Message: upstreamMsg, }) - s.handleGrokAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody) + errCtx := withGrokTeamRateLimitModel(ctx, upstreamModel) + s.handleGrokAccountUpstreamError(errCtx, account, resp.StatusCode, resp.Header, respBody) + // 429 / free-usage: stamp team+model cool so sibling accounts skip this model. + if resp.StatusCode == http.StatusTooManyRequests || + classifyGrokUpstreamFailure(resp.StatusCode, respBody, upstreamModel).Class == GrokFailureFreeUsage { + markGrokTeamModelRateLimit(account, upstreamModel, resolveGrokTeamRateLimitUntil(time.Now().Add(grokTeamRateLimitDefaultTTL), time.Now())) + } if s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) { return nil, &UpstreamFailoverError{ StatusCode: resp.StatusCode, @@ -173,11 +183,16 @@ func (s *OpenAIGatewayService) forwardGrokResponses( return s.handleErrorResponse(ctx, resp, c, account, patchedBody, upstreamModel) } - s.updateGrokUsageFromResponse(ctx, account, resp.Header, resp.StatusCode) + // Attach model so rate-limit snapshots can fan out a team+model cool. + stateCtx := withGrokTeamRateLimitModel(ctx, upstreamModel) + s.updateGrokUsageFromResponse(stateCtx, account, resp.Header, resp.StatusCode) var usage *OpenAIUsage var firstTokenMs *int responseID := "" + searchCount := 0 + imageCount := 0 + var imageOutputSizes []string if reqStream { maxLineSize := defaultMaxLineSize if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { @@ -194,6 +209,9 @@ func (s *OpenAIGatewayService) forwardGrokResponses( usage = streamResult.usage firstTokenMs = streamResult.firstTokenMs responseID = strings.TrimSpace(streamResult.responseID) + searchCount = streamResult.searchCount + imageCount = streamResult.imageCount + imageOutputSizes = streamResult.imageOutputSizes } else { nonStreamResult, err := s.handleNonStreamingResponse(ctx, resp, c, account, originalModel, upstreamModel) if err != nil { @@ -201,13 +219,16 @@ func (s *OpenAIGatewayService) forwardGrokResponses( } usage = nonStreamResult.usage responseID = strings.TrimSpace(nonStreamResult.responseID) + searchCount = nonStreamResult.searchCount + imageCount = nonStreamResult.imageCount + imageOutputSizes = nonStreamResult.imageOutputSizes } if usage == nil { usage = &OpenAIUsage{} } reasoningEffort := extractOpenAIReasoningEffortFromBody(patchedBody, originalModel) - return &OpenAIForwardResult{ + result := &OpenAIForwardResult{ RequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")), ResponseID: responseID, Usage: *usage, @@ -219,7 +240,17 @@ func (s *OpenAIGatewayService) forwardGrokResponses( ResponseHeaders: resp.Header.Clone(), Duration: time.Since(startTime), FirstTokenMs: firstTokenMs, - }, nil + } + // Propagate search/image counters from the shared Responses handler — without + // this, stream/JSON counting runs but search_price_per_1k / image bills never apply. + if searchCount > 0 { + result.SearchCount = searchCount + } + if imageCount > 0 { + result.ImageCount = imageCount + result.ImageOutputSizes = imageOutputSizes + } + return result, nil } func isGrokInvalidEncryptedContentResponse(statusCode int, body []byte) bool { @@ -410,7 +441,13 @@ func patchGrokResponsesBodyBase(body []byte, upstreamModel string) ([]byte, erro if !json.Valid(body) { return nil, fmt.Errorf("invalid json request body") } - out, err := sjson.SetBytes(body, "model", upstreamModel) + // sjson may reuse the input backing array; keep the caller's request bytes + // unchanged because the same body can be inspected for billing/retry paths. + out, err := sjson.SetBytes(append([]byte(nil), body...), "model", upstreamModel) + if err != nil { + return nil, err + } + out, err = normalizeGrokResponsesReasoningEffort(out, upstreamModel) if err != nil { return nil, err } @@ -436,6 +473,16 @@ func patchGrokResponsesBodyBase(body []byte, upstreamModel string) ([]byte, erro } } } + if grokModelRejectsLogprobs(upstreamModel) { + for _, unsupportedField := range []string{"logprobs", "top_logprobs"} { + if gjson.GetBytes(out, unsupportedField).Exists() { + out, err = sjson.DeleteBytes(out, unsupportedField) + if err != nil { + return nil, err + } + } + } + } out, err = sanitizeGrokResponsesUnsupportedFields(out) if err != nil { return nil, err @@ -459,6 +506,17 @@ func patchGrokResponsesBodyBase(body []byte, upstreamModel string) ([]byte, erro return out, nil } +// xAI's Grok 4.20 family and newer models do not support OpenAI's logprobs +// fields. Remove them before egress instead of forwarding a request the +// upstream rejects. Older Grok models retain the fields for compatibility. +func grokModelRejectsLogprobs(model string) bool { + model = strings.ToLower(strings.TrimSpace(model)) + if slash := strings.LastIndex(model, "/"); slash >= 0 { + model = strings.TrimSpace(model[slash+1:]) + } + return strings.HasPrefix(model, "grok-4.20") +} + func sanitizeGrokResponsesModelCapabilities(body []byte, upstreamModel string) ([]byte, error) { if !grokModelRejectsReasoningEffort(upstreamModel) { return body, nil @@ -491,6 +549,98 @@ func grokModelRejectsReasoningEffort(model string) bool { } } +func normalizeGrokResponsesReasoningEffort(body []byte, upstreamModel string) ([]byte, error) { + supportsEffort := grokSupportsReasoningEffort(upstreamModel) + out := body + var err error + for _, field := range []string{"reasoning.effort", "reasoning_effort"} { + value := gjson.GetBytes(out, field) + if !value.Exists() { + continue + } + normalized, keep := normalizeGrokReasoningEffortValue(value.String()) + if !supportsEffort || !keep { + out, err = sjson.DeleteBytes(out, field) + } else { + out, err = sjson.SetBytes(out, field, normalized) + } + if err != nil { + return nil, fmt.Errorf("normalize Grok reasoning field %s: %w", field, err) + } + } + if camel := gjson.GetBytes(out, "reasoningEffort"); camel.Exists() { + normalized, keep := normalizeGrokReasoningEffortValue(camel.String()) + out, err = sjson.DeleteBytes(out, "reasoningEffort") + if err != nil { + return nil, fmt.Errorf("remove Grok reasoningEffort: %w", err) + } + if supportsEffort && keep && !gjson.GetBytes(out, "reasoning_effort").Exists() { + out, err = sjson.SetBytes(out, "reasoning_effort", normalized) + if err != nil { + return nil, fmt.Errorf("set Grok reasoning_effort: %w", err) + } + } + } + if reasoning := gjson.GetBytes(out, "reasoning"); reasoning.Exists() && reasoning.IsObject() && len(reasoning.Map()) == 0 { + out, err = sjson.DeleteBytes(out, "reasoning") + if err != nil { + return nil, fmt.Errorf("remove empty Grok reasoning: %w", err) + } + } + return out, nil +} + +func normalizeGrokChatReasoningEffort(body []byte, upstreamModel string) ([]byte, error) { + raw := strings.TrimSpace(gjson.GetBytes(body, "reasoning_effort").String()) + if raw == "" { + raw = strings.TrimSpace(gjson.GetBytes(body, "reasoningEffort").String()) + } + normalized, keep := normalizeGrokReasoningEffortValue(raw) + keep = keep && grokSupportsReasoningEffort(upstreamModel) + out := body + var err error + if gjson.GetBytes(out, "reasoningEffort").Exists() { + out, err = sjson.DeleteBytes(out, "reasoningEffort") + if err != nil { + return nil, err + } + } + if !keep { + if gjson.GetBytes(out, "reasoning_effort").Exists() { + out, err = sjson.DeleteBytes(out, "reasoning_effort") + } + return out, err + } + out, err = sjson.SetBytes(out, "reasoning_effort", normalized) + return out, err +} + +func normalizeGrokReasoningEffortValue(raw string) (string, bool) { + value := strings.NewReplacer("-", "", "_", "", " ", "").Replace(strings.ToLower(strings.TrimSpace(raw))) + switch value { + case "none", "low", "medium", "high": + return value, true + case "minimal": + return "low", true + case "xhigh", "extrahigh", "max", "ultra": + return "high", true + default: + return "", false + } +} + +func grokSupportsReasoningEffort(model string) bool { + model = strings.ToLower(xai.StripGrokProviderPrefix(strings.TrimSpace(model))) + switch model { + case xai.DefaultTextModel, "grok-4.5-latest", "grok-4.3", "grok-4.3-latest", + "grok-3-mini", "grok-3-mini-fast", "grok-4.20-0309-reasoning", + "grok-4.20-reasoning", "grok-4.20-multi-agent-0309": + return true + default: + return false + } +} + var grokResponsesUnsupportedRecursiveFields = map[string]struct{}{ "external_web_access": {}, } @@ -676,15 +826,30 @@ func sanitizeGrokResponsesTools(body []byte) ([]byte, error) { rawTools := tools.Array() filteredTools := make([]json.RawMessage, 0, len(rawTools)) + toolsChanged := false for _, tool := range rawTools { toolType := strings.TrimSpace(tool.Get("type").String()) if _, ok := grokResponsesSupportedToolTypes[toolType]; ok { - filteredTools = append(filteredTools, json.RawMessage(tool.Raw)) + raw := json.RawMessage(tool.Raw) + if toolType == "function" && (!tool.Get("parameters").Exists() || tool.Get("parameters").Type == gjson.Null) { + var payload map[string]any + if err := json.Unmarshal(raw, &payload); err != nil { + return nil, err + } + payload["parameters"] = map[string]any{"type": "object", "properties": map[string]any{}} + encoded, err := json.Marshal(payload) + if err != nil { + return nil, err + } + raw = encoded + toolsChanged = true + } + filteredTools = append(filteredTools, raw) } } var err error - if len(filteredTools) != len(rawTools) { + if len(filteredTools) != len(rawTools) || toolsChanged { if len(filteredTools) == 0 { body, err = sjson.DeleteBytes(body, "tools") } else { @@ -892,7 +1057,7 @@ func (s *OpenAIGatewayService) describeGrokComposerImage( upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) // Image-description probes are auxiliary requests, not conversation turns. // Do not bind them to the caller's Grok prompt-cache identity. - upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token, "", s.cfg) + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, body, token, "", s.cfg, s.settingService) releaseUpstreamCtx() if err != nil { return "", OpenAIUsage{}, fmt.Errorf("build grok composer image bridge request: %w", err) @@ -1062,8 +1227,8 @@ func addOpenAIUsage(dst *OpenAIUsage, usage OpenAIUsage) { dst.ImageOutputTokens += usage.ImageOutputTokens } -func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, cacheIdentity string, cfg *config.Config) (*http.Request, error) { - targetURL, err := buildGrokResponsesURL(account, cfg) +func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Account, body []byte, token, cacheIdentity string, cfg *config.Config, settings ...*SettingService) (*http.Request, error) { + targetURL, err := buildGrokResponsesURL(account, cfg, settings...) if err != nil { return nil, err } @@ -1091,12 +1256,18 @@ func buildGrokResponsesRequest(ctx context.Context, c *gin.Context, account *Acc // applyGrokCLIHeaders identifies subscription traffic as a supported Grok CLI // version. The CLI gateway rejects otherwise valid OAuth requests without it. +// Identity pins come from package xai so service-layer headers match the final +// transport rewrite on cli-chat-proxy.grok.com. func applyGrokCLIHeaders(headers http.Header) { if headers == nil { return } - headers.Set("User-Agent", grokUpstreamUserAgent) - headers.Set("X-Grok-Client-Version", grokCLIVersion) + version := xai.ResolveCLIVersion() + headers.Set("User-Agent", xai.CLIUserAgent(version)) + headers.Set("X-Grok-Client-Version", version) + headers.Set("x-grok-client-version", version) + headers.Set("x-grok-client-identifier", xai.CLIClientIdentifier) + // Historical mode value expected by some unit tests / older CLI probes. headers.Set("X-Grok-Client-Mode", "interactive") } @@ -1119,6 +1290,15 @@ func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, acco } } + updates := map[string]any{ + grokQuotaSnapshotExtraKey: snapshot, + } + // Also derive the scheduling-threshold extras (grok_sched_*) the evaluator + // reads in grokThresholdCandidates. Without this writer the admin-configured + // Grok auto-pause threshold could never fire (the read side was dead config). + for k, v := range buildGrokSchedulerExtraUpdates(snapshot) { + updates[k] = v + } stateCtx := ctx if hasActiveLimit { var cancel context.CancelFunc @@ -1126,9 +1306,7 @@ func (s *OpenAIGatewayService) updateGrokUsageSnapshot(ctx context.Context, acco defer cancel() } if s.accountRepo != nil { - _ = s.accountRepo.UpdateExtra(stateCtx, accountID, map[string]any{ - grokQuotaSnapshotExtraKey: snapshot, - }) + _ = s.accountRepo.UpdateExtra(stateCtx, accountID, updates) } // Error responses are reconciled by handleGrokAccountUpstreamError. Pool-mode // API keys retain the snapshot for observability but leave account health to @@ -1337,7 +1515,8 @@ func (s *OpenAIGatewayService) rateLimitGrok(ctx context.Context, account *Accou if s == nil || account == nil { return } - resetAt = normalizeGrokRateLimitResetAt(account, resetAt, time.Now()) + now := time.Now() + resetAt = normalizeGrokRateLimitResetAt(account, resetAt, now) runtimeUntil := resetAt if account.TempUnschedulableUntil != nil && account.TempUnschedulableUntil.After(runtimeUntil) { @@ -1345,6 +1524,106 @@ func (s *OpenAIGatewayService) rateLimitGrok(ctx context.Context, account *Accou } s.BlockAccountScheduling(account, runtimeUntil, "429") persistGrokRateLimit(ctx, s.accountRepo, account, resetAt) + + // Propagate a short team+model cool so sibling OAuth accounts on the same + // xAI team skip the hot model without waiting for each to hit 429 alone. + // Model is taken from the latest request context when available; empty is a + // no-op inside markGrokTeamModelRateLimit. + if model, _ := ctx.Value(grokTeamRateLimitModelContextKey{}).(string); model != "" { + markGrokTeamModelRateLimit(account, model, resolveGrokTeamRateLimitUntil(resetAt, now)) + } +} + +// buildGrokSchedulerExtraUpdates derives the grok_sched_* scheduling snapshot +// (utilization percent + reset time) consumed by EvaluateAccountSchedulingThreshold. +// Utilization is the most-constrained of the requests/tokens windows. +func buildGrokSchedulerExtraUpdates(snapshot *xai.QuotaSnapshot) map[string]any { + if snapshot == nil { + return nil + } + util, reset, ok := grokSnapshotUtilization(snapshot) + if !ok { + return nil + } + updates := map[string]any{ + "grok_sched_utilization": util, + "grok_sched_usage_updated_at": time.Now().UTC().Format(time.RFC3339), + } + if reset != nil { + // 防御:调度阈值暂停时长由 grok_sched_reset_at 决定。若上游返回脏的 + // reset 头(例如把相对毫秒 "6000" 误当相对秒解析出 ~33h 的未来时刻), + // 不设上限会把耗尽账号长时间锁死。xAI 配额窗口不会超过一天,因此对 + // 未来时刻做 grokMaxSchedulingResetHorizon 钳制;过去/无效值直接不写。 + now := time.Now() + if reset.After(now) { + capped := *reset + if horizon := now.Add(grokMaxSchedulingResetHorizon); capped.After(horizon) { + capped = horizon + } + updates["grok_sched_reset_at"] = capped.UTC().Format(time.RFC3339) + } + } + return updates +} + +// grokSnapshotUtilization returns the highest window utilization (0-100) across +// the requests/tokens quota windows and the reset time of that window. +func grokSnapshotUtilization(snapshot *xai.QuotaSnapshot) (float64, *time.Time, bool) { + if snapshot == nil { + return 0, nil, false + } + best := -1.0 + var bestReset *time.Time + consider := func(window *xai.QuotaWindow) { + if window == nil || window.Limit == nil || *window.Limit <= 0 || window.Remaining == nil { + return + } + remaining := *window.Remaining + if remaining < 0 { + remaining = 0 + } + util := (1 - float64(remaining)/float64(*window.Limit)) * 100 + if util < 0 { + util = 0 + } + if util > 100 { + util = 100 + } + if util > best { + best = util + if window.ResetUnix != nil { + t := time.Unix(*window.ResetUnix, 0).UTC() + bestReset = &t + } else { + bestReset = nil + } + } + } + consider(snapshot.Requests) + consider(snapshot.Tokens) + if best < 0 { + return 0, nil, false + } + return best, bestReset, true +} + +// grokMaxSchedulingResetHorizon bounds how far into the future a Grok +// scheduling-threshold pause (grok_sched_reset_at) may be set, so a malformed +// upstream reset header can't park an over-threshold account for days. xAI quota +// windows do not exceed ~a day. +const grokMaxSchedulingResetHorizon = 25 * time.Hour + +// grokTeamRateLimitModelContextKey carries the upstream model for team cools. +type grokTeamRateLimitModelContextKey struct{} + +// withGrokTeamRateLimitModel attaches the upstream model name for rate-limit +// side effects (team+model cool). Safe when model is empty. +func withGrokTeamRateLimitModel(ctx context.Context, model string) context.Context { + model = strings.TrimSpace(model) + if model == "" || ctx == nil { + return ctx + } + return context.WithValue(ctx, grokTeamRateLimitModelContextKey{}, model) } func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte) { @@ -1356,6 +1635,35 @@ func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Contex } now := time.Now() s.updateGrokUsageSnapshot(ctx, account, parseGrokQuotaSnapshot(headers, statusCode, now)) + + // Body-first free-usage / empty / billing / capacity must run before the + // status switch so non-429 free-usage bodies still cool the account. + // Pool-mode still skips durable mutation unless an explicit temp rule matches. + decision := classifyGrokUpstreamFailure(statusCode, responseBody, "") + if decision.ShouldCooldown && decision.Class != GrokFailureNone && decision.Class != GrokFailureRateLimit { + if account.IsPoolMode() { + // Allow configured temp rules (403) below; skip default body cools. + } else { + // A free-tier exhaustion message describes a rolling usage window. Use + // an upstream absolute reset (or Retry-After) when available; otherwise + // apply only a short probe cooldown. Never start a fabricated 24h window + // at the instant this error was received. + if decision.Class == GrokFailureFreeUsage { + if resetAt, limited := grokRateLimitResetAtForAccount(account, parseGrokQuotaSnapshot(headers, statusCode, now), now); limited && resetAt.After(now) { + if decision.Model != "" && isGrokModelSpecificFreeUsage(strings.ToLower(decision.Reason), decision.Model) { + markGrokModelQuotaBlock(account.ID, decision.Model, resetAt) + return + } + s.rateLimitGrok(ctx, account, resetAt) + return + } + } + if s.applyGrokUpstreamFailureDecision(ctx, account, decision) { + return + } + } + } + if statusCode == http.StatusForbidden && s.applyGrokForbiddenPolicy(ctx, account, responseBody) { return } @@ -1367,17 +1675,44 @@ func (s *OpenAIGatewayService) handleGrokAccountUpstreamError(ctx context.Contex case http.StatusUnauthorized: s.tempUnscheduleGrok(ctx, account, 10*time.Minute, "grok credentials unauthorized") case http.StatusPaymentRequired: + // 402 without a body-classified billing decision: keep the legacy 30m cool. s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok payment required") case http.StatusForbidden: + // Spending-limit already handled by body classifier when phrasing matches. + if isGrokSpendingLimitError(responseBody) { + s.rateLimitGrok(ctx, account, grokSpendingLimitResetAt(account, time.Now())) + return + } s.tempUnscheduleGrok(ctx, account, 30*time.Minute, "grok access or entitlement denied") case http.StatusTooManyRequests: // updateGrokUsageSnapshot installs rate-limit state for non-pool accounts. + // Free-usage 429 was already cooled above via body classification. default: if statusCode >= 500 { s.tempUnscheduleGrok(ctx, account, 2*time.Minute, "grok upstream temporary error") } } - _ = responseBody +} + +// isGrokSpendingLimitError detects xAI billing exhaustion bodies (often 403, sometimes 402). +func isGrokSpendingLimitError(responseBody []byte) bool { + if len(responseBody) == 0 { + return false + } + code := strings.ToLower(strings.TrimSpace(firstNonEmpty( + gjson.GetBytes(responseBody, "code").String(), + gjson.GetBytes(responseBody, "error.code").String(), + ))) + if code == "personal-team-blocked:spending-limit" { + return true + } + message := strings.ToLower(strings.TrimSpace(firstNonEmpty( + gjson.GetBytes(responseBody, "error").String(), + gjson.GetBytes(responseBody, "error.message").String(), + gjson.GetBytes(responseBody, "message").String(), + ))) + return strings.Contains(message, "spending limit") || + strings.Contains(message, "run out of credits") } func (s *OpenAIGatewayService) tempUnscheduleGrok(ctx context.Context, account *Account, cooldown time.Duration, reason string) { diff --git a/backend/internal/service/openai_gateway_grok_405_test.go b/backend/internal/service/openai_gateway_grok_405_test.go new file mode 100644 index 0000000000..10ed8a599e --- /dev/null +++ b/backend/internal/service/openai_gateway_grok_405_test.go @@ -0,0 +1,29 @@ +package service + +import ( + "net/http" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestShouldFailoverUpstreamError_405IsFailoverEligible(t *testing.T) { + svc := &OpenAIGatewayService{} + + assert.True(t, svc.shouldFailoverUpstreamError(http.StatusMethodNotAllowed), + "405 should trigger failover so sticky sessions can escape to healthy accounts") +} + +func TestShouldFailoverUpstreamError_ExistingCodesStillWork(t *testing.T) { + svc := &OpenAIGatewayService{} + + failoverCodes := []int{401, 402, 403, 405, 429, 529, 500, 502, 503, 504} + for _, code := range failoverCodes { + assert.True(t, svc.shouldFailoverUpstreamError(code), "status %d should trigger failover", code) + } + + nonFailoverCodes := []int{200, 201, 400, 404, 408, 422} + for _, code := range nonFailoverCodes { + assert.False(t, svc.shouldFailoverUpstreamError(code), "status %d should NOT trigger failover", code) + } +} diff --git a/backend/internal/service/openai_gateway_grok_cache.go b/backend/internal/service/openai_gateway_grok_cache.go index 4c6923e799..82fa39438e 100644 --- a/backend/internal/service/openai_gateway_grok_cache.go +++ b/backend/internal/service/openai_gateway_grok_cache.go @@ -118,6 +118,12 @@ func explicitGrokCacheSeed(c *gin.Context, body []byte, explicitKey string) stri if seed == "" { seed = strings.TrimSpace(explicitKey) } + // previous_response_id is last-resort: multi-turn Responses without an + // explicit session still share one cache identity (model is already in the + // isolated seed). Message ids are rejected by the seed helper. + if seed == "" && len(body) > 0 { + seed = grokPreviousResponseSessionSeed(body) + } return seed } @@ -298,6 +304,9 @@ func applyGrokFreeToolCacheRoute(body, intentSourceBody []byte, account *Account return appendGrokFreeCacheNativeToolsWithPolicy(body, allowPureClientTools, allowFunctionSearch) } +// isKnownGrokFreeAccount recognizes free-tier Grok accounts, used for +// Free cache routing / media free_tier blocks (broader than soft-gate). +// Soft-gate uses isExplicitGrokFreeOAuthAccount (exact "free" only). func isKnownGrokFreeAccount(account *Account) bool { if account == nil || !account.IsGrokOAuth() { return false @@ -313,14 +322,12 @@ func isKnownGrokFreeAccount(account *Account) bool { paidSignal = true } } + // Usage % or a monthly dollar cap is evidence of a paid plan. if billing.UsagePercent != nil || billing.UsedPercent != nil || (billing.MonthlyLimitCents != nil && *billing.MonthlyLimitCents > 0) { paidSignal = true } - // xAI deliberately reports an empty plan for Free accounts; only paid - // subscriptions receive a SuperGrok plan/monthly limit. A successful - // monthly billing observation with no paid signal is therefore positive - // Free evidence, not an unknown tier. Keep partial probes fail-closed. + // Empty plan + successful monthly observation → inferred free (no paid plan/limit). if strings.TrimSpace(billing.MonthlyUpdatedAt) != "" || (billing.StatusCode >= http.StatusOK && billing.StatusCode < http.StatusMultipleChoices && !billing.Partial && len(billing.FailedWindows) == 0) { @@ -340,6 +347,7 @@ func isKnownGrokFreeAccount(account *Account) bool { inferredFreeSignal = true } } + // Only credentials subscription_tier is authoritative here (not plan_type / extra keys). if tier := strings.TrimSpace(account.GetCredential("subscription_tier")); tier != "" { if isGrokFreeSubscriptionTier(tier) { freeSignal = true @@ -347,9 +355,7 @@ func isKnownGrokFreeAccount(account *Account) bool { paidSignal = true } } - // Explicit paid evidence always wins over an inferred Free signal. This - // protects upgraded/stale accounts whose previous quota snapshot still - // carries the historical 2M Free token limit. + // Explicit paid evidence always wins over an inferred Free signal. return !paidSignal && (freeSignal || inferredFreeSignal) } diff --git a/backend/internal/service/openai_gateway_grok_cache_test.go b/backend/internal/service/openai_gateway_grok_cache_test.go index dc3beaa461..4cc3521c38 100644 --- a/backend/internal/service/openai_gateway_grok_cache_test.go +++ b/backend/internal/service/openai_gateway_grok_cache_test.go @@ -24,6 +24,36 @@ func newGrokCacheTestContext(apiKeyID int64) *gin.Context { return c } +func TestGrokPreviousResponseSessionSeed(t *testing.T) { + require.Equal(t, "grok-prev-resp:resp_abc123", grokPreviousResponseSessionSeed([]byte(`{"previous_response_id":"resp_abc123"}`))) + require.Empty(t, grokPreviousResponseSessionSeed([]byte(`{"previous_response_id":"msg_abc123"}`))) + require.Empty(t, grokPreviousResponseSessionSeed([]byte(`{"previous_response_id":""}`))) + require.Empty(t, grokPreviousResponseSessionSeed([]byte(`{}`))) +} + +func TestResolveGrokCacheIdentityUsesPreviousResponseIDWhenNoOtherSeed(t *testing.T) { + gin.SetMode(gin.TestMode) + c := newGrokCacheTestContext(301) + // No prompt_cache_key / headers / reusable prefix — only previous_response_id. + body := []byte(`{"model":"grok","input":[{"role":"user","content":"follow up"}],"previous_response_id":"resp_chain_001"}`) + got := resolveGrokCacheIdentity(c, body, "", "grok-4.5") + require.NotEmpty(t, got) + + // Same previous_response_id → same identity (model already in isolated seed). + again := resolveGrokCacheIdentity(c, body, "", "grok-4.5") + require.Equal(t, got, again) + + // Different model → different identity (model scope). + otherModel := resolveGrokCacheIdentity(c, body, "", "grok-4.3") + require.NotEqual(t, got, otherModel) + + // prompt_cache_key still wins over previous_response_id. + withCache := []byte(`{"model":"grok","prompt_cache_key":"client-session","previous_response_id":"resp_chain_001","input":[{"role":"user","content":"x"}]}`) + cacheID := resolveGrokCacheIdentity(c, withCache, "", "grok-4.5") + require.NotEmpty(t, cacheID) + require.NotEqual(t, got, cacheID) +} + func TestResolveGrokCacheIdentityStableAcrossAppendOnlyTurns(t *testing.T) { gin.SetMode(gin.TestMode) c := newGrokCacheTestContext(101) @@ -858,6 +888,17 @@ func TestGrokFreeMessagesFunctionToolCacheRouteRequiresKnownFreeTier(t *testing. return a }(), }, + { + name: "paid billing overrides stale free credentials", + account: func() *Account { + a := healthyGrokOAuthGatewayTestAccount(9123, "access-token") + a.Credentials["subscription_tier"] = "free" + a.Extra = map[string]any{ + grokBillingExtraKey: map[string]any{"plan": "SuperGrok", "status_code": http.StatusOK}, + } + return a + }(), + }, { name: "partial billing without monthly evidence remains unknown", account: func() *Account { diff --git a/backend/internal/service/openai_gateway_grok_chat_bridge.go b/backend/internal/service/openai_gateway_grok_chat_bridge.go index 19aa4a8fce..ed910eab43 100644 --- a/backend/internal/service/openai_gateway_grok_chat_bridge.go +++ b/backend/internal/service/openai_gateway_grok_chat_bridge.go @@ -586,7 +586,7 @@ func (s *OpenAIGatewayService) forwardGrokChatCompletionsViaResponses( return nil, fmt.Errorf("get grok access token: %w", err) } upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) - upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, cacheIdentity, s.cfg) + upstreamReq, err := buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, cacheIdentity, s.cfg, s.settingService) releaseUpstreamCtx() if err != nil { return nil, fmt.Errorf("build grok responses bridge request: %w", err) diff --git a/backend/internal/service/openai_gateway_grok_search_billing_test.go b/backend/internal/service/openai_gateway_grok_search_billing_test.go new file mode 100644 index 0000000000..dc998ec5dd --- /dev/null +++ b/backend/internal/service/openai_gateway_grok_search_billing_test.go @@ -0,0 +1,187 @@ +//go:build unit + +package service + +import ( + "bytes" + "context" + "io" + "net/http" + "net/http/httptest" + "sync" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/usagestats" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestForwardGrokResponses_PropagatesSearchCountFromJSON(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"grok","input":"search something","tools":[{"type":"web_search"}],"stream":false}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + + account := healthyGrokOAuthGatewayTestAccount(9901, "access-token") + repo := &mockAccountRepoForPlatform{accountsByID: map[int64]*Account{account.ID: account}} + upstreamBody := `{ + "id":"resp_search_bill", + "object":"response", + "model":"grok-4.5", + "status":"completed", + "output":[ + {"type":"web_search_call","id":"ws1","call_id":"c1","status":"completed"}, + {"type":"x_search_call","id":"xs1","call_id":"c2"}, + {"type":"message","role":"assistant","content":[{"type":"output_text","text":"ok"}]} + ], + "usage":{"input_tokens":10,"output_tokens":5} + }` + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(bytes.NewReader([]byte(upstreamBody))), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", false, time.Now()) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 2, result.SearchCount, "Grok Responses must surface search tool calls for surcharge billing") + require.Equal(t, 10, result.Usage.InputTokens) + require.Equal(t, 5, result.Usage.OutputTokens) +} + +func TestForwardGrokResponses_PropagatesSearchCountFromSSE(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"grok","input":"search","tools":[{"type":"web_search"}],"stream":true}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + + account := healthyGrokOAuthGatewayTestAccount(9902, "access-token") + repo := &mockAccountRepoForPlatform{accountsByID: map[int64]*Account{account.ID: account}} + // item.done + response.completed for same call_id must count once after wire-up. + sse := "data: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"web_search_call\",\"id\":\"ws1\",\"call_id\":\"c1\"}}\n\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_s\",\"status\":\"completed\",\"output\":[{\"type\":\"web_search_call\",\"id\":\"ws1\",\"call_id\":\"c1\"}],\"usage\":{\"input_tokens\":3,\"output_tokens\":1}}}\n\n" + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(bytes.NewReader([]byte(sse))), + }} + svc := &OpenAIGatewayService{ + httpUpstream: upstream, + grokTokenProvider: NewGrokTokenProvider(repo, nil), + accountRepo: repo, + } + + result, err := svc.forwardGrokResponses(context.Background(), c, account, body, "grok", true, time.Now()) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 1, result.SearchCount, "stream SearchCount must be wired and deduped") +} + +func TestGetSchedulableAccount_AppliesGrokFreeSoftGate(t *testing.T) { + // Sticky/non-list path must not return over-gate free OAuth accounts once cache is warm. + // First sticky hit fail-opens and schedules async refresh; subsequent hits use the cache. + cfg := &config.Config{} + cfg.Gateway.Grok.FreeQuotaSoftGateEnabled = true + cfg.Gateway.Grok.FreeQuotaTokenLimit = 500_000 + cfg.Gateway.Grok.FreeQuotaSoftGatePercent = 95 + cfg.Gateway.Grok.FreeQuotaWindowHours = 24 + cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds = 60 + + account := healthyGrokOAuthGatewayTestAccount(8801, "tok") + account.Credentials["subscription_tier"] = "free" + account.Status = StatusActive + account.Schedulable = true + + repo := &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + } + usageRepo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + account.ID: {Tokens: 480_000}, // above 95% of 500k + }} + // Clear shared gateway free-gate cache so this test is deterministic. + gatewayGrokFreeQuotaGateCache.Range(func(key, _ any) bool { + gatewayGrokFreeQuotaGateCache.Delete(key) + return true + }) + if root, ok := freeQuotaRefreshInFlight.Load(&gatewayGrokFreeQuotaGateCache); ok { + if m, ok := root.(*sync.Map); ok { + m.Delete(account.ID) + } + } + svc := &GatewayService{ + cfg: cfg, + accountRepo: repo, + usageLogRepo: usageRepo, + } + + // Miss: fail open + schedule refresh. + got, err := svc.getSchedulableAccount(context.Background(), account.ID) + require.NoError(t, err) + require.NotNil(t, got, "first sticky hit fail-opens while free-gate stats refresh") + + require.Eventually(t, func() bool { + got, err := svc.getSchedulableAccount(context.Background(), account.ID) + return err == nil && got == nil + }, 2*time.Second, 10*time.Millisecond, "over free soft-gate sticky hit must miss after cache warm") +} + +func TestOpenAIGetSchedulableAccount_AppliesGrokFreeSoftGate(t *testing.T) { + // Legacy OpenAI-compatible sticky (advanced scheduler off) must free-gate Grok. + cfg := &config.Config{} + cfg.Gateway.Grok.FreeQuotaSoftGateEnabled = true + cfg.Gateway.Grok.FreeQuotaTokenLimit = 500_000 + cfg.Gateway.Grok.FreeQuotaSoftGatePercent = 95 + cfg.Gateway.Grok.FreeQuotaWindowHours = 24 + cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds = 60 + + account := healthyGrokOAuthGatewayTestAccount(8802, "tok") + account.Credentials["subscription_tier"] = "free" + account.Status = StatusActive + account.Schedulable = true + + repo := &mockAccountRepoForPlatform{ + accountsByID: map[int64]*Account{account.ID: account}, + } + usageRepo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{ + account.ID: {Tokens: 480_000}, + }} + openaiGrokFreeQuotaGateCache.Range(func(key, _ any) bool { + openaiGrokFreeQuotaGateCache.Delete(key) + return true + }) + if root, ok := freeQuotaRefreshInFlight.Load(&openaiGrokFreeQuotaGateCache); ok { + if m, ok := root.(*sync.Map); ok { + m.Delete(account.ID) + } + } + svc := &OpenAIGatewayService{ + cfg: cfg, + accountRepo: repo, + usageLogRepo: usageRepo, + } + + got, err := svc.getSchedulableAccount(context.Background(), account.ID) + require.NoError(t, err) + require.NotNil(t, got, "first sticky hit fail-opens while free-gate stats refresh") + + require.Eventually(t, func() bool { + got, err := svc.getSchedulableAccount(context.Background(), account.ID) + return err == nil && got == nil + }, 2*time.Second, 10*time.Millisecond, "OpenAI legacy sticky must apply free soft-gate after cache warm") +} + +func TestCountGrokNativeSearchCallsFromJSON_MessagesStyleBody(t *testing.T) { + // Proves the same counter used by Anthropic-buffered Grok /v1/messages path. + body := []byte(`{"id":"r1","output":[{"type":"web_search_call","id":"ws1"},{"type":"message","role":"assistant"}]}`) + require.Equal(t, 1, countGrokNativeSearchCallsFromJSONBytes(body)) +} diff --git a/backend/internal/service/openai_gateway_grok_test.go b/backend/internal/service/openai_gateway_grok_test.go index d087bd5184..3e6ce06bc9 100644 --- a/backend/internal/service/openai_gateway_grok_test.go +++ b/backend/internal/service/openai_gateway_grok_test.go @@ -58,7 +58,7 @@ func TestPatchGrokResponsesBodySanitizesComposerReasoningParameters(t *testing.T {name: "grok 4.5", upstreamModel: "grok-4.5", wantReasoning: true}, } - body := []byte(`{ + bodyTemplate := []byte(`{ "model": "grok", "input": "hello", "reasoning": {"effort": "medium", "summary": "auto"}, @@ -68,7 +68,7 @@ func TestPatchGrokResponsesBodySanitizesComposerReasoningParameters(t *testing.T for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - patched, err := patchGrokResponsesBody(body, tt.upstreamModel) + patched, err := patchGrokResponsesBody(append([]byte(nil), bodyTemplate...), tt.upstreamModel) require.NoError(t, err) require.True(t, json.Valid(patched)) require.Equal(t, tt.upstreamModel, gjson.GetBytes(patched, "model").String()) @@ -76,7 +76,7 @@ func TestPatchGrokResponsesBodySanitizesComposerReasoningParameters(t *testing.T if tt.wantReasoning { require.Equal(t, "medium", gjson.GetBytes(patched, "reasoning.effort").String()) require.Equal(t, "medium", gjson.GetBytes(patched, "reasoning_effort").String()) - require.Equal(t, "medium", gjson.GetBytes(patched, "reasoningEffort").String()) + require.False(t, gjson.GetBytes(patched, "reasoningEffort").Exists()) return } @@ -142,6 +142,61 @@ func TestPatchGrokResponsesBodyKeepsPenaltyAndStopFieldsForNon45Models(t *testin require.Len(t, gjson.GetBytes(patched, "stop").Array(), 1) } +func TestPatchGrokResponsesBodyDropsLogprobsForGrok420Family(t *testing.T) { + t.Parallel() + body := []byte(`{"model":"grok-4.20-0309-reasoning","input":"hello","logprobs":true,"top_logprobs":5}`) + patched, err := patchGrokResponsesBody(body, "grok-4.20-0309-reasoning") + require.NoError(t, err) + require.False(t, gjson.GetBytes(patched, "logprobs").Exists()) + require.False(t, gjson.GetBytes(patched, "top_logprobs").Exists()) +} + +func TestPatchGrokResponsesBodyNormalizesReasoningEffortAliases(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + body string + path string + want string + }{ + {name: "minimal nested", body: `{"input":"hi","reasoning":{"effort":"minimal"}}`, path: "reasoning.effort", want: "low"}, + {name: "xhigh snake", body: `{"input":"hi","reasoning_effort":"xhigh"}`, path: "reasoning_effort", want: "high"}, + {name: "max camel", body: `{"input":"hi","reasoningEffort":"max"}`, path: "reasoning_effort", want: "high"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + patched, err := patchGrokResponsesBody([]byte(tt.body), "grok-4.5") + require.NoError(t, err) + require.Equal(t, tt.want, gjson.GetBytes(patched, tt.path).String(), string(patched)) + require.False(t, gjson.GetBytes(patched, "reasoningEffort").Exists()) + }) + } +} + +func TestPatchGrokResponsesBodyAddsDefaultFunctionParameters(t *testing.T) { + patched, err := patchGrokResponsesBody( + []byte(`{"input":"hi","tools":[{"type":"function","name":"lookup"},{"type":"function","name":"wait","parameters":null}]}`), + "grok-4.5", + ) + require.NoError(t, err) + for _, tool := range gjson.GetBytes(patched, "tools").Array() { + require.Equal(t, "object", tool.Get("parameters.type").String(), string(patched)) + require.True(t, tool.Get("parameters.properties").IsObject(), string(patched)) + } +} + +func TestNormalizeGrokChatReasoningEffort(t *testing.T) { + patched, err := normalizeGrokChatReasoningEffort([]byte(`{"reasoningEffort":"ultra"}`), "grok-4.3") + require.NoError(t, err) + require.Equal(t, "high", gjson.GetBytes(patched, "reasoning_effort").String()) + require.False(t, gjson.GetBytes(patched, "reasoningEffort").Exists()) + + patched, err = normalizeGrokChatReasoningEffort([]byte(`{"reasoning_effort":"high"}`), "grok-composer-2.5-fast") + require.NoError(t, err) + require.False(t, gjson.GetBytes(patched, "reasoning_effort").Exists()) +} + func TestPatchGrokResponsesBodyDropsNestedUnsupportedFields(t *testing.T) { t.Parallel() @@ -857,6 +912,33 @@ func TestCanonicalizeGrokMediaImageURLFieldsReplacesEmptyOfficialURL(t *testing. require.False(t, gjson.GetBytes(out, "image.image_url").Exists()) } +func TestPrepareGrokImageEditNormalizesOfficialImageObjects(t *testing.T) { + body := []byte(`{ + "model":"grok-imagine-image-quality", + "image":{"image_url":{"url":"https://example.com/first.png"}}, + "images":["https://example.com/second.png"], + "mask":{"image_url":"https://example.com/mask.png"} + }`) + + out, contentType, err := prepareGrokMediaForwardBody(GrokMediaEndpointImagesEdits, body, "application/json") + require.NoError(t, err) + require.Equal(t, "application/json", contentType) + for _, path := range []string{"image", "images.0", "mask"} { + require.Equal(t, "image_url", gjson.GetBytes(out, path+".type").String()) + require.NotEmpty(t, gjson.GetBytes(out, path+".url").String()) + require.False(t, gjson.GetBytes(out, path+".image_url").Exists()) + } +} + +func TestPrepareGrokImageEditRejectsMoreThanThreeSources(t *testing.T) { + body := []byte(`{"images":["https://example.com/1.png","https://example.com/2.png","https://example.com/3.png","https://example.com/4.png"]}`) + + out, _, err := prepareGrokMediaForwardBody(GrokMediaEndpointImagesEdits, body, "application/json") + require.Error(t, err) + require.Nil(t, out) + require.Contains(t, err.Error(), "maximum of 3 source images") +} + func TestNormalizeGrokMediaModelForEndpoint(t *testing.T) { tests := []struct { name string @@ -870,7 +952,7 @@ func TestNormalizeGrokMediaModelForEndpoint(t *testing.T) { {name: "image quality passthrough", endpoint: GrokMediaEndpointImagesGenerations, model: "grok-imagine-image-quality", want: "grok-imagine-image-quality"}, {name: "image fast passthrough", endpoint: GrokMediaEndpointImagesGenerations, model: "grok-imagine-image", want: "grok-imagine-image"}, {name: "video passthrough", endpoint: GrokMediaEndpointVideosGenerations, model: "grok-imagine-video", want: "grok-imagine-video"}, - {name: "video 1.5 text-only fallback", endpoint: GrokMediaEndpointVideosGenerations, model: "grok-imagine-video-1.5", want: "grok-imagine-video"}, + {name: "video 1.5 text-only remains explicit", endpoint: GrokMediaEndpointVideosGenerations, model: "grok-imagine-video-1.5", want: "grok-imagine-video-1.5"}, {name: "video 1.5 image-to-video passthrough", endpoint: GrokMediaEndpointVideosGenerations, model: "grok-imagine-video-1.5", hasInputImage: true, want: "grok-imagine-video-1.5"}, } @@ -962,9 +1044,9 @@ func TestForwardGrokMediaAppliesAccountModelMappingAfterEndpointNormalization(t path: "/v1/videos/generations", body: `{"model":"grok-imagine-video-1.5","prompt":"waves"}`, modelMapping: map[string]any{"grok-imagine-video": "grok-image-video"}, - wantRequestModel: "grok-imagine-video", - wantUpstream: "grok-image-video", - wantBody: `{"model":"grok-image-video","prompt":"waves"}`, + wantRequestModel: "grok-imagine-video-1.5", + wantUpstream: "grok-imagine-video-1.5", + wantBody: `{"model":"grok-imagine-video-1.5","prompt":"waves"}`, responseBody: `{"request_id":"video-request-mapped"}`, }, { @@ -1204,18 +1286,64 @@ func TestForwardGrokMediaVideoGenerationReturnsUsageAndResponseID(t *testing.T) result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointVideosGenerations, "", body, "application/json") require.NoError(t, err) require.Equal(t, "https://xai.test/v1/videos/generations", upstream.lastReq.URL.String()) - require.JSONEq(t, `{"model":"grok-imagine-video","prompt":"waves","resolution":"720p","duration":10}`, string(upstream.lastBody)) + require.JSONEq(t, `{"model":"grok-imagine-video-1.5","prompt":"waves","resolution":"720p","duration":10}`, string(upstream.lastBody)) require.Equal(t, "video-request-123", result.ResponseID) - require.Equal(t, "grok-imagine-video", result.BillingModel) + require.Equal(t, "grok-imagine-video-1.5", result.BillingModel) require.Equal(t, 3, result.Usage.InputTokens) require.Equal(t, 4, result.Usage.OutputTokens) - require.Equal(t, 1, result.ImageCount) + // Create accepts the job only — VideoCount stays 0 until status returns video.url. + require.Equal(t, 0, result.ImageCount) require.Empty(t, result.ImageSize) - require.Equal(t, 1, result.VideoCount) + require.Equal(t, 0, result.VideoCount) require.Equal(t, VideoBillingResolution720P, result.VideoResolution) require.Equal(t, 10, result.VideoDurationSeconds) } +func TestForwardGrokMediaVideoGenerationReturnsTaskIDAsResponseID(t *testing.T) { + t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") + gin.SetMode(gin.TestMode) + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + body := []byte(`{"model":"grok-imagine-video","prompt":"waves"}`) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/videos/generations", bytes.NewReader(body)) + c.Request.Header.Set("Content-Type", "application/json") + + account := &Account{ + ID: 63, + Name: "grok", + Platform: PlatformGrok, + Type: AccountTypeAPIKey, + Concurrency: 1, + Credentials: map[string]any{ + "api_key": "api-key", + "base_url": "https://xai.test/v1", + }, + } + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"task_id":"video-task-123"}`)), + }} + svc := &OpenAIGatewayService{httpUpstream: upstream} + + result, err := svc.ForwardGrokMedia(context.Background(), c, account, GrokMediaEndpointVideosGenerations, "", body, "application/json") + require.NoError(t, err) + require.Equal(t, "video-task-123", result.ResponseID) +} + +func TestExtractGrokMediaVideoRequestIDPreservesExistingPrecedence(t *testing.T) { + body := []byte(`{ + "request_id":"request-id", + "id":"id", + "task_id":"task-id", + "data":{"request_id":"data-request-id","id":"data-id","task_id":"data-task-id"}, + "video":{"request_id":"video-request-id","id":"video-id","task_id":"video-task-id"} + }`) + + require.Equal(t, "request-id", extractGrokMediaVideoRequestID(body)) +} + func TestForwardGrokMediaVideoGenerationPreservesImageToVideoModel(t *testing.T) { t.Setenv(xai.EnvAllowUnsafeURLOverrides, "true") gin.SetMode(gin.TestMode) @@ -1383,7 +1511,7 @@ func TestForwardGrokMediaVideoMutationEndpoints(t *testing.T) { require.Equal(t, http.MethodPost, upstream.lastReq.Method) require.JSONEq(t, `{"model":"vendor-video-mutation","prompt":"continue","video":{"url":"https://example.com/in.mp4"},"duration":6}`, string(upstream.lastBody)) require.Equal(t, "video-mutation-123", result.ResponseID) - require.Equal(t, 1, result.VideoCount) + require.Equal(t, 0, result.VideoCount) require.Equal(t, 6, result.VideoDurationSeconds) require.Equal(t, "grok-imagine-video", result.BillingModel) require.Equal(t, "vendor-video-mutation", result.UpstreamModel) @@ -1964,7 +2092,7 @@ func TestAccountTestServiceGrokAPIKeyUsesXAIResponses(t *testing.T) { c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/54/test", nil) - err := svc.testGrokAccountConnection(c, account, "grok") + err := svc.testGrokAccountConnection(c, account, "grok", "", AccountTestModeDefault, AccountTestOptions{}) require.NoError(t, err) require.Equal(t, "https://api.x.ai/v1/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer xai-test-key", upstream.lastReq.Header.Get("Authorization")) @@ -1998,7 +2126,7 @@ func TestAccountTestServiceGrokAPIKeyAllowsConfiguredHTTPWhenGlobalPolicyDoes(t c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/55/test", nil) - err := svc.testGrokAccountConnection(c, account, "grok") + err := svc.testGrokAccountConnection(c, account, "grok", "", AccountTestModeDefault, AccountTestOptions{}) require.NoError(t, err) require.Equal(t, "http://grok.example.test/v1/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer third-party-key", upstream.lastReq.Header.Get("Authorization")) @@ -2026,13 +2154,13 @@ func TestAccountTestServiceGrokOAuthPaymentRequiredTemporarilyUnschedulesAccount c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/admin/accounts/56/test", nil) before := time.Now() - err := svc.testGrokAccountConnection(c, account, "grok") + err := svc.testGrokAccountConnection(c, account, "grok", "", AccountTestModeDefault, AccountTestOptions{}) require.Error(t, err) - require.Equal(t, 1, repo.tempUnschedCalls) - require.Equal(t, account.ID, repo.lastTempUnschedID) - require.Equal(t, "grok payment required", repo.lastTempUnschedReason) - require.WithinDuration(t, before.Add(30*time.Minute), repo.lastTempUnschedUntil, time.Second) + require.Zero(t, repo.tempUnschedCalls) + require.Equal(t, 1, repo.rateLimitedCalls) + require.Equal(t, account.ID, repo.lastRateLimitedID) + require.WithinDuration(t, before.Add(grokSpendingLimitProbeCooldown), repo.lastRateLimitResetAt, time.Second) require.Contains(t, recorder.Body.String(), `"type":"error"`) require.Contains(t, recorder.Body.String(), "Grok Responses API returned 402") } @@ -2082,7 +2210,7 @@ func TestForwardAsChatCompletionsForGrokStreamingUsesRawXAIChatCompletions(t *te require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept")) - require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.lastBody, "model").String()) require.True(t, gjson.GetBytes(upstream.lastBody, "stream_options.include_usage").Bool()) require.True(t, result.Stream) @@ -2255,7 +2383,7 @@ func TestForwardAsChatCompletionsForGrokStreamingStopFallsBackToRawXAIChatComple require.Equal(t, xai.DefaultCLIBaseURL+"/chat/completions", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept")) - require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) require.NotEmpty(t, upstream.lastReq.Header.Get(grokConversationIDHeader)) require.NotEqual(t, "native-client-conversation", upstream.lastReq.Header.Get(grokConversationIDHeader)) @@ -2359,7 +2487,7 @@ func TestForwardAsAnthropicForGrokUsesXAIResponses(t *testing.T) { require.NoError(t, err) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) - require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) require.Equal(t, "grok-experimental", upstream.lastReq.Header.Get("OpenAI-Beta")) require.Empty(t, upstream.lastReq.Header.Get("originator")) @@ -2501,6 +2629,20 @@ func grokMessagesSSECompletedResponse(responseID string, cachedTokens int) *http } } +func TestHandleGrokAccountUpstreamErrorSpendingLimitUsesRecoverableProbeCool(t *testing.T) { + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + account := &Account{ID: 2570, Platform: PlatformGrok, Type: AccountTypeOAuth} + before := time.Now() + body := []byte(`{"code":"personal-team-blocked:spending-limit","error":"You have run out of credits"}`) + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusForbidden, nil, body) + + require.Equal(t, 1, repo.rateLimitedCalls) + require.WithinDuration(t, before.Add(grokSpendingLimitProbeCooldown), repo.lastRateLimitResetAt, 2*time.Second) + require.Zero(t, repo.tempUnschedCalls) +} + func TestHandleGrokAccountUpstreamErrorTempUnschedulesNonRateLimitStates(t *testing.T) { tests := []struct { name string @@ -2560,6 +2702,23 @@ func TestHandleGrokAccountUpstreamErrorTempUnschedulesNonRateLimitStates(t *test } } +func TestHandleGrokAccountUpstreamErrorSpendingLimit403RateLimits(t *testing.T) { + account := &Account{ID: 614, Platform: PlatformGrok, Type: AccountTypeOAuth} + repo := &grokQuotaAccountRepo{} + svc := &OpenAIGatewayService{accountRepo: repo} + before := time.Now() + body := []byte(`{"code":"personal-team-blocked:spending-limit","error":"You have run out of credits"}`) + + svc.handleGrokAccountUpstreamError(context.Background(), account, http.StatusForbidden, nil, body) + + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + require.Equal(t, 1, repo.rateLimitedCalls) + require.Equal(t, account.ID, repo.lastRateLimitedID) + require.WithinDuration(t, before.Add(grokSpendingLimitProbeCooldown), repo.lastRateLimitResetAt, 2*time.Second) + require.Zero(t, repo.tempUnschedCalls) + require.True(t, isGrokSpendingLimitError(body)) +} + func TestHandleGrokAccountUpstreamError5xxRespectsPoolMode(t *testing.T) { t.Run("pool mode keeps scheduling state", func(t *testing.T) { account := &Account{ @@ -3190,3 +3349,30 @@ func TestIsGrokImageGenerationModel(t *testing.T) { }) } } + +func TestBuildGrokSchedulerExtraUpdates_FeedsThresholdEvaluator(t *testing.T) { + int64p := func(v int64) *int64 { return &v } + resetUnix := time.Now().Add(90 * time.Minute).Unix() + snapshot := &xai.QuotaSnapshot{ + Requests: &xai.QuotaWindow{Limit: int64p(100), Remaining: int64p(30)}, // 70% used + Tokens: &xai.QuotaWindow{Limit: int64p(1000), Remaining: int64p(50), ResetUnix: &resetUnix}, // 95% used (most constrained) + } + + updates := buildGrokSchedulerExtraUpdates(snapshot) + require.NotNil(t, updates) + require.InDelta(t, 95.0, updates["grok_sched_utilization"], 0.001, "picks the most-constrained window") + require.Contains(t, updates, "grok_sched_reset_at") + + // The written extras must actually drive EvaluateAccountSchedulingThreshold + // (proves the previously-dead read side is now fed). + account := &Account{Platform: PlatformGrok, Extra: updates} + decision := EvaluateAccountSchedulingThreshold(account, map[string]int{PlatformGrok: 90}, time.Now()) + require.True(t, decision.ShouldPause) + require.InDelta(t, 95.0, decision.UsedPercent, 0.001) + require.NotNil(t, decision.Until) +} + +func TestBuildGrokSchedulerExtraUpdates_NilWhenNoQuotaWindows(t *testing.T) { + require.Nil(t, buildGrokSchedulerExtraUpdates(&xai.QuotaSnapshot{})) + require.Nil(t, buildGrokSchedulerExtraUpdates(nil)) +} diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 066fc82cec..2aa226dcd4 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -34,6 +34,8 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( promptCacheKey string, defaultMappedModel string, ) (*OpenAIForwardResult, error) { + beginUpstreamResponseModelObservation(c) + // 入口分流:APIKey 账号 + 上游不支持 Responses API → 走 CC 直转(与 // ForwardAsChatCompletions 对称)。缺少此分流时,/v1/messages 入站请求 // 会被无条件转为 Responses 格式发往上游 /v1/responses,导致只支持 @@ -303,7 +305,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx) var upstreamReq *http.Request if account.Platform == PlatformGrok { - upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, grokCacheIdentity, s.cfg) + upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, responsesBody, token, grokCacheIdentity, s.cfg, s.settingService) } else { upstreamReq, err = s.buildUpstreamRequest(upstreamCtx, c, account, responsesBody, token, isStream, promptCacheKey, false) } @@ -355,7 +357,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( break } upstreamCtxRetry, releaseRetry := detachUpstreamContext(ctx) - upstreamReq, err = buildGrokResponsesRequest(upstreamCtxRetry, c, account, responsesBody, token, grokCacheIdentity, s.cfg) + upstreamReq, err = buildGrokResponsesRequest(upstreamCtxRetry, c, account, responsesBody, token, grokCacheIdentity, s.cfg, s.settingService) releaseRetry() if err != nil { return nil, fmt.Errorf("build grok retry request: %w", err) @@ -549,6 +551,11 @@ func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse( writeAnthropicError(c, http.StatusBadGateway, "api_error", "Upstream stream ended without a terminal response event") return nil, fmt.Errorf("upstream stream ended without terminal event") } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + observer.Observe(finalResponse.Model, true) if strings.TrimSpace(finalResponse.Status) == "failed" { payload, _ := json.Marshal(gin.H{"type": "response.failed", "response": finalResponse}) @@ -601,16 +608,27 @@ func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse( c.Header("Content-Type", "application/json; charset=utf-8") c.JSON(http.StatusOK, anthropicResp) - return &OpenAIForwardResult{ - RequestID: requestID, - ResponseID: finalResponse.ID, - Usage: usage, - Model: originalModel, - BillingModel: billingModel, - UpstreamModel: upstreamModel, - Stream: false, - Duration: time.Since(startTime), - }, nil + result := &OpenAIForwardResult{ + RequestID: requestID, + ResponseID: finalResponse.ID, + Usage: usage, + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: false, + Duration: time.Since(startTime), + } + // Grok /v1/messages uses Responses upstream; count native search for surcharge. + if account != nil && account.IsGrok() && finalResponse != nil { + if body, err := json.Marshal(finalResponse); err == nil { + if n := countGrokNativeSearchCallsFromJSONBytes(body); n > 0 { + result.SearchCount = n + } + } + } + return result, nil } func isOpenAICompatResponsesTerminalEvent(eventType string) bool { @@ -830,6 +848,9 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( clientOutputStarted := false var streamFailoverErr error var streamNonFailoverErr error + searchCount := 0 + streamSearchSeen := make(map[string]struct{}) + countSearch := account != nil && account.IsGrok() scanner := s.newUpstreamSSEScanner(resp.Body) @@ -846,21 +867,31 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( if intervalTicker != nil { intervalCh = intervalTicker.C } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } // resultWithUsage builds the final result snapshot. resultWithUsage := func() *OpenAIForwardResult { - return &OpenAIForwardResult{ - RequestID: requestID, - ResponseID: responseID, - Usage: usage, - Model: originalModel, - BillingModel: billingModel, - UpstreamModel: upstreamModel, - Stream: true, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, - ClientDisconnect: clientDisconnected, + out := &OpenAIForwardResult{ + RequestID: requestID, + ResponseID: responseID, + Usage: usage, + Model: originalModel, + BillingModel: billingModel, + UpstreamModel: upstreamModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + Stream: true, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, + ClientDisconnect: clientDisconnected, } + if searchCount > 0 { + out.SearchCount = searchCount + } + return out } // processDataLine handles a single "data: ..." SSE line from upstream. @@ -870,6 +901,9 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( ms := int(time.Since(startTime).Milliseconds()) firstTokenMs = &ms } + if countSearch { + searchCount += countGrokNativeSearchCallsInSSEDataDedup([]byte(payload), streamSearchSeen) + } var event apicompat.ResponsesStreamEvent if err := json.Unmarshal([]byte(payload), &event); err != nil { @@ -879,6 +913,7 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( ) return false } + observer.ObserveOpenAI([]byte(payload), event.Type) eventType := strings.TrimSpace(event.Type) isBareErrorEvent := eventType == "error" diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index b1a3d189cf..44c9670a3e 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -278,17 +278,19 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( } forwardResult := &OpenAIForwardResult{ - RequestID: resp.Header.Get("x-request-id"), - ResponseID: responseID, - Usage: *usage, - Model: reqModel, - UpstreamModel: upstreamPassthroughModel, - ServiceTier: serviceTier, - ReasoningEffort: reasoningEffort, - Stream: reqStream, - OpenAIWSMode: false, - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, + RequestID: resp.Header.Get("x-request-id"), + ResponseID: responseID, + Usage: *usage, + Model: reqModel, + UpstreamModel: upstreamPassthroughModel, + UpstreamResponseModel: observedUpstreamResponseModel(c), + UpstreamResponseModelConflict: observedUpstreamResponseModelConflict(c), + ServiceTier: serviceTier, + ReasoningEffort: reasoningEffort, + Stream: reqStream, + OpenAIWSMode: false, + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, } if imageCount > 0 { forwardResult.ImageCount = imageCount @@ -390,6 +392,10 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( // OAuth 透传到 ChatGPT internal API 时补齐必要头。 if account.Type == AccountTypeOAuth { + // Current Codex OAuth HTTP no longer negotiates the legacy Responses + // experiment. Passthrough may receive it from an older client, so remove + // only that token while preserving any independent beta negotiation. + stripOpenAILegacyResponsesBeta(req.Header) promptCacheKey := strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) req.Host = "chatgpt.com" if err := resolveAndSetOpenAIChatGPTAccountHeaders(ctx, s.accountRepo, req.Header, account); err != nil { @@ -410,9 +416,6 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( } else if req.Header.Get("accept") == "" { req.Header.Set("accept", "text/event-stream") } - if req.Header.Get("OpenAI-Beta") == "" { - req.Header.Set("OpenAI-Beta", "responses=experimental") - } if req.Header.Get("originator") == "" { req.Header.Set("originator", openai.CodexDefaultOriginator) } @@ -456,10 +459,43 @@ func (s *OpenAIGatewayService) buildUpstreamRequestOpenAIPassthrough( // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op) account.ApplyHeaderOverrides(req.Header) + setOpenAICodexRoutingHintFromBody(req.Header, account, body) + logOpenAIRoutingDiagnosticsFromBody(ctx, account, "http_passthrough", req.Header, body, "not_applicable") return req, nil } +func stripOpenAILegacyResponsesBeta(headers http.Header) { + if headers == nil { + return + } + + preserved := make([]string, 0) + for key, values := range headers { + if !strings.EqualFold(strings.TrimSpace(key), "OpenAI-Beta") { + continue + } + delete(headers, key) + for _, value := range values { + parts := strings.Split(value, ",") + kept := parts[:0] + for _, part := range parts { + part = strings.TrimSpace(part) + if part == "" || strings.EqualFold(part, "responses=experimental") { + continue + } + kept = append(kept, part) + } + if len(kept) > 0 { + preserved = append(preserved, strings.Join(kept, ", ")) + } + } + } + for _, value := range preserved { + headers.Add("OpenAI-Beta", value) + } +} + func shouldFailoverOpenAIPassthroughResponse(account *Account, statusCode int, responseBody []byte) bool { if isOpenAIContextWindowError("", responseBody) { return false @@ -757,8 +793,17 @@ func openAIStreamDataStartsClientOutput(data, eventType string) bool { if trimmed == "" { return false } - if strings.TrimSpace(eventType) == "response.failed" { + switch strings.TrimSpace(eventType) { + case "response.failed": return false + case "error": + // 上游降载/瞬时故障会先推 {"type":"error"} 帧、再以 response.failed 收尾。 + // 可重试类错误帧不能算客户端输出:一旦把它当首输出 flush, + // clientOutputStarted 即被固化,随后的 failed 事件永远进不了 pre-output + // failover 分支,只能把致命错误原样转发给客户端。不可重试类 + // (content_policy / invalid_request 等)维持原样转发,保留上游错误细节。 + payload := []byte(trimmed) + return !openAIStreamFailedEventShouldFailover(payload, extractOpenAISSEErrorMessage(payload)) } return !openAIStreamEventIsPreamble(eventType) } @@ -785,6 +830,41 @@ func isOpenAIUpstreamCapacityShedEvent(payload []byte) bool { } } +// openAICapacityShedRetryableClientCode 是把上游容量降载错误转发给客户端时改写 +// 使用的错误码。Codex CLI 按闭集对错误码分类:server_is_overloaded / slow_down +// 被判为致命错误(客户端提示 "Selected model is at capacity. Please try a +// different model." 并直接终止会话),而 server_error 等致命集之外的错误码会进入 +// 客户端内置的退避重试。 +const openAICapacityShedRetryableClientCode = "server_error" + +// sanitizeOpenAICapacityShedErrorCodeForClient 把即将写给下游客户端的 +// error / response.failed 事件中的容量降载错误码改写为客户端可重试的错误码。 +// 走到转发这一步说明网关侧 failover 已不可用(流中途)或已用尽;保留原始降载码 +// 只会让客户端就地终止会话。错误消息原样保留;监控与账号状态判定都基于改写前 +// 的原始 payload,不受影响。rate_limit 等其他错误码一律不动(客户端依赖 +// rate_limit_exceeded 原码解析重试延时)。 +func sanitizeOpenAICapacityShedErrorCodeForClient(payload []byte) ([]byte, bool) { + if len(payload) == 0 || !gjson.ValidBytes(payload) || !isOpenAIUpstreamCapacityShedEvent(payload) { + return payload, false + } + updated := payload + changed := false + for _, path := range []string{"response.error.code", "error.code"} { + switch strings.ToLower(strings.TrimSpace(gjson.GetBytes(updated, path).String())) { + case "server_is_overloaded", "slow_down": + default: + continue + } + next, err := sjson.SetBytes(updated, path, openAICapacityShedRetryableClientCode) + if err != nil { + return payload, false + } + updated = next + changed = true + } + return updated, changed +} + func openAIStreamFailedEventSemanticStatus(payload []byte, message string) int { if isOpenAIContextWindowError(message, payload) { return http.StatusBadRequest @@ -1051,6 +1131,10 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( originalModel string, mappedModel string, ) (*openaiStreamingResultPassthrough, error) { + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } writeOpenAIPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) // SSE headers @@ -1131,6 +1215,8 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( if data, ok := extractOpenAISSEDataLine(line); ok { dataBytes := []byte(data) trimmedData := strings.TrimSpace(data) + rawEventType := strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String()) + observer.ObserveOpenAI(dataBytes, rawEventType) if needModelReplace && strings.Contains(data, mappedModel) { line = s.replaceModelInSSELine(line, mappedModel, originalModel) if replacedData, replaced := extractOpenAISSEDataLine(line); replaced { @@ -1317,6 +1403,15 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough( if err != nil { return nil, err } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + if bodyHasSSEFraming(body) { + observeOpenAISSEBody(observer, string(body)) + } else { + observer.ObserveOpenAI(body, strings.TrimSpace(gjson.GetBytes(body, "type").String())) + } // Detect SSE responses from upstream and convert to JSON. // Some upstreams (e.g. other sub2api instances) may return SSE even when diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index beb674474e..4ae09575e3 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -2041,11 +2041,12 @@ func TestGrokVideoBillingUsesSeparateVideoRateMultiplier(t *testing.T) { err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ Result: &OpenAIForwardResult{ - RequestID: "video-request-123", - ResponseID: "video-request-123", - Model: "grok-imagine-video-1.5", - BillingModel: "grok-imagine-video-1.5", - ImageCount: 1, + RequestID: "video-request-123", + ResponseID: "video-request-123", + Model: "grok-imagine-video-1.5", + BillingModel: "grok-imagine-video-1.5", + // Pure video completion clears ImageCount (handler contract). + ImageCount: 0, VideoCount: 1, VideoResolution: VideoBillingResolution480P, VideoDurationSeconds: 1, @@ -2073,7 +2074,7 @@ func TestGrokVideoBillingUsesSeparateVideoRateMultiplier(t *testing.T) { require.NoError(t, err) require.NotNil(t, usageRepo.lastLog) require.Equal(t, "grok-imagine-video-1.5", usageRepo.lastLog.Model) - require.Equal(t, 1, usageRepo.lastLog.ImageCount) + require.Equal(t, 0, usageRepo.lastLog.ImageCount) require.Nil(t, usageRepo.lastLog.ImageSize) require.InDelta(t, 0.08, usageRepo.lastLog.TotalCost, 1e-12) require.InDelta(t, 0.02, usageRepo.lastLog.ActualCost, 1e-12) @@ -2098,7 +2099,7 @@ func TestOpenAIGatewayServiceRecordUsage_GrokVideoUsesDefaultRateCard(t *testing ResponseID: "video-default-rate-card", Model: "grok-imagine-video-1.5", BillingModel: "grok-imagine-video-1.5", - ImageCount: 1, + ImageCount: 0, VideoCount: 1, VideoResolution: VideoBillingResolution720P, Duration: time.Second, @@ -2122,7 +2123,7 @@ func TestOpenAIGatewayServiceRecordUsage_GrokVideoUsesDefaultRateCard(t *testing // 结果未携带 duration 时按上游默认 8 秒计费:0.14 USD/s × 8s。 require.InDelta(t, 0.14*8, usageRepo.lastLog.TotalCost, 1e-12) require.InDelta(t, 0.14*8, usageRepo.lastLog.ActualCost, 1e-12) - require.Equal(t, 1, usageRepo.lastLog.ImageCount) + require.Equal(t, 0, usageRepo.lastLog.ImageCount) require.NotNil(t, usageRepo.lastLog.BillingMode) require.Equal(t, string(BillingModeVideo), *usageRepo.lastLog.BillingMode) require.Equal(t, 1, usageRepo.lastLog.VideoCount) @@ -2186,7 +2187,7 @@ func TestOpenAIGatewayServiceRecordUsage_GroupVideoPriceOverridesChannelImagePri RequestID: "resp_grok_video_group_price", Model: "grok-imagine-video", BillingModel: "grok-imagine-video", - ImageCount: 1, + ImageCount: 0, VideoCount: 1, VideoResolution: VideoBillingResolution720P, VideoDurationSeconds: 1, @@ -2210,7 +2211,7 @@ func TestOpenAIGatewayServiceRecordUsage_GroupVideoPriceOverridesChannelImagePri require.NoError(t, err) require.NotNil(t, usageRepo.lastLog) - require.Equal(t, 1, usageRepo.lastLog.ImageCount) + require.Equal(t, 0, usageRepo.lastLog.ImageCount) require.Nil(t, usageRepo.lastLog.ImageSize) require.InDelta(t, 0.037, usageRepo.lastLog.TotalCost, 1e-12) require.InDelta(t, 0.037, usageRepo.lastLog.ActualCost, 1e-12) @@ -2218,6 +2219,53 @@ func TestOpenAIGatewayServiceRecordUsage_GroupVideoPriceOverridesChannelImagePri require.Equal(t, string(BillingModeVideo), *usageRepo.lastLog.BillingMode) } +func TestOpenAIGatewayServiceRecordUsage_GroupVideoModelPriceOverridesFlatAndChannelPrice(t *testing.T) { + groupID := int64(129) + channelPrice := 0.201 + flatVideoPrice720P := 0.037 + modelVideoPrice720P := 0.123 + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil) + svc.resolver = newOpenAIImageChannelPricingResolverForTest(t, groupID, "grok-imagine-video-1.5-preview", channelPrice) + + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_grok_video_model_price", + Model: "grok-imagine-video-1.5-preview", + BillingModel: "grok-imagine-video-1.5-preview", + ImageCount: 0, + VideoCount: 1, + VideoResolution: VideoBillingResolution720P, + VideoDurationSeconds: 2, + Duration: time.Second, + }, + APIKey: &APIKey{ + ID: 10129, + GroupID: i64p(groupID), + Group: &Group{ + ID: groupID, + Platform: PlatformGrok, + RateMultiplier: 1, + VideoRateIndependent: true, + VideoRateMultiplier: 1, + VideoPrice720P: &flatVideoPrice720P, + VideoModelPrices: map[string]map[string]float64{ + VideoPriceFamilyGrokImagineVideo15: {VideoBillingResolution720P: modelVideoPrice720P}, + }, + }, + }, + User: &User{ID: 20129}, + Account: &Account{ID: 30129, Platform: PlatformGrok}, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + require.InDelta(t, modelVideoPrice720P*2, usageRepo.lastLog.TotalCost, 1e-12) + require.InDelta(t, modelVideoPrice720P*2, usageRepo.lastLog.ActualCost, 1e-12) + require.NotNil(t, usageRepo.lastLog.BillingMode) + require.Equal(t, string(BillingModeVideo), *usageRepo.lastLog.BillingMode) +} + func TestOpenAIGatewayServiceRecordUsage_HydratesGroupImagePriceWhenAuthSnapshotOmitsIt(t *testing.T) { groupID := int64(130) groupImagePrice2K := 0.021 @@ -2288,7 +2336,7 @@ func TestOpenAIGatewayServiceRecordUsage_HydratesGroupVideoPriceWhenAuthSnapshot RequestID: "resp_grok_video_hydrated_price", Model: "grok-imagine-video", BillingModel: "grok-imagine-video", - ImageCount: 1, + ImageCount: 0, VideoCount: 1, VideoResolution: VideoBillingResolution720P, VideoDurationSeconds: 1, @@ -2329,7 +2377,7 @@ func TestOpenAIGatewayServiceRecordUsage_GrokVideoWithTokenChannelPricingKeepsVi RequestID: "resp_grok_video_token_channel", Model: "grok-imagine-video", BillingModel: "grok-imagine-video", - ImageCount: 1, + ImageCount: 0, VideoCount: 1, VideoResolution: VideoBillingResolution720P, VideoDurationSeconds: 5, @@ -2354,7 +2402,7 @@ func TestOpenAIGatewayServiceRecordUsage_GrokVideoWithTokenChannelPricingKeepsVi require.NotNil(t, usageRepo.lastLog.BillingMode) require.Equal(t, string(BillingModeToken), *usageRepo.lastLog.BillingMode) require.Nil(t, usageRepo.lastLog.ImageSize) - require.Equal(t, 1, usageRepo.lastLog.ImageCount) + require.Equal(t, 0, usageRepo.lastLog.ImageCount) require.Equal(t, 1, usageRepo.lastLog.VideoCount) require.NotNil(t, usageRepo.lastLog.VideoResolution) require.Equal(t, VideoBillingResolution720P, *usageRepo.lastLog.VideoResolution) diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index 2bb586ce52..04b3e7e050 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -308,6 +308,7 @@ func normalizeOpenAICompactRequestBody(body []byte) ([]byte, bool, error) { "tools", "parallel_tool_calls", "reasoning", + "service_tier", "text", "previous_response_id", } { diff --git a/backend/internal/service/openai_gateway_response_flush_test.go b/backend/internal/service/openai_gateway_response_flush_test.go index 2c35374596..ebfa0f22f1 100644 --- a/backend/internal/service/openai_gateway_response_flush_test.go +++ b/backend/internal/service/openai_gateway_response_flush_test.go @@ -387,13 +387,30 @@ func TestOpenAIResponseFlush_FailedAndErrorEventsFlushAtBoundaries(t *testing.T) require.Contains(t, flushes[1], "response.failed") }) - t.Run("error event", func(t *testing.T) { + t.Run("retryable error event buffered until terminal", func(t *testing.T) { + // 可重试类 error 帧不算客户端输出:保持在 attempt 缓冲中不单独 flush, + // 为随后可能到达的 response.failed 保留 pre-output failover 能力, + // 与终止帧一起出站。 body := "data: {\"type\":\"error\",\"error\":{\"message\":\"failed\"}}\n\n" + "data: [DONE]\n\n" recorder := newOpenAIResponseFlushRecorder() result, err := runOpenAIResponseFlushTest(recorder, io.NopCloser(strings.NewReader(body)), config.GatewayConfig{}) + require.NoError(t, err) + require.NotNil(t, result) + gotBody, flushes := recorder.snapshot() + require.Equal(t, body, gotBody) + require.Len(t, flushes, 1) + }) + + t.Run("non-retryable error event flushes at boundary", func(t *testing.T) { + body := "data: {\"type\":\"error\",\"error\":{\"code\":\"invalid_request\",\"message\":\"bad request\"}}\n\n" + + "data: [DONE]\n\n" + recorder := newOpenAIResponseFlushRecorder() + + result, err := runOpenAIResponseFlushTest(recorder, io.NopCloser(strings.NewReader(body)), config.GatewayConfig{}) + require.NoError(t, err) require.NotNil(t, result) gotBody, flushes := recorder.snapshot() diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index 2a1993c463..cda23fa07b 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -28,6 +28,7 @@ type openaiStreamingResult struct { responseID string imageCount int imageOutputSizes []string + searchCount int } type openaiNonStreamingResult struct { @@ -36,6 +37,7 @@ type openaiNonStreamingResult struct { responseID string imageCount int imageOutputSizes []string + searchCount int } func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, startTime time.Time, originalModel, mappedModel string) (*openaiStreamingResult, error) { @@ -43,6 +45,10 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp } func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, startTime time.Time, originalModel, mappedModel, reasoningEffort string) (*openaiStreamingResult, error) { + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } firstOutputTimeout := time.Duration(0) if account != nil && account.Platform == PlatformOpenAI { firstOutputTimeout = s.openAIFirstOutputTimeout(reasoningEffort) @@ -150,6 +156,16 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 { streamInterval = time.Duration(s.cfg.Gateway.StreamDataIntervalTimeout) * time.Second } + // Grok: always enforce an upstream-read idle so hung SSE bodies fail over + // instead of holding the OAuth slot until the client cancels. Prefer the + // global gateway setting when set; otherwise apply a Grok-only default. + if account != nil && account.Platform == PlatformGrok { + cfgSec := 0 + if s.cfg != nil { + cfgSec = s.cfg.Gateway.StreamDataIntervalTimeout + } + streamInterval = resolveGrokStreamIdleTimeout(cfgSec) + } // 仅监控上游数据间隔超时,不被下游写入阻塞影响 var intervalTicker *time.Ticker if streamInterval > 0 { @@ -288,6 +304,10 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. streamOutputAccumulator := apicompat.NewBufferedResponseAccumulator() streamImageOutputs := make([]json.RawMessage, 0, 1) streamSeenImages := make(map[string]struct{}) + searchCounter := 0 + // Dedup search tool calls across SSE events (item.done + response.completed + // both list the same call_id — counting both would ~2× the surcharge). + streamSearchSeen := make(map[string]struct{}) resultWithUsage := func() *openaiStreamingResult { return &openaiStreamingResult{ usage: usage, @@ -295,6 +315,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. responseID: responseID, imageCount: imageCounter.Count(), imageOutputSizes: imageCounter.Sizes(), + searchCount: searchCounter, } } flushPending := func(disconnectMessage string) { @@ -406,6 +427,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. dataBytes := []byte(data) eventTypeRaw := gjson.GetBytes(dataBytes, "type").String() eventType := strings.TrimSpace(eventTypeRaw) + observer.ObserveOpenAI(dataBytes, eventTypeRaw) // 初始上游 data 的 type 只解析一次:原始值保持终止事件的精确匹配,规范化值供后续分支复用。 if openAIStreamEventIsTerminalWithType(data, eventTypeRaw) { sawTerminalEvent = true @@ -462,6 +484,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. line = "data: " + data } imageCounter.AddSSEData(dataBytes) + searchCounter += countGrokNativeSearchCallsInSSEDataDedup(dataBytes, streamSearchSeen) // Correct Codex tool calls if needed (apply_patch -> edit, etc.) if correctedData, corrected := s.toolCorrector.CorrectToolCallsInSSEBytes(dataBytes); corrected { @@ -687,6 +710,16 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. if s.rateLimitService != nil { s.rateLimitService.HandleStreamTimeout(ctx, account, originalModel) } + // Grok: short cool + account failover when no client-visible bytes + // were committed yet (pre-commit). After output started we keep the + // legacy stream_timeout path so partial SSE is not dual-written. + if account != nil && account.Platform == PlatformGrok { + s.tempUnscheduleGrok(ctx, account, grokStreamIdleCooldown, "grok stream idle timeout") + if !openAIStreamClientOutputStarted(c, clientOutputStarted) && !eventShouldFlush { + _ = resp.Body.Close() + return resultWithUsage(), grokStreamIdleFailoverError(account, streamInterval) + } + } sendErrorEvent("stream_timeout") return resultWithUsage(), fmt.Errorf("stream data interval timeout") @@ -1110,6 +1143,15 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r if err != nil { return nil, err } + observer := upstreamResponseModelObserverFromContext(c) + if observer == nil { + observer = beginUpstreamResponseModelObservation(c) + } + if bodyHasSSEFraming(body) { + observeOpenAISSEBody(observer, string(body)) + } else { + observer.ObserveOpenAI(body, strings.TrimSpace(gjson.GetBytes(body, "type").String())) + } // Detect SSE responses for ALL account types via Content-Type header. // Some OpenAI-compatible upstreams (including other sub2api instances) @@ -1180,6 +1222,7 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r responseID: extractOpenAIResponseIDFromJSONBytes(body), imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body), imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body), + searchCount: countGrokNativeSearchCallsFromJSONBytes(body), }, nil } @@ -1274,6 +1317,7 @@ func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Conte responseID: extractOpenAIResponseIDFromJSONBytes(body), imageCount: countOpenAIImageOutputsFromSSEBody(bodyText), imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText), + searchCount: countGrokNativeSearchCallsFromSSEBody(bodyText), }, nil } @@ -1310,10 +1354,21 @@ func extractOpenAISSEErrorMessage(payload []byte) string { } func sanitizeOpenAIResponseFailedEventForClient(payload []byte, eventType string, clientOutputStarted bool) ([]byte, bool) { - if eventType != "response.failed" || len(payload) == 0 || !gjson.ValidBytes(payload) { + eventType = strings.TrimSpace(eventType) + isFailedEvent := eventType == "response.failed" + if (!isFailedEvent && eventType != "error") || len(payload) == 0 || !gjson.ValidBytes(payload) { return payload, false } updated := payload + // 容量降载码对 Codex CLI 是致命错误;事件既然要写给客户端(failover 已不可用), + // 就改写为客户端可重试的错误码。error 帧与 response.failed 都要改:上游降载 + // 总是先推 error 帧再收 failed,两帧携带同一个错误。 + if rewritten, changed := sanitizeOpenAICapacityShedErrorCodeForClient(updated); changed { + updated = rewritten + } + if !isFailedEvent { + return updated, !bytes.Equal(updated, payload) + } if clientOutputStarted && isOpenAIContextWindowError(extractOpenAISSEErrorMessage(payload), payload) { errorPath := "" switch { diff --git a/backend/internal/service/openai_gateway_scheduling.go b/backend/internal/service/openai_gateway_scheduling.go index 23c04daf0a..4e76069474 100644 --- a/backend/internal/service/openai_gateway_scheduling.go +++ b/backend/internal/service/openai_gateway_scheduling.go @@ -75,6 +75,12 @@ func explicitOpenAISessionID(c *gin.Context, body []byte) string { // with Grok's native conversation header only for requests authenticated to a // Grok group. This keeps an unrelated x-grok-conv-id header from changing // scheduling or upstream session behavior for non-Grok groups. +// +// For Grok groups only, previous_response_id is a last-resort sticky seed so +// multi-turn Responses chains stay on the same OAuth account when no explicit +// session/conversation/prompt_cache_key is present. Non-Grok groups omit this +// so HTTP OpenAI paths that delete previous_response_id before upstream are +// unchanged. func explicitOpenAIRequestSessionID(c *gin.Context, body []byte) string { if c == nil { return "" @@ -87,9 +93,27 @@ func explicitOpenAIRequestSessionID(c *gin.Context, body []byte) string { if sessionID == "" && len(body) > 0 { sessionID = strings.TrimSpace(gjson.GetBytes(body, "prompt_cache_key").String()) } + if sessionID == "" && isGrokRequestContext(c) && len(body) > 0 { + sessionID = grokPreviousResponseSessionSeed(body) + } return sessionID } +// grokPreviousResponseSessionSeed returns a stable sticky seed from a Responses +// previous_response_id. Only resp_* response ids are accepted; message ids and +// unknown shapes must not pin sticky routing or prompt-cache identity. +func grokPreviousResponseSessionSeed(body []byte) string { + id := strings.TrimSpace(gjson.GetBytes(body, "previous_response_id").String()) + if id == "" { + return "" + } + if ClassifyOpenAIPreviousResponseIDKind(id) != OpenAIPreviousResponseIDKindResponseID { + return "" + } + // Namespace so content-derived seeds never collide with response ids. + return "grok-prev-resp:" + id +} + // GenerateExplicitSessionHash generates a sticky-session hash only from explicit // client session signals. It intentionally skips content-derived fallback and is // used by stateless endpoints such as /v1/images. @@ -114,6 +138,13 @@ func (s *OpenAIGatewayService) GenerateExplicitSessionHash(c *gin.Context, body // 5. Header: x-grok-conv-id (Grok groups only) // 6. Body: prompt_cache_key // 7. Body: content-based fallback (model + system + tools + first user message) +// +// Grok sticky affinity is intentionally separate from the upstream +// prompt_cache_key identity (resolveGrokCacheIdentity): sticky pins an OAuth +// account for multi-turn routing, while the cache identity is tenant+model +// isolated for xAI server-side prompt cache. For Grok groups we scope the +// sticky seed with the client-requested model so switching models does not +// inherit a stale account binding (grok2api affinityKey pattern). func (s *OpenAIGatewayService) GenerateSessionHash(c *gin.Context, body []byte) string { if c == nil { return "" @@ -127,11 +158,32 @@ func (s *OpenAIGatewayService) GenerateSessionHash(c *gin.Context, body []byte) return "" } + if isGrokRequestContext(c) { + sessionID = grokStickyAffinitySeed(sessionID, body) + } + currentHash, legacyHash := deriveOpenAISessionHashes(sessionID) attachOpenAILegacySessionHashToGin(c, legacyHash) return currentHash } +// grokStickyAffinitySeed scopes sticky routing by model without changing the +// upstream prompt_cache_key written by applyGrokResponsesCacheIdentity. +func grokStickyAffinitySeed(sessionID string, body []byte) string { + sessionID = strings.TrimSpace(sessionID) + if sessionID == "" { + return "" + } + model := "" + if len(body) > 0 { + model = strings.ToLower(strings.TrimSpace(gjson.GetBytes(body, "model").String())) + } + if model == "" { + return "grok-affinity:v1:" + sessionID + } + return "grok-affinity:v1:" + model + ":" + sessionID +} + // GenerateSessionHashWithFallback 先按常规信号生成会话哈希; // 当未携带 session_id/conversation_id/prompt_cache_key 时,使用 fallbackSeed 生成稳定哈希。 // 该方法用于 WS ingress,避免会话信号缺失时发生跨账号漂移。 @@ -1206,7 +1258,14 @@ func (s *OpenAIGatewayService) listSchedulableAccounts(ctx context.Context, grou platform = normalizeOpenAICompatiblePlatform(platform) if s.schedulerSnapshot != nil { accounts, _, err := s.schedulerSnapshot.ListSchedulableAccounts(ctx, groupID, platform, false) - return accounts, err + if err != nil { + return accounts, err + } + accounts = s.filterOpenAIAccountsBySchedulingThreshold(ctx, accounts) + if platform == PlatformGrok { + accounts = s.filterGrokFreeQuotaAccountsForOpenAI(ctx, accounts) + } + return accounts, nil } var accounts []Account var err error @@ -1220,6 +1279,10 @@ func (s *OpenAIGatewayService) listSchedulableAccounts(ctx context.Context, grou if err != nil { return nil, fmt.Errorf("query accounts failed: %w", err) } + accounts = s.filterOpenAIAccountsBySchedulingThreshold(ctx, accounts) + if platform == PlatformGrok { + accounts = s.filterGrokFreeQuotaAccountsForOpenAI(ctx, accounts) + } return accounts, nil } @@ -1254,6 +1317,9 @@ func (s *OpenAIGatewayService) resolveFreshSchedulableOpenAIAccount(ctx context. if s.isOpenAIAccountRequestRuntimeBlocked(fresh, requestedModel) { return nil } + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, fresh) { + return nil + } if s.isOpenAIProxyStreamQuarantined(ctx, fresh) { return nil } @@ -1283,6 +1349,9 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co if !isOpenAICompatibleAccountEligibleForRequest(ctx, account, platform, requestedModel, requireCompact, requiredCapability) { return nil } + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, account) { + return nil + } if !parentHealthyForShadow(account, s.parentAccountLookup(ctx)) { return nil } @@ -1308,6 +1377,9 @@ func (s *OpenAIGatewayService) recheckSelectedOpenAIAccountFromDB(ctx context.Co if s.isOpenAIAccountRequestRuntimeBlocked(latest, requestedModel) { return nil } + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, latest) { + return nil + } if s.isOpenAIProxyStreamQuarantined(ctx, latest) { return nil } @@ -1334,9 +1406,49 @@ func (s *OpenAIGatewayService) getSchedulableAccount(ctx context.Context, accoun if err != nil || account == nil { return account, err } + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, account) { + return nil, nil + } + // Legacy sticky (advanced scheduler off) must still free-gate Grok OAuth. + if account.IsGrok() { + if gated := s.filterGrokFreeQuotaAccountsForOpenAI(ctx, []Account{*account}); len(gated) == 0 { + return nil, nil + } + } return account, nil } +// filterGrokFreeQuotaAccountsForOpenAI applies the same local free soft-gate as +// GatewayService / advanced scheduler, for OpenAI-compatible legacy selection. +func (s *OpenAIGatewayService) filterGrokFreeQuotaAccountsForOpenAI(ctx context.Context, accounts []Account) []Account { + if s == nil { + return accounts + } + return filterGrokFreeQuotaAccountsCore(ctx, s.cfg, s.usageLogRepo, &openaiGrokFreeQuotaGateCache, accounts) +} + +func (s *OpenAIGatewayService) filterOpenAIAccountsBySchedulingThreshold(ctx context.Context, accounts []Account) []Account { + if len(accounts) == 0 { + return accounts + } + + filtered := make([]Account, 0, len(accounts)) + for i := range accounts { + if s.isOpenAIAccountBlockedBySchedulingThreshold(ctx, &accounts[i]) { + continue + } + filtered = append(filtered, accounts[i]) + } + return filtered +} + +func (s *OpenAIGatewayService) isOpenAIAccountBlockedBySchedulingThreshold(ctx context.Context, account *Account) bool { + if s == nil || s.rateLimitService == nil || account == nil { + return false + } + return s.rateLimitService.ApplyAccountSchedulingThreshold(ctx, account) +} + func (s *OpenAIGatewayService) hydrateSelectedAccount(ctx context.Context, account *Account) (*Account, error) { if account == nil || s.schedulerSnapshot == nil { return account, nil diff --git a/backend/internal/service/openai_gateway_search_surcharge_test.go b/backend/internal/service/openai_gateway_search_surcharge_test.go new file mode 100644 index 0000000000..fd3eefe4b4 --- /dev/null +++ b/backend/internal/service/openai_gateway_search_surcharge_test.go @@ -0,0 +1,118 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCalculateOpenAIRecordUsageCost_SearchIsAdditiveToTokens(t *testing.T) { + t.Parallel() + + price := 10.0 // $10 / 1k searches → 100 searches = $1.0 + svc := &OpenAIGatewayService{ + billingService: newTestBillingService(), + } + apiKey := &APIKey{ + Group: &Group{ + SearchPricePer1k: &price, + }, + } + + // claude-sonnet-4 fallback: Input $3/MTok, Output $15/MTok + // 1000 in + 500 out → 0.003 + 0.0075 = 0.0105 + // + 100 searches → +1.0 → total 1.0105 + cost, err := svc.calculateOpenAIRecordUsageCost( + context.Background(), + &OpenAIForwardResult{SearchCount: 100}, + apiKey, + []string{"claude-sonnet-4"}, + 1.0, + 1.0, + 1.0, + 1.0, + UsageTokens{InputTokens: 1000, OutputTokens: 500}, + "", + false, + ) + require.NoError(t, err) + require.NotNil(t, cost) + require.InDelta(t, 1.0105, cost.ActualCost, 1e-9) + require.InDelta(t, 1.0105, cost.TotalCost, 1e-9) +} + +func TestCalculateOpenAIRecordUsageCost_SearchOnlyWhenNoTokenPricing(t *testing.T) { + t.Parallel() + + price := 10.0 + svc := &OpenAIGatewayService{ + billingService: newTestBillingService(), + } + apiKey := &APIKey{ + Group: &Group{SearchPricePer1k: &price}, + } + // Empty model list: token path fails; search-only surcharge still bills. + cost, err := svc.calculateOpenAIRecordUsageCost( + context.Background(), + &OpenAIForwardResult{SearchCount: 100}, + apiKey, + nil, + 1.0, + 1.0, + 1.0, + 1.0, + UsageTokens{}, + "", + false, + ) + require.NoError(t, err) + require.NotNil(t, cost) + require.InDelta(t, 1.0, cost.ActualCost, 1e-9) +} + +func TestGroupMediaPricingLooksIncomplete_VideoModelPricesComplete(t *testing.T) { + t.Parallel() + require.True(t, groupMediaPricingLooksIncomplete(nil)) + require.True(t, groupMediaPricingLooksIncomplete(&Group{})) + require.False(t, groupMediaPricingLooksIncomplete(&Group{ + VideoModelPrices: map[string]map[string]float64{ + "grok-imagine-video": {"720p": 0.1}, + }, + })) + price := 10.0 + require.False(t, groupMediaPricingLooksIncomplete(&Group{SearchPricePer1k: &price})) + require.False(t, groupMediaPricingLooksIncomplete(&Group{AudioRealtimePricePerMin: &price})) + // Legacy video price alone still marks complete (existing path). + require.False(t, groupMediaPricingLooksIncomplete(&Group{VideoPrice720P: &price})) +} + +func TestCalculateOpenAIRecordUsageCost_TokenPricingErrorNotSwallowedBySearch(t *testing.T) { + t.Parallel() + + price := 10.0 + svc := &OpenAIGatewayService{ + billingService: newTestBillingService(), + } + apiKey := &APIKey{ + Group: &Group{SearchPricePer1k: &price}, + } + // Unknown model → token pricing fails; search must not replace that with $0/$search bill. + cost, err := svc.calculateOpenAIRecordUsageCost( + context.Background(), + &OpenAIForwardResult{SearchCount: 100}, + apiKey, + []string{"totally-unknown-model-xyz-no-pricing"}, + 1.0, + 1.0, + 1.0, + 1.0, + UsageTokens{InputTokens: 1000, OutputTokens: 500}, + "", + false, + ) + require.Error(t, err) + require.Nil(t, cost) +} diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 135f63aa81..55353e1d0a 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -243,6 +243,10 @@ type OpenAIForwardResult struct { // UpstreamModel is the actual model sent to the upstream provider after mapping. // Empty when no mapping was applied (requested model was used as-is). UpstreamModel string + // UpstreamResponseModel is captured from the raw successful upstream + // response before any client-facing rewrite or protocol conversion. + UpstreamResponseModel string + UpstreamResponseModelConflict bool // UpstreamEndpoint is the actual upstream API path used for this request. // It avoids guessing when one downstream protocol can use multiple upstream endpoints. UpstreamEndpoint string @@ -275,6 +279,10 @@ type OpenAIForwardResult struct { // WebSearchCalls 是 Codex alpha/search 网页搜索调用次数(每次成功请求为 1)。 // 上游不返回 usage 字段,>0 时走按次计费(分组单价 × 次数 × 倍率)。 WebSearchCalls int + // SearchCount is Grok-native web_search / tool search call count (per 1k pricing). + SearchCount int + // AudioUsage carries Voice billing units when present. + AudioUsage *AudioUsage wsReplayInput []json.RawMessage wsReplayInputExists bool diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 3ac83b4d0d..b3cb5fbd97 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -68,6 +68,29 @@ func (r stubOpenAIAccountRepo) GetByID(ctx context.Context, id int64) (*Account, return nil, errors.New("account not found") } +func (r stubOpenAIAccountRepo) GetByIDs(ctx context.Context, ids []int64) ([]*Account, error) { + if len(ids) == 0 { + return []*Account{}, nil + } + index := make(map[int64]*Account, len(r.accounts)) + for i := range r.accounts { + account := &r.accounts[i] + index[account.ID] = account + } + out := make([]*Account, 0, len(ids)) + seen := make(map[int64]struct{}, len(ids)) + for _, id := range ids { + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + if account, ok := index[id]; ok { + out = append(out, account) + } + } + return out, nil +} + func (r stubOpenAIAccountRepo) ListSchedulableByGroupIDAndPlatform(ctx context.Context, groupID int64, platform string) ([]Account, error) { var result []Account for _, acc := range r.accounts { @@ -641,6 +664,20 @@ func (c *stubGatewayCache) DeleteSessionAccountID(ctx context.Context, groupID i return nil } +func (c *stubGatewayCache) SetGrokVideoPendingBilling(_ context.Context, _ string, _ []byte, _ time.Duration) error { + return nil +} +func (c *stubGatewayCache) GetGrokVideoPendingBilling(_ context.Context, _ string) ([]byte, error) { + return nil, nil +} +func (c *stubGatewayCache) ClaimGrokVideoBilled(_ context.Context, _ string, _ time.Duration) (bool, error) { + return true, nil +} + +func (c *stubGatewayCache) ReleaseGrokVideoBilled(_ context.Context, _ string) error { + return nil +} + func TestOpenAISelectAccountWithLoadAwareness_FiltersUnschedulable(t *testing.T) { now := time.Now() resetAt := now.Add(10 * time.Minute) @@ -2893,10 +2930,30 @@ func TestOpenAIBuildUpstreamRequestOpenAIPassthroughPreservesCompactPath(t *test require.Equal(t, chatgptCodexURL+"/compact", req.URL.String()) require.Equal(t, "application/json", req.Header.Get("Accept")) require.Equal(t, codexCLIVersion, req.Header.Get("Version")) + require.Empty(t, req.Header.Get("OpenAI-Beta"), "Codex OAuth HTTP must not synthesize the legacy responses beta header") require.NotEmpty(t, req.Header.Get("Session_Id")) require.Equal(t, HTTPUpstreamProfileOpenAI, HTTPUpstreamProfileFromContext(req.Context())) } +func TestOpenAIBuildUpstreamRequestOpenAIPassthroughPreservesExplicitAPIKeyBetaHeader(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader([]byte(`{"model":"gpt-5"}`))) + c.Request.Header.Set("OpenAI-Beta", "api-key-specific-beta") + + svc := &OpenAIGatewayService{cfg: &config.Config{ + Security: config.SecurityConfig{ + URLAllowlist: config.URLAllowlistConfig{Enabled: false}, + }, + }} + account := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + + req, err := svc.buildUpstreamRequestOpenAIPassthrough(c.Request.Context(), c, account, []byte(`{"model":"gpt-5"}`), "token") + require.NoError(t, err) + require.Equal(t, "api-key-specific-beta", req.Header.Get("OpenAI-Beta"), "OAuth-only backport must not alter API-key passthrough headers") +} + func TestOpenAIBuildUpstreamRequestCompactForcesJSONAcceptForOAuth(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() @@ -2914,6 +2971,7 @@ func TestOpenAIBuildUpstreamRequestCompactForcesJSONAcceptForOAuth(t *testing.T) require.Equal(t, chatgptCodexURL+"/compact", req.URL.String()) require.Equal(t, "application/json", req.Header.Get("Accept")) require.Equal(t, codexCLIVersion, req.Header.Get("Version")) + require.Empty(t, req.Header.Get("OpenAI-Beta"), "Codex OAuth HTTP must not synthesize the legacy responses beta header") require.NotEmpty(t, req.Header.Get("Session_Id")) require.Equal(t, HTTPUpstreamProfileOpenAI, HTTPUpstreamProfileFromContext(req.Context())) } @@ -3461,6 +3519,32 @@ func TestHandleNonStreamingResponse_OAuthJSONBodyWithDataEventTextKeepsJSONUsage require.Contains(t, rec.Body.String(), "processing data: 1,2,3 then event: click finished") } +func TestHandleNonStreamingResponse_ObservesUpstreamModelBeforeClientRewrite(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + + svc := &OpenAIGatewayService{cfg: &config.Config{}} + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"id":"resp_model_audit","object":"response","model":"gpt-5.5","status":"completed","output":[],"usage":{"input_tokens":3,"output_tokens":2,"total_tokens":5}}`, + )), + } + account := &Account{ID: 1, Type: AccountTypeAPIKey} + + result, err := svc.handleNonStreamingResponse(context.Background(), resp, c, account, "gpt-5.6-sol", "gpt-5.5") + require.NoError(t, err) + require.NotNil(t, result) + + // 客户端仍看到自己请求的模型名,审计观察器则保留改写前的上游声明。 + require.Equal(t, "gpt-5.6-sol", gjson.Get(rec.Body.String(), "model").String()) + require.Equal(t, "gpt-5.5", observedUpstreamResponseModel(c)) + require.False(t, observedUpstreamResponseModelConflict(c)) +} + func TestHandleSSEToJSON_ReconstructsImageGenerationOutputItemDone(t *testing.T) { gin.SetMode(gin.TestMode) rec := httptest.NewRecorder() diff --git a/backend/internal/service/openai_gateway_upstream_errors.go b/backend/internal/service/openai_gateway_upstream_errors.go index f56246e120..e0d3b8cf83 100644 --- a/backend/internal/service/openai_gateway_upstream_errors.go +++ b/backend/internal/service/openai_gateway_upstream_errors.go @@ -211,7 +211,7 @@ func isOpenAIContextWindowError(upstreamMsg string, upstreamBody []byte) bool { func (s *OpenAIGatewayService) shouldFailoverUpstreamError(statusCode int) bool { switch statusCode { - case 401, 402, 403, 429, 529: + case 401, 402, 403, 405, 429, 529: return true default: return statusCode >= 500 diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go index 0725c34cf7..1bba5fce72 100644 --- a/backend/internal/service/openai_gateway_usage.go +++ b/backend/internal/service/openai_gateway_usage.go @@ -257,37 +257,61 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec requestID = upstreamRequestID } } + // Async Grok video: always use the stable task id for dedup (status + content polls + // share one bill). Context-local client/local IDs would otherwise create a new row + // per poll if Redis claim is lost. + if result.VideoCount > 0 { + if stable := StableGrokVideoBillingRequestID(firstNonEmpty( + strings.TrimPrefix(strings.TrimSpace(result.RequestID), "grok-video:"), + strings.TrimSpace(result.ResponseID), + strings.TrimPrefix(strings.TrimSpace(requestID), "grok-video:"), + )); stable != "" { + requestID = stable + } + } // 确定 RequestedModel(渠道映射前的原始模型) requestedModel := result.Model if input.OriginalModel != "" { requestedModel = input.OriginalModel } + sentModel := upstreamSentModel(result.Model, result.UpstreamModel) + if result.UpstreamResponseModelConflict { + logger.L().Warn("upstream_response_model_conflict", + zap.String("platform", account.Platform), + zap.Int64("account_id", account.ID), + zap.String("request_id", requestID), + zap.String("sent_model", sentModel), + zap.String("selected_response_model", strings.TrimSpace(result.UpstreamResponseModel)), + ) + } usageLog := &UsageLog{ - UserID: user.ID, - APIKeyID: apiKey.ID, - AccountID: account.ID, - RequestID: requestID, - Model: result.Model, - RequestedModel: requestedModel, - UpstreamModel: optionalTrimmedStringPtr(result.UpstreamModel), - ServiceTier: result.ServiceTier, - ReasoningEffort: result.ReasoningEffort, - InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint), - UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint), - InputTokens: actualInputTokens, - OutputTokens: result.Usage.OutputTokens, - CacheCreationTokens: result.Usage.CacheCreationInputTokens, - CacheReadTokens: result.Usage.CacheReadInputTokens, - ImageInputTokens: result.Usage.ImageInputTokens, - ImageOutputTokens: result.Usage.ImageOutputTokens, - ImageCount: result.ImageCount, - ImageSize: optionalTrimmedStringPtr(result.ImageSize), - ImageInputSize: optionalTrimmedStringPtr(result.ImageInputSize), - ImageOutputSize: optionalTrimmedStringPtr(result.ImageOutputSize), - ImageSizeSource: optionalTrimmedStringPtr(result.ImageSizeSource), - ImageSizeBreakdown: result.ImageSizeBreakdown, + UserID: user.ID, + APIKeyID: apiKey.ID, + AccountID: account.ID, + RequestID: requestID, + Model: result.Model, + RequestedModel: requestedModel, + UpstreamModel: optionalTrimmedStringPtr(result.UpstreamModel), + UpstreamResponseModel: optionalTrimmedStringPtr(result.UpstreamResponseModel), + UpstreamModelMismatch: upstreamModelMismatch(sentModel, result.UpstreamResponseModel), + ServiceTier: result.ServiceTier, + ReasoningEffort: result.ReasoningEffort, + InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint), + UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint), + InputTokens: actualInputTokens, + OutputTokens: result.Usage.OutputTokens, + CacheCreationTokens: result.Usage.CacheCreationInputTokens, + CacheReadTokens: result.Usage.CacheReadInputTokens, + ImageInputTokens: result.Usage.ImageInputTokens, + ImageOutputTokens: result.Usage.ImageOutputTokens, + ImageCount: result.ImageCount, + ImageSize: optionalTrimmedStringPtr(result.ImageSize), + ImageInputSize: optionalTrimmedStringPtr(result.ImageInputSize), + ImageOutputSize: optionalTrimmedStringPtr(result.ImageOutputSize), + ImageSizeSource: optionalTrimmedStringPtr(result.ImageSizeSource), + ImageSizeBreakdown: result.ImageSizeBreakdown, } isVideoUsage := isGrokVideoUsageResult(result, billingModels) if isVideoUsage { @@ -435,39 +459,83 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( return s.calculateOpenAIVideoCost(ctx, billingModel, apiKey, result, videoMultiplier), nil } } + if result != nil && result.AudioUsage != nil { + cfg := groupAudioPriceConfigFromAPIKey(apiKey) + return s.billingService.CalculateAudioCost(result.AudioUsage.Mode, result.AudioUsage.DurationOrUnits, cfg, webSearchMultiplier), nil + } + if result != nil && result.ImageCount > 0 { // 渠道定价为 token 计费时走 token 路径,否则走图片计费 if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved == nil || resolved.Mode != BillingModeToken { return s.calculateOpenAIImageCost(ctx, billingModel, apiKey, result, imageMultiplier), nil } } - if len(billingModels) == 0 || billingModel == "" { - return nil, errors.New("openai usage billing model is empty") - } + + // Token path (optional search surcharge is additive — never replaces token cost). + var tokenCost *CostBreakdown var lastErr error - for _, candidate := range billingModels { - candidate = strings.TrimSpace(candidate) - if candidate == "" { - continue + if len(billingModels) > 0 && billingModel != "" { + for _, candidate := range billingModels { + candidate = strings.TrimSpace(candidate) + if candidate == "" { + continue + } + cost, err := s.calculateOpenAIRecordUsageTokenCost( + ctx, + apiKey, + candidate, + multiplier, + tokens, + serviceTier, + longContextBillingEnabled, + ) + if err == nil { + tokenCost = cost + break + } + lastErr = err } - cost, err := s.calculateOpenAIRecordUsageTokenCost( - ctx, - apiKey, - candidate, - multiplier, - tokens, - serviceTier, - longContextBillingEnabled, - ) - if err == nil { - return cost, nil + } + // Search surcharge is additive. Never let a zero/default search cost mask a + // real token-pricing failure for requests that attempted token billing. + searchCost := (*CostBreakdown)(nil) + if result != nil && result.SearchCount > 0 { + price := groupSearchPricePer1kFromAPIKey(apiKey) + if price != nil && *price == 0 { + logger.L().Info("openai_usage.search_price_per_1k_explicit_free", + zap.Int("search_count", result.SearchCount), + zap.String("model", billingModel), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + ) } - lastErr = err + searchCost = s.billingService.CalculateSearchCost(result.SearchCount, price, webSearchMultiplier) } - if lastErr == nil { - lastErr = errors.New("no non-empty billing model candidates") + + tokenBillingAttempted := len(billingModels) > 0 && billingModel != "" + if tokenCost == nil { + if tokenBillingAttempted { + if lastErr == nil { + lastErr = errors.New("no non-empty billing model candidates") + } + return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr) + } + // Search-only (no model / pure tool path): allow search billing alone. + if searchCost != nil { + return searchCost, nil + } + if lastErr == nil { + lastErr = errors.New("openai usage billing model is empty") + } + return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr) } - return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr) + if searchCost == nil || (searchCost.TotalCost == 0 && searchCost.ActualCost == 0) { + return tokenCost, nil + } + // Additive: tokens + search surcharge. + tokenCost.TotalCost += searchCost.TotalCost + tokenCost.ActualCost += searchCost.ActualCost + return tokenCost, nil } func isGrokVideoBillingModel(model string) bool { @@ -478,6 +546,8 @@ func isGrokVideoUsageResult(result *OpenAIForwardResult, billingModels []string) if result == nil || result.VideoCount <= 0 { return false } + // VideoCount alone is authoritative for async video completion billing. + // Prefer model-family match when present; never drop video mode on rename/mapping. candidates := append([]string{}, billingModels...) candidates = append(candidates, result.BillingModel, result.Model, result.UpstreamModel) for _, candidate := range candidates { @@ -485,7 +555,7 @@ func isGrokVideoUsageResult(result *OpenAIForwardResult, billingModels []string) return true } } - return false + return true } func isUsagePricingUnavailableError(err error) bool { @@ -586,13 +656,13 @@ func (s *OpenAIGatewayService) calculateOpenAIVideoCost( resolution := NormalizeVideoBillingResolutionOrDefault(result.VideoResolution) durationSeconds := NormalizeVideoBillingDurationSecondsOrDefault(result.VideoDurationSeconds) groupConfig := videoPriceConfigFromAPIKey(apiKey) - if apiKeyHasConfiguredVideoPrice(apiKey, resolution) { + if apiKeyHasConfiguredVideoPrice(apiKey, billingModel, resolution) { return s.billingService.CalculateVideoCost(billingModel, resolution, videoCount, durationSeconds, groupConfig, multiplier) } if refreshed := s.apiKeyWithFreshGroupMediaPricing(ctx, apiKey); refreshed != apiKey { apiKey = refreshed groupConfig = videoPriceConfigFromAPIKey(apiKey) - if apiKeyHasConfiguredVideoPrice(apiKey, resolution) { + if apiKeyHasConfiguredVideoPrice(apiKey, billingModel, resolution) { return s.billingService.CalculateVideoCost(billingModel, resolution, videoCount, durationSeconds, groupConfig, multiplier) } } @@ -639,10 +709,14 @@ func (s *OpenAIGatewayService) apiKeyWithFreshGroupMediaPricing(ctx context.Cont return &clone } -// groupMediaPricingLooksIncomplete 判断分组对象是否可能缺失媒体计费字段(例如由不含 -// 这些字段的旧快照或手工构造的上下文对象生成)。image/video 独立倍率在数据库中的 -// 默认值均为 1.0,正常加载的分组不可能两个倍率同时为 0 且未开启独立倍率、全部媒体 -// 价为 nil——只有这种情况才回源查库,避免对未配置覆盖价的分组每条媒体用量都多打一次 DB 查询。 +// groupMediaPricingLooksIncomplete 判断分组对象是否可能缺失媒体/搜索/语音计费字段 +// (例如由不含这些字段的旧快照或手工构造的上下文对象生成)。image/video 独立倍率在 +// 数据库中的默认值均为 1.0;正常加载的分组不可能两个倍率同时为 0 且未开启独立倍率、 +// 全部媒体/搜索/语音价为 nil——只有这种情况才回源查库,避免对未配置覆盖价的分组每条 +// 用量都多打一次 DB 查询。 +// +// 注意:apiKeyAuthSnapshotVersion 升级会强制刷新存量快照;本函数是热路径上的二次兜底, +// 不能仅凭 legacy video_price_* 判定完整而跳过 VideoModelPrices/search/audio 的回源。 func groupMediaPricingLooksIncomplete(group *Group) bool { if group == nil { return true @@ -653,6 +727,17 @@ func groupMediaPricingLooksIncomplete(group *Group) bool { if group.ImageRateMultiplier != 0 || group.VideoRateMultiplier != 0 { return false } + // Any first-class pricing field present means the projection is not a blank shell. + if len(group.VideoModelPrices) > 0 { + return false + } + if group.SearchPricePer1k != nil || + group.AudioRealtimePricePerMin != nil || + group.AudioTTSPricePerMillionChars != nil || + group.AudioSTTPricePerHour != nil || + group.WebSearchPricePerCall != nil { + return false + } return group.ImagePrice1K == nil && group.ImagePrice2K == nil && group.ImagePrice4K == nil && group.VideoPrice480P == nil && group.VideoPrice720P == nil && group.VideoPrice1080P == nil } diff --git a/backend/internal/service/openai_images_test.go b/backend/internal/service/openai_images_test.go index 82f9c3addc..e3edbaa5ae 100644 --- a/backend/internal/service/openai_images_test.go +++ b/backend/internal/service/openai_images_test.go @@ -767,7 +767,7 @@ func TestOpenAIGatewayServiceForwardImages_OAuthPassesNAndReturnsAllImages(t *te require.Equal(t, "application/json", upstream.lastReq.Header.Get("Content-Type")) require.Equal(t, "text/event-stream", upstream.lastReq.Header.Get("Accept")) require.Equal(t, "acct-123", upstream.lastReq.Header.Get("chatgpt-account-id")) - require.Equal(t, "responses=experimental", upstream.lastReq.Header.Get("OpenAI-Beta")) + require.Empty(t, upstream.lastReq.Header.Get("OpenAI-Beta")) require.Equal(t, openAIImagesResponsesMainModel, gjson.GetBytes(upstream.lastBody, "model").String()) require.True(t, gjson.GetBytes(upstream.lastBody, "stream").Bool()) diff --git a/backend/internal/service/openai_messages_dispatch.go b/backend/internal/service/openai_messages_dispatch.go index f72b84a7f3..aedfb1b3f7 100644 --- a/backend/internal/service/openai_messages_dispatch.go +++ b/backend/internal/service/openai_messages_dispatch.go @@ -69,10 +69,14 @@ func (g *Group) ResolveMessagesDispatchModel(requestedModel string) string { } if g.Platform == PlatformGrok { - if claudeMessagesDispatchFamily(requestedModel) != "" { - return xai.DefaultModelMapping()["grok"] + if claudeMessagesDispatchFamily(requestedModel) == "" { + return "" } - return "" + opts := xai.RuntimeModelMappingOptions() + if !opts.EnableCrossClientMap { + return "" + } + return xai.ModelMappingWithOptions(opts)["claude-*"] } cfg := normalizeOpenAIMessagesDispatchModelConfig(g.MessagesDispatchModelConfig) diff --git a/backend/internal/service/openai_messages_dispatch_test.go b/backend/internal/service/openai_messages_dispatch_test.go index db7804a4f3..fd6dea86fa 100644 --- a/backend/internal/service/openai_messages_dispatch_test.go +++ b/backend/internal/service/openai_messages_dispatch_test.go @@ -1,8 +1,11 @@ package service -import "testing" +import ( + "testing" -import "github.com/stretchr/testify/require" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" + "github.com/stretchr/testify/require" +) func TestNormalizeOpenAIMessagesDispatchModelConfig(t *testing.T) { t.Parallel() @@ -26,14 +29,21 @@ func TestNormalizeOpenAIMessagesDispatchModelConfig(t *testing.T) { }, cfg.ExactModelMappings) } -func TestGroupResolveMessagesDispatchModel_GrokMapsClaudeFamilyToGrok(t *testing.T) { - t.Parallel() - +func TestGroupResolveMessagesDispatchModel_GrokRequiresCrossClientMapping(t *testing.T) { + original := xai.RuntimeModelMappingOptions() + t.Cleanup(func() { xai.SetRuntimeModelMappingOptions(original) }) group := &Group{Platform: PlatformGrok} - require.Equal(t, "grok-4.5", group.ResolveMessagesDispatchModel("claude-sonnet-4-5")) - require.Equal(t, "grok-4.5", group.ResolveMessagesDispatchModel("claude-opus-4-6")) - require.Equal(t, "grok-4.5", group.ResolveMessagesDispatchModel("claude-haiku-4-5")) + xai.SetRuntimeModelMappingOptions(xai.ModelMappingOptions{}) + require.Empty(t, group.ResolveMessagesDispatchModel("claude-sonnet-4-5")) + + xai.SetRuntimeModelMappingOptions(xai.ModelMappingOptions{ + DefaultText: "grok-build-0.1", + EnableCrossClientMap: true, + }) + require.Equal(t, "grok-build-0.1", group.ResolveMessagesDispatchModel("claude-sonnet-4-5")) + require.Equal(t, "grok-build-0.1", group.ResolveMessagesDispatchModel("claude-opus-4-6")) + require.Equal(t, "grok-build-0.1", group.ResolveMessagesDispatchModel("claude-haiku-4-5")) require.Empty(t, group.ResolveMessagesDispatchModel("grok")) require.Empty(t, group.ResolveMessagesDispatchModel("gpt-5.3-codex")) } diff --git a/backend/internal/service/openai_responses_tool_schema.go b/backend/internal/service/openai_responses_tool_schema.go new file mode 100644 index 0000000000..8c7e7b8d18 --- /dev/null +++ b/backend/internal/service/openai_responses_tool_schema.go @@ -0,0 +1,127 @@ +package service + +import ( + "bytes" + "sort" + + "github.com/tidwall/gjson" +) + +const ( + // 工具定义在多轮历史里最多再嵌套一层 tools,留出余量后截断,避免畸形请求体 + // 造成无界递归。 + openAIResponsesToolSchemaMaxDepth = 4 + // JSON Schema 里 type 只能是字符串或字符串数组;显式 null 无论哪个方言都非法, + // 补成 object 与 upstream 对该工具的实际期望一致。 + openAIResponsesToolSchemaFallbackType = `"object"` + // 显式 null 在 JSON 里只有这一种字面量形态。 + openAIResponsesToolSchemaNullLiteral = "null" +) + +// openAIResponsesToolSchemaNullType 记录一处待修正的 null,用原始 body 上的 +// 绝对字节偏移表示,便于最后一次性拼接。 +type openAIResponsesToolSchemaNullType struct { + offset int + length int +} + +// sanitizeOpenAIResponsesToolParameterTypes 修正请求体中显式为 null 的 +// tools[].parameters.type。 +// +// Codex Desktop 内置的 automation_update 工具会带 parameters.type = null, +// OpenAI 直接回 400 invalid_function_parameters,而网关把该状态归一成可重试的 +// 502 upstream_error;该工具定义又会沉进多轮历史,导致之后每一轮继续失败并在 +// 账号池里反复重放同一份坏 Schema。 +// +// 只修正显式 null:缺失 type 的 Schema 本身合法(等价于不约束),补写会收窄 +// 客户端语义,因此保持原样。 +// +// 实现上先收集全部命中的绝对偏移,再一次性拼出新 body:逐个 sjson.SetBytes 每次 +// 都会重扫并全量拷贝整个文档,命中 N 处就是 N 次全量拷贝,而 /v1/responses 的 +// body 上限是 gateway.max_body_size(默认 256MB),构造请求能塞进百万级命中。 +func sanitizeOpenAIResponsesToolParameterTypes(body []byte) ([]byte, bool, error) { + if len(body) == 0 { + return body, false, nil + } + + hits := make([]openAIResponsesToolSchemaNullType, 0, 2) + collectOpenAIResponsesToolSchemaNullTypes(body, gjson.GetBytes(body, "tools"), 0, &hits) + if input := gjson.GetBytes(body, "input"); input.IsArray() { + input.ForEach(func(_, item gjson.Result) bool { + if item.IsObject() { + collectOpenAIResponsesToolSchemaNullTypes(body, item.Get("tools"), 0, &hits) + } + return true + }) + } + if len(hits) == 0 { + return body, false, nil + } + + // tools 与 input 在 body 里的先后顺序由客户端决定,收集顺序不保证单调。 + sort.Slice(hits, func(i, j int) bool { return hits[i].offset < hits[j].offset }) + + sanitized := make([]byte, 0, len(body)+len(hits)*len(openAIResponsesToolSchemaFallbackType)) + cursor := 0 + for _, hit := range hits { + // 收集阶段已逐个校验过区间,这里再挡一次重叠,保证拼接严格单调向前。 + if hit.offset < cursor { + continue + } + sanitized = append(sanitized, body[cursor:hit.offset]...) + sanitized = append(sanitized, openAIResponsesToolSchemaFallbackType...) + cursor = hit.offset + hit.length + } + sanitized = append(sanitized, body[cursor:]...) + return sanitized, true, nil +} + +// collectOpenAIResponsesToolSchemaNullTypes 收集一个 tools 数组里所有需要修正的 +// parameters.type 位置。不按 tool type 过滤:null 的 schema type 在 function、 +// custom 以及任何 hosted 工具上都同样非法。 +func collectOpenAIResponsesToolSchemaNullTypes( + body []byte, tools gjson.Result, depth int, hits *[]openAIResponsesToolSchemaNullType, +) { + if depth > openAIResponsesToolSchemaMaxDepth || !tools.IsArray() { + return + } + tools.ForEach(func(_, tool gjson.Result) bool { + if !tool.IsObject() { + return true + } + // Responses 形态用顶层 parameters,ChatCompletions 形态用 function.parameters, + // 两种都可能出现在 Responses 请求里(见 normalizeCodexTools)。 + for _, suffix := range []string{"parameters", "function.parameters"} { + params := tool.Get(suffix) + if !params.IsObject() { + continue + } + // gjson 用 Type==Null 同时表示「显式 null」和「路径不存在」,靠 Raw + // 区分:不存在时 Raw 为空串。 + if typ := params.Get("type"); typ.Type == gjson.Null && typ.Raw == openAIResponsesToolSchemaNullLiteral { + appendOpenAIResponsesToolSchemaNullType(body, typ, hits) + } + } + // 历史输入里的工具定义会再嵌套一层 tools(upstream 报错路径形如 + // input[234].tools[0].tools[3].parameters)。 + collectOpenAIResponsesToolSchemaNullTypes(body, tool.Get("tools"), depth+1, hits) + return true + }) +} + +// appendOpenAIResponsesToolSchemaNullType 先校验 gjson 给出的偏移确实指向原始 +// body 上那段 null,再记录。gjson 对嵌套取值同样返回相对原始文档的绝对偏移,但 +// Index 为 0 表示未知;偏移不可用时跳过该处而不是猜位置——少修一个工具只是维持 +// 现状,拼错位置会损坏整个请求体。 +func appendOpenAIResponsesToolSchemaNullType( + body []byte, typ gjson.Result, hits *[]openAIResponsesToolSchemaNullType, +) { + end := typ.Index + len(typ.Raw) + if typ.Index <= 0 || end > len(body) { + return + } + if !bytes.Equal(body[typ.Index:end], []byte(typ.Raw)) { + return + } + *hits = append(*hits, openAIResponsesToolSchemaNullType{offset: typ.Index, length: len(typ.Raw)}) +} diff --git a/backend/internal/service/openai_responses_tool_schema_test.go b/backend/internal/service/openai_responses_tool_schema_test.go new file mode 100644 index 0000000000..2dbc7fb600 --- /dev/null +++ b/backend/internal/service/openai_responses_tool_schema_test.go @@ -0,0 +1,249 @@ +package service + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +// issue #5364 的最小复现体:Codex Desktop 内置 automation_update 带 +// parameters.type = null,upstream 回 400 invalid_function_parameters。 +func TestSanitizeOpenAIResponsesToolParameterTypes_TopLevelFunctionTool(t *testing.T) { + body := []byte(`{ + "model": "gpt-5.6-sol", + "input": "Reply with OK.", + "stream": false, + "tools": [ + { + "type": "function", + "name": "automation_update", + "description": "Update an automation.", + "parameters": {"type": null, "properties": {}} + } + ] + }`) + + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes(body) + + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(sanitized, "tools.0.parameters.type").String()) + // 只改 type,工具其余定义原样保留。 + require.Equal(t, "automation_update", gjson.GetBytes(sanitized, "tools.0.name").String()) + require.Equal(t, "Update an automation.", gjson.GetBytes(sanitized, "tools.0.description").String()) + require.True(t, gjson.GetBytes(sanitized, "tools.0.parameters.properties").IsObject()) + // 请求体其余字段不受影响。 + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(sanitized, "model").String()) + require.Equal(t, "Reply with OK.", gjson.GetBytes(sanitized, "input").String()) +} + +// 合法 Schema 必须原样返回:changed=false 且字节不变,避免无谓重写打散 +// prompt cache 前缀。 +func TestSanitizeOpenAIResponsesToolParameterTypes_ValidSchemaUntouched(t *testing.T) { + body := []byte(`{"tools":[{"type":"function","name":"ok","parameters":{"type":"object","properties":{}}}]}`) + + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes(body) + + require.NoError(t, err) + require.False(t, changed) + require.Equal(t, string(body), string(sanitized)) +} + +// 缺失 type 的 Schema 本身合法(等价于不约束),不得补写——补写会收窄客户端语义。 +func TestSanitizeOpenAIResponsesToolParameterTypes_MissingTypeNotInvented(t *testing.T) { + body := []byte(`{"tools":[{"type":"function","name":"ok","parameters":{"properties":{}}}]}`) + + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes(body) + + require.NoError(t, err) + require.False(t, changed) + require.False(t, gjson.GetBytes(sanitized, "tools.0.parameters.type").Exists()) +} + +// 多轮历史:工具定义沉进 input 后,upstream 报错路径形如 +// input[N].tools[i].tools[j].parameters,两层都要修。 +func TestSanitizeOpenAIResponsesToolParameterTypes_NestedHistoryTools(t *testing.T) { + body := []byte(`{ + "input": [ + {"type": "message", "role": "user", "content": "hi"}, + { + "type": "additional_tools", + "role": "developer", + "tools": [ + { + "type": "namespace", + "name": "codex_app", + "tools": [ + {"type": "function", "name": "noop", "parameters": {"type": "object"}}, + {"type": "function", "name": "automation_update", "parameters": {"type": null}} + ] + }, + {"type": "function", "name": "outer", "parameters": {"type": null}} + ] + } + ] + }`) + + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes(body) + + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(sanitized, "input.1.tools.0.tools.1.parameters.type").String()) + require.Equal(t, "object", gjson.GetBytes(sanitized, "input.1.tools.1.parameters.type").String()) + // 原本合法的兄弟条目保持不变。 + require.Equal(t, "object", gjson.GetBytes(sanitized, "input.1.tools.0.tools.0.parameters.type").String()) + require.Equal(t, "hi", gjson.GetBytes(sanitized, "input.0.content").String()) +} + +// ChatCompletions 形态的工具({type:"function", function:{...}})同样可能出现在 +// Responses 请求里,见 normalizeCodexTools。 +func TestSanitizeOpenAIResponsesToolParameterTypes_ChatCompletionsShape(t *testing.T) { + body := []byte(`{"tools":[{"type":"function","function":{"name":"legacy","parameters":{"type":null}}}]}`) + + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes(body) + + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(sanitized, "tools.0.function.parameters.type").String()) + require.Equal(t, "legacy", gjson.GetBytes(sanitized, "tools.0.function.name").String()) +} + +// 索引映射:只有坏条目被改,前后兄弟条目按原下标保持不变。 +func TestSanitizeOpenAIResponsesToolParameterTypes_OnlyOffendingIndexRewritten(t *testing.T) { + body := []byte(`{"tools":[ + {"type":"function","name":"a","parameters":{"type":"object"}}, + {"type":"function","name":"b","parameters":{"type":null}}, + {"type":"function","name":"c","parameters":{"type":"object"}} + ]}`) + + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes(body) + + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "a", gjson.GetBytes(sanitized, "tools.0.name").String()) + require.Equal(t, "b", gjson.GetBytes(sanitized, "tools.1.name").String()) + require.Equal(t, "c", gjson.GetBytes(sanitized, "tools.2.name").String()) + require.Equal(t, "object", gjson.GetBytes(sanitized, "tools.1.parameters.type").String()) + require.Equal(t, 3, int(gjson.GetBytes(sanitized, "tools.#").Int())) +} + +// 畸形/非常规形态不得 panic,且一律按不变处理。 +func TestSanitizeOpenAIResponsesToolParameterTypes_MalformedShapesAreNoOps(t *testing.T) { + cases := []struct { + name string + body string + }{ + {"empty_body", ``}, + {"no_tools", `{"model":"gpt-5.6-sol","input":"hi"}`}, + {"tools_null", `{"tools":null}`}, + {"tools_object", `{"tools":{"type":"function"}}`}, + {"tool_is_string", `{"tools":["freeform"]}`}, + {"parameters_is_string", `{"tools":[{"type":"function","parameters":"nope"}]}`}, + {"parameters_null", `{"tools":[{"type":"function","parameters":null}]}`}, + {"input_string", `{"input":"hi","tools":[]}`}, + {"input_item_not_object", `{"input":["hi"]}`}, + {"type_already_array", `{"tools":[{"type":"function","parameters":{"type":["object","null"]}}]}`}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes([]byte(tc.body)) + require.NoError(t, err) + require.False(t, changed) + require.Equal(t, tc.body, string(sanitized)) + }) + } +} + +// 递归深度守卫:超深嵌套只做截断,不递归到栈溢出,也不报错。 +func TestSanitizeOpenAIResponsesToolParameterTypes_DepthGuard(t *testing.T) { + tool := map[string]any{"type": "function", "name": "deep", "parameters": map[string]any{"type": nil}} + for i := 0; i < 12; i++ { + tool = map[string]any{"type": "namespace", "tools": []any{tool}} + } + body, err := json.Marshal(map[string]any{"tools": []any{tool}}) + require.NoError(t, err) + + require.NotPanics(t, func() { + _, _, sanitizeErr := sanitizeOpenAIResponsesToolParameterTypes(body) + require.NoError(t, sanitizeErr) + }) +} + +// 输出必须是合法 JSON,且除目标字段外与输入等价。 +func TestSanitizeOpenAIResponsesToolParameterTypes_OutputStaysValidJSON(t *testing.T) { + body := []byte(`{"model":"gpt-5.5","tool_choice":"none","store":false,"tools":[{"type":"function","name":"automation_update","parameters":{"type":null,"properties":{}}}]}`) + + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes(body) + require.NoError(t, err) + require.True(t, changed) + + var decoded map[string]any + require.NoError(t, json.Unmarshal(sanitized, &decoded)) + require.Equal(t, "gpt-5.5", decoded["model"]) + require.Equal(t, "none", decoded["tool_choice"]) + require.Equal(t, false, decoded["store"]) +} + +// 输入 body 是调用方持有的缓冲区(Forward 里 canonicalImageIntentBody 与它同源), +// 净化必须返回新切片,绝不能就地改写。 +func TestSanitizeOpenAIResponsesToolParameterTypes_DoesNotMutateInputBody(t *testing.T) { + body := []byte(`{"model":"gpt-5.6-sol","tools":[{"type":"function","name":"a","parameters":{"type":null}}]}`) + original := append([]byte(nil), body...) + + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes(body) + + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, string(original), string(body), "调用方的 body 不得被就地改写") + require.NotEqual(t, string(original), string(sanitized)) +} + +func buildToolSchemaNullTypeBody(t *testing.T, hits int) []byte { + t.Helper() + tools := make([]any, 0, hits) + for i := 0; i < hits; i++ { + tools = append(tools, map[string]any{ + "type": "function", + "name": "automation_update", + "parameters": map[string]any{"type": nil, "properties": map[string]any{}}, + }) + } + body, err := json.Marshal(map[string]any{"model": "gpt-5.6-sol", "tools": tools}) + require.NoError(t, err) + return body +} + +// 复杂度守卫:重写次数必须与命中数无关。 +// +// 逐个 sjson.SetBytes 的写法每命中一处就重扫并全量拷贝一次文档,命中 N 处即 N 次 +// 全量拷贝;/v1/responses 的 body 上限是 gateway.max_body_size(默认 256MB), +// 构造请求可以塞进百万级命中,会被放大成 TB 级 memcpy。这里用分配次数锁死该行为: +// 命中数放大 500 倍,分配次数不得随之增长。 +func TestSanitizeOpenAIResponsesToolParameterTypes_RewriteCountIndependentOfHits(t *testing.T) { + small := buildToolSchemaNullTypeBody(t, 4) + large := buildToolSchemaNullTypeBody(t, 2000) + + smallAllocs := testing.AllocsPerRun(2, func() { + _, _, _ = sanitizeOpenAIResponsesToolParameterTypes(small) + }) + largeAllocs := testing.AllocsPerRun(2, func() { + _, _, _ = sanitizeOpenAIResponsesToolParameterTypes(large) + }) + + // 命中切片扩容是对数级,留出充裕余量;线性写法在这里会是 2000 量级。 + require.Less(t, largeAllocs, smallAllocs+40, + "分配次数随命中数线性增长,说明退回了逐路径全量重写 (small=%v large=%v)", smallAllocs, largeAllocs) + + // 同时确认大 body 的结果确实全部修好了。 + sanitized, changed, err := sanitizeOpenAIResponsesToolParameterTypes(large) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, 2000, int(gjson.GetBytes(sanitized, "tools.#").Int())) + gjson.GetBytes(sanitized, "tools").ForEach(func(_, tool gjson.Result) bool { + require.Equal(t, "object", tool.Get("parameters.type").String()) + return true + }) +} diff --git a/backend/internal/service/openai_routing_hint.go b/backend/internal/service/openai_routing_hint.go new file mode 100644 index 0000000000..0a1cc928cd --- /dev/null +++ b/backend/internal/service/openai_routing_hint.go @@ -0,0 +1,128 @@ +package service + +import ( + "context" + "net/http" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/tidwall/gjson" + "go.uber.org/zap" + "golang.org/x/net/http/httpguts" +) + +const openAICodexRoutingHintHeader = "x-codex-routing-hint" + +// setOpenAICodexRoutingHint mirrors the Codex backend routing-hint contract for +// OpenAI OAuth requests. The request model must already be the final upstream +// slug and serviceTier must already reflect any local policy rewrite/filter. +func setOpenAICodexRoutingHint(headers http.Header, account *Account, model string, serviceTier string) { + if headers == nil { + return + } + + // The routing hint is gateway-owned. Strip every spelling before deciding + // whether to synthesize it so API-key/provider credential paths cannot pass + // through a caller- or account-override-supplied hint. http.Header.Del only + // removes the canonical map key; inbound maps can contain raw lowercase keys. + deleteOpenAIHeaderEqualFold(headers, openAICodexRoutingHintHeader) + if account == nil || !account.IsOpenAIOAuth() { + return + } + + model = strings.TrimSpace(model) + if model == "" || strings.ContainsAny(model, ";=") { + return + } + + // Codex treats "default" as an explicit standard-routing sentinel, not as a + // service tier sent to the backend. Fast follows the gateway's existing + // canonicalization and therefore becomes "priority"; flex stays "flex". + canonicalTier := normalizedOpenAIServiceTierValue(serviceTier) + // This backport has no Codex model-catalog snapshot with which to validate + // arbitrary tier ids. Keep the hint to the two effective tiers Codex itself + // selects; default, missing, and other gateway-compatible API values remain + // model-only rather than expanding the ChatGPT routing protocol here. + switch canonicalTier { + case OpenAIFastTierPriority, OpenAIFastTierFlex: + default: + canonicalTier = "" + } + + hint := "model=" + model + if canonicalTier != "" { + hint += ";tier=" + canonicalTier + } + if !httpguts.ValidHeaderFieldValue(hint) { + return + } + headers.Set(openAICodexRoutingHintHeader, hint) +} + +func deleteOpenAIHeaderEqualFold(headers http.Header, name string) { + if headers == nil { + return + } + name = strings.TrimSpace(name) + for key := range headers { + if strings.EqualFold(strings.TrimSpace(key), name) { + delete(headers, key) + } + } +} + +func setOpenAICodexRoutingHintFromBody(headers http.Header, account *Account, body []byte) { + fields := gjson.GetManyBytes(body, "model", "service_tier") + setOpenAICodexRoutingHint(headers, account, fields[0].String(), fields[1].String()) +} + +// logOpenAIRoutingDiagnostics records only gateway-derived routing state. In +// particular, it deliberately does not include any header values, tokens, or +// credentials because these diagnostics run on authentication-bearing paths. +func logOpenAIRoutingDiagnostics( + ctx context.Context, + account *Account, + transport string, + model string, + serviceTier string, + hintGenerated bool, + wsAffinityDecision string, +) { + if ctx == nil { + ctx = context.Background() + } + accountID := int64(0) + if account != nil { + accountID = account.ID + } + + logger.FromContext(ctx).Debug("openai routing decision", + zap.String("component", "service.openai_routing"), + zap.String("transport", strings.TrimSpace(transport)), + zap.Int64("account_id", accountID), + zap.String("final_model", strings.TrimSpace(model)), + zap.String("final_service_tier", normalizedOpenAIServiceTierValue(serviceTier)), + zap.Bool("routing_hint_generated", hintGenerated), + zap.String("ws_affinity_decision", strings.TrimSpace(wsAffinityDecision)), + ) +} + +func logOpenAIRoutingDiagnosticsFromBody( + ctx context.Context, + account *Account, + transport string, + headers http.Header, + body []byte, + wsAffinityDecision string, +) { + fields := gjson.GetManyBytes(body, "model", "service_tier") + logOpenAIRoutingDiagnostics( + ctx, + account, + transport, + fields[0].String(), + fields[1].String(), + strings.TrimSpace(headers.Get(openAICodexRoutingHintHeader)) != "", + wsAffinityDecision, + ) +} diff --git a/backend/internal/service/openai_routing_hint_test.go b/backend/internal/service/openai_routing_hint_test.go new file mode 100644 index 0000000000..5801b15fa4 --- /dev/null +++ b/backend/internal/service/openai_routing_hint_test.go @@ -0,0 +1,369 @@ +package service + +import ( + "bytes" + "context" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestSetOpenAICodexRoutingHintCanonicalizesOfficialServiceTiers(t *testing.T) { + oauthAccount := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth} + tests := []struct { + name string + model string + serviceTier string + want string + }{ + {name: "fast alias", model: "gpt-5.6", serviceTier: " fast ", want: "model=gpt-5.6;tier=priority"}, + {name: "priority", model: "gpt-5.6", serviceTier: "priority", want: "model=gpt-5.6;tier=priority"}, + {name: "flex", model: "gpt-5.6", serviceTier: "flex", want: "model=gpt-5.6;tier=flex"}, + {name: "explicit default sentinel", model: "gpt-5.6", serviceTier: "default", want: "model=gpt-5.6"}, + {name: "omitted tier", model: "gpt-5.6", want: "model=gpt-5.6"}, + {name: "auto is not expanded without catalog support", model: "gpt-5.6", serviceTier: "auto", want: "model=gpt-5.6"}, + {name: "scale is not expanded without catalog support", model: "gpt-5.6", serviceTier: "scale", want: "model=gpt-5.6"}, + {name: "unknown tier does not expand protocol", model: "gpt-5.6", serviceTier: "turbo", want: "model=gpt-5.6"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + headers := make(http.Header) + setOpenAICodexRoutingHint(headers, oauthAccount, tt.model, tt.serviceTier) + require.Equal(t, tt.want, headers.Get(openAICodexRoutingHintHeader)) + }) + } + + t.Run("invalid header value is omitted", func(t *testing.T) { + headers := make(http.Header) + setOpenAICodexRoutingHint(headers, oauthAccount, "gpt-5.6\ninvalid", "priority") + require.Empty(t, headers.Get(openAICodexRoutingHintHeader)) + }) + + for _, model := range []string{"gpt-5.6;evil", "gpt=5.6"} { + t.Run("delimiter in model is omitted: "+model, func(t *testing.T) { + headers := make(http.Header) + setOpenAICodexRoutingHint(headers, oauthAccount, model, "priority") + require.Empty(t, headers.Get(openAICodexRoutingHintHeader)) + }) + } + + t.Run("api key strips gateway-owned hint in every key casing", func(t *testing.T) { + headers := make(http.Header) + headers[openAICodexRoutingHintHeader] = []string{"lowercase-spoof"} + headers["X-Codex-Routing-Hint"] = []string{"canonical-spoof"} + setOpenAICodexRoutingHint(headers, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, "gpt-5.6", "priority") + for key := range headers { + require.False(t, strings.EqualFold(key, openAICodexRoutingHintHeader)) + } + }) + + t.Run("oauth replaces spoofed lowercase hint", func(t *testing.T) { + headers := make(http.Header) + headers[openAICodexRoutingHintHeader] = []string{"model=spoof;tier=flex"} + setOpenAICodexRoutingHint(headers, oauthAccount, "gpt-5.6", "priority") + require.Equal(t, "model=gpt-5.6;tier=priority", headers.Get(openAICodexRoutingHintHeader)) + require.Len(t, headers, 1) + }) +} + +func TestOpenAIOAuthHTTPBuildersSendRoutingHintFromFinalBody(t *testing.T) { + gin.SetMode(gin.TestMode) + oauthAccount := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "test-account", + }, + } + svc := &OpenAIGatewayService{} + + tests := []struct { + name string + body []byte + want string + }{ + {name: "fast", body: []byte(`{"model":"gpt-5.6-codex","service_tier":"fast"}`), want: "model=gpt-5.6-codex;tier=priority"}, + {name: "flex", body: []byte(`{"model":"gpt-5.6-codex","service_tier":"flex"}`), want: "model=gpt-5.6-codex;tier=flex"}, + {name: "default", body: []byte(`{"model":"gpt-5.6-codex","service_tier":"default"}`), want: "model=gpt-5.6-codex"}, + {name: "omitted", body: []byte(`{"model":"gpt-5.6-codex"}`), want: "model=gpt-5.6-codex"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + for _, passthrough := range []bool{false, true} { + mode := "ordinary" + if passthrough { + mode = "passthrough" + } + t.Run(mode, func(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(tt.body)) + + var req *http.Request + var err error + if passthrough { + req, err = svc.buildUpstreamRequestOpenAIPassthrough(context.Background(), c, oauthAccount, tt.body, "test-token") + } else { + req, err = svc.buildUpstreamRequest(context.Background(), c, oauthAccount, tt.body, "test-token", false, "", true) + } + require.NoError(t, err) + require.Equal(t, tt.want, req.Header.Get(openAICodexRoutingHintHeader)) + }) + } + }) + } +} + +func TestOpenAIHTTPPassthroughStripsOnlyOAuthLegacyResponsesBeta(t *testing.T) { + gin.SetMode(gin.TestMode) + svc := &OpenAIGatewayService{cfg: &config.Config{ + Security: config.SecurityConfig{ + URLAllowlist: config.URLAllowlistConfig{Enabled: false}, + }, + }} + body := []byte(`{"model":"gpt-5.6-codex","service_tier":"priority"}`) + + build := func(t *testing.T, account *Account, betaValues []string, rawLowercaseKey bool) http.Header { + t.Helper() + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + if rawLowercaseKey { + c.Request.Header["openai-beta"] = append([]string(nil), betaValues...) + } else { + for _, value := range betaValues { + c.Request.Header.Add("OpenAI-Beta", value) + } + } + + req, err := svc.buildUpstreamRequestOpenAIPassthrough(context.Background(), c, account, body, "test-token") + require.NoError(t, err) + return req.Header + } + + oauth := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "test-account", + }, + } + apiKey := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "test-api-key", + }, + } + + t.Run("oauth legacy only is removed including raw lowercase key", func(t *testing.T) { + headers := build(t, oauth, []string{"responses=experimental"}, true) + require.Empty(t, headers.Values("OpenAI-Beta")) + }) + + t.Run("oauth mixed beta preserves independent tokens", func(t *testing.T) { + headers := build(t, oauth, []string{ + "responses=experimental, future_feature=v1", + "another_feature=v2, RESPONSES=EXPERIMENTAL", + }, false) + require.Equal(t, []string{"future_feature=v1", "another_feature=v2"}, headers.Values("OpenAI-Beta")) + }) + + t.Run("api key explicit beta remains caller controlled", func(t *testing.T) { + headers := build(t, apiKey, []string{"responses=experimental, future_feature=v1"}, false) + require.Equal(t, []string{"responses=experimental, future_feature=v1"}, headers.Values("OpenAI-Beta")) + }) +} + +func TestBuildOpenAIWSHeadersSendsOAuthRoutingHintOnly(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + svc := &OpenAIGatewayService{} + decision := OpenAIWSProtocolDecision{Transport: OpenAIUpstreamTransportResponsesWebsocketV2} + + build := func(t *testing.T, account *Account, tier string) http.Header { + headers, _, err := svc.buildOpenAIWSHeaders( + context.Background(), + c, + account, + "test-token", + decision, + true, + "", + "", + "", + "gpt-5.6-codex", + tier, + ) + require.NoError(t, err) + return headers + } + + oauthAccount := &Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "test-account", + }, + } + require.Equal(t, "model=gpt-5.6-codex;tier=priority", build(t, oauthAccount, "fast").Get(openAICodexRoutingHintHeader)) + require.Equal(t, "model=gpt-5.6-codex", build(t, oauthAccount, "default").Get(openAICodexRoutingHintHeader)) + require.Empty(t, build(t, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}, "priority").Get(openAICodexRoutingHintHeader)) +} + +func TestOpenAIRoutingDiagnosticsUseFinalDerivedValuesOnly(t *testing.T) { + gin.SetMode(gin.TestMode) + logSink, restore := captureStructuredLog(t) + defer restore() + + account := &Account{ + ID: 917, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + "chatgpt_account_id": "chatgpt-account", + }, + } + body := []byte(`{"model":"gpt-5.6-codex","service_tier":"fast"}`) + svc := &OpenAIGatewayService{} + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body)) + c.Request.Header.Set("Authorization", "Bearer caller-secret") + c.Request.Header.Set(openAICodexRoutingHintHeader, "model=caller-secret") + _, err := svc.buildUpstreamRequest(context.Background(), c, account, body, "oauth-secret", false, "", true) + require.NoError(t, err) + + decision := OpenAIWSProtocolDecision{Transport: OpenAIUpstreamTransportResponsesWebsocketV2} + _, _, err = svc.buildOpenAIWSHeaders( + context.Background(), c, account, "oauth-secret", decision, true, + "", "", "", "gpt-5.6-codex", "fast", + ) + require.NoError(t, err) + + require.True(t, logSink.ContainsMessageAtLevel("openai routing decision", "debug")) + require.True(t, logSink.ContainsFieldValue("account_id", "917")) + require.True(t, logSink.ContainsFieldValue("final_model", "gpt-5.6-codex")) + require.True(t, logSink.ContainsFieldValue("final_service_tier", "priority")) + require.True(t, logSink.ContainsFieldValue("routing_hint_generated", "true")) + require.True(t, logSink.ContainsFieldValue("transport", "http")) + require.True(t, logSink.ContainsFieldValue("transport", string(OpenAIUpstreamTransportResponsesWebsocketV2))) + require.True(t, logSink.ContainsFieldValue("ws_affinity_decision", "not_applicable")) + require.True(t, logSink.ContainsFieldValue("ws_affinity_decision", "soft_routing_hint")) + require.False(t, logSink.ContainsFieldValue("authorization", "caller-secret")) + require.False(t, logSink.ContainsFieldValue("credentials", "oauth-secret")) + require.False(t, logSink.ContainsFieldValue("routing_hint", "caller-secret")) +} + +func TestOpenAIWSConnPoolPreferredContinuationIgnoresRoutingHintChanges(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 2 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + account := &Account{ID: 913, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + + acquire := func(t *testing.T, hint, preferred string, forcePreferred bool) *openAIWSConnLease { + t.Helper() + headers := make(http.Header) + if hint != "" { + headers.Set(openAICodexRoutingHintHeader, hint) + } + lease, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + Headers: headers, + PreferredConnID: preferred, + ForcePreferredConn: forcePreferred, + }) + require.NoError(t, err) + require.NotNil(t, lease) + return lease + } + + standard := acquire(t, "model=gpt-5.6-codex", "", false) + connID := standard.ConnID() + standard.Release() + + priority := acquire(t, "model=gpt-5.6-codex;tier=priority", connID, true) + require.True(t, priority.Reused()) + require.Equal(t, connID, priority.ConnID()) + priority.Release() + + standardAgain := acquire(t, "model=gpt-5.6-codex", connID, true) + require.True(t, standardAgain.Reused()) + require.Equal(t, connID, standardAgain.ConnID()) + standardAgain.Release() + + require.Equal(t, 1, dialer.DialCount(), "routing hint is dial-time advisory, not continuation compatibility") +} + +func TestOpenAIWSConnPoolUsesRoutingHintAsSoftDialAffinity(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 4 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 4 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + account := &Account{ID: 913, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + + acquire := func(t *testing.T, hint string) *openAIWSConnLease { + headers := make(http.Header) + headers.Set(openAICodexRoutingHintHeader, hint) + lease, err := pool.Acquire(context.Background(), openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + Headers: headers, + }) + require.NoError(t, err) + require.NotNil(t, lease) + return lease + } + + priority := acquire(t, "model=gpt-5.6-codex;tier=priority") + priorityConnID := priority.ConnID() + priority.Release() + + priorityAgain := acquire(t, "model=gpt-5.6-codex;tier=priority") + require.True(t, priorityAgain.Reused()) + require.Equal(t, priorityConnID, priorityAgain.ConnID()) + priorityAgain.Release() + + flex := acquire(t, "model=gpt-5.6-codex;tier=flex") + require.False(t, flex.Reused()) + require.NotEqual(t, priorityConnID, flex.ConnID()) + flex.Release() + + otherModel := acquire(t, "model=gpt-5.5-codex;tier=priority") + require.False(t, otherModel.Reused()) + require.NotEqual(t, priorityConnID, otherModel.ConnID()) + otherModel.Release() + + defaultTier := acquire(t, "model=gpt-5.6-codex") + require.False(t, defaultTier.Reused()) + require.NotEqual(t, priorityConnID, defaultTier.ConnID()) + defaultConnID := defaultTier.ConnID() + defaultTier.Release() + + defaultAgain := acquire(t, "model=gpt-5.6-codex") + require.True(t, defaultAgain.Reused()) + require.Equal(t, defaultConnID, defaultAgain.ConnID()) + defaultAgain.Release() + + require.Equal(t, 4, dialer.DialCount()) +} diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index 400ae714b8..ce7e8f0386 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -617,7 +617,20 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } } - wsHeaders, _, buildHdrErr := s.buildOpenAIWSHeaders(ctx, c, account, token, wsDecision, isCodexCLI, turnState, strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)), firstPayload.promptCacheKey) + firstRoutingFields := gjson.GetManyBytes(firstPayload.payloadRaw, "model", "service_tier") + wsHeaders, _, buildHdrErr := s.buildOpenAIWSHeaders( + ctx, + c, + account, + token, + wsDecision, + isCodexCLI, + turnState, + strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)), + firstPayload.promptCacheKey, + firstRoutingFields[0].String(), + firstRoutingFields[1].String(), + ) if buildHdrErr != nil { return fmt.Errorf("build ws headers: %w", buildHdrErr) } @@ -785,6 +798,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } sendAndRelay := func(turn int, lease *openAIWSConnLease, payload []byte, payloadBytes int, originalModel string, imageBillingModel string, imageSizeTier string, imageInputSize string) (*OpenAIForwardResult, error) { + responseModelObserver := &upstreamResponseModelObserver{} if lease == nil { return nil, errors.New("upstream websocket lease is nil") } @@ -852,6 +866,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } eventType, eventResponseID, _ := parseOpenAIWSEventEnvelope(upstreamMessage) + responseModelObserver.ObserveOpenAI(upstreamMessage, eventType) if responseID == "" && eventResponseID != "" { responseID = eventResponseID } @@ -1026,18 +1041,20 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } imageCount := imageCounter.Count() result := &OpenAIForwardResult{ - RequestID: responseID, - Usage: usage, - Model: originalModel, - UpstreamModel: mappedModel, - ServiceTier: extractOpenAIServiceTierFromBody(payload), - ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(payload, mappedModel, originalModel), payload, mappedModel), - Stream: reqStream, - OpenAIWSMode: true, - UpstreamTerminalEvent: terminalEvent, - ResponseHeaders: lease.HandshakeHeaders(), - Duration: time.Since(turnStart), - FirstTokenMs: firstTokenMs, + RequestID: responseID, + Usage: usage, + Model: originalModel, + UpstreamModel: mappedModel, + UpstreamResponseModel: responseModelObserver.Model(), + UpstreamResponseModelConflict: responseModelObserver.Conflict(), + ServiceTier: extractOpenAIServiceTierFromBody(payload), + ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(payload, mappedModel, originalModel), payload, mappedModel), + Stream: reqStream, + OpenAIWSMode: true, + UpstreamTerminalEvent: terminalEvent, + ResponseHeaders: lease.HandshakeHeaders(), + Duration: time.Since(turnStart), + FirstTokenMs: firstTokenMs, } if replayInput := replayCollector.Items(); len(replayInput) > 0 { result.wsReplayInput = replayInput @@ -1632,16 +1649,30 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( if parseErr != nil { return parseErr } + nextRoutingFields := gjson.GetManyBytes(nextPayload.payloadRaw, "model", "service_tier") if nextPayload.promptCacheKey != "" { // ingress 会话在整个客户端 WS 生命周期内复用同一上游连接; // prompt_cache_key 对握手头的更新仅在未来需要重新建连时生效。 - updatedHeaders, _, updHdrErr := s.buildOpenAIWSHeaders(ctx, c, account, token, wsDecision, isCodexCLI, turnState, strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)), nextPayload.promptCacheKey) + updatedHeaders, _, updHdrErr := s.buildOpenAIWSHeaders( + ctx, + c, + account, + token, + wsDecision, + isCodexCLI, + turnState, + strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)), + nextPayload.promptCacheKey, + nextRoutingFields[0].String(), + nextRoutingFields[1].String(), + ) if updHdrErr != nil { logOpenAIWSModeInfo("ingress_ws_update_headers_failed account_id=%d err=%v", account.ID, updHdrErr) } else { baseAcquireReq.Headers = updatedHeaders } } + setOpenAICodexRoutingHint(baseAcquireReq.Headers, account, nextRoutingFields[0].String(), nextRoutingFields[1].String()) if nextPayload.previousResponseID != "" { expectedPrev := strings.TrimSpace(lastTurnResponseID) chainedFromLast := expectedPrev != "" && nextPayload.previousResponseID == expectedPrev diff --git a/backend/internal/service/openai_ws_forwarder_ingress_test.go b/backend/internal/service/openai_ws_forwarder_ingress_test.go index 7753ea9598..a13248e393 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress_test.go +++ b/backend/internal/service/openai_ws_forwarder_ingress_test.go @@ -551,13 +551,22 @@ func TestNormalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(t *testing.T) t.Parallel() normalized, err := normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID( - []byte(`{"model":"gpt-5.1","input":[1],"previous_response_id":"resp_x","metadata":{"b":2,"a":1}}`), + []byte(`{"model":"gpt-5.1","input":[1],"previous_response_id":"resp_x","client_metadata":{"request_start_ms":"1"},"stream_options":{"include_usage":true},"generate":false,"metadata":{"b":2,"a":1}}`), ) require.NoError(t, err) require.False(t, gjson.GetBytes(normalized, "input").Exists()) require.False(t, gjson.GetBytes(normalized, "previous_response_id").Exists()) + require.False(t, gjson.GetBytes(normalized, "client_metadata").Exists()) + require.False(t, gjson.GetBytes(normalized, "stream_options").Exists()) + require.False(t, gjson.GetBytes(normalized, "generate").Exists()) require.Equal(t, float64(1), gjson.GetBytes(normalized, "metadata.a").Float()) + normalized, err = normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID( + []byte(`{"model":"gpt-5.1","generate":true}`), + ) + require.NoError(t, err) + require.True(t, gjson.GetBytes(normalized, "generate").Bool()) + _, err = normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(nil) require.Error(t, err) @@ -662,6 +671,36 @@ func TestShouldKeepIngressPreviousResponseID(t *testing.T) { require.Equal(t, "strict_incremental_ok", reason) }) + t.Run("codex_prewarm_to_business_keep", func(t *testing.T) { + prewarmPayload := []byte(`{ + "type":"response.create", + "model":"gpt-5.1", + "store":false, + "generate":false, + "client_metadata":{"x-codex-ws-stream-request-start-ms":"100"}, + "stream_options":{"include_usage":true}, + "input":[{"type":"input_text","text":"hello"}] + }`) + businessPayload := []byte(`{ + "type":"response.create", + "model":"gpt-5.1", + "store":false, + "client_metadata":{"x-codex-ws-stream-request-start-ms":"200"}, + "previous_response_id":"resp_prewarm", + "input":[{"type":"input_text","text":"hello"}] + }`) + + keep, reason, err := shouldKeepIngressPreviousResponseID( + prewarmPayload, + businessPayload, + "resp_prewarm", + false, + ) + require.NoError(t, err) + require.True(t, keep) + require.Equal(t, "strict_incremental_ok", reason) + }) + t.Run("missing_previous_response_id", func(t *testing.T) { payload := []byte(`{"type":"response.create","model":"gpt-5.1","input":[]}`) keep, reason, err := shouldKeepIngressPreviousResponseID(previousPayload, payload, "resp_turn_1", false) diff --git a/backend/internal/service/openai_ws_forwarder_payload.go b/backend/internal/service/openai_ws_forwarder_payload.go index a4418527cb..c53d9e8241 100644 --- a/backend/internal/service/openai_ws_forwarder_payload.go +++ b/backend/internal/service/openai_ws_forwarder_payload.go @@ -75,6 +75,8 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders( turnState string, turnMetadata string, promptCacheKey string, + routingModel string, + routingServiceTier string, ) (http.Header, openAIWSSessionHeaderResolution, error) { headers := make(http.Header) if account == nil || !account.IsOpenAIAgentIdentity() { @@ -157,6 +159,16 @@ func (s *OpenAIGatewayService) buildOpenAIWSHeaders( // 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)。 // 覆盖所有 WS 模式(ctx_pool/dedicated/passthrough)的握手头。 account.ApplyHeaderOverrides(headers) + setOpenAICodexRoutingHint(headers, account, routingModel, routingServiceTier) + logOpenAIRoutingDiagnostics( + ctx, + account, + string(decision.Transport), + routingModel, + routingServiceTier, + strings.TrimSpace(headers.Get(openAICodexRoutingHintHeader)) != "", + "soft_routing_hint", + ) return headers, sessionResolution, nil } @@ -449,6 +461,17 @@ func normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(payload []byte) ( } delete(decoded, "input") delete(decoded, "previous_response_id") + // Codex changes transport-only metadata for every response.create. These fields + // do not alter the context referenced by previous_response_id and are excluded + // from Codex's own websocket reuse comparison. + delete(decoded, "client_metadata") + delete(decoded, "stream_options") + // Official Codex prewarms a connection with generate=false, then omits the + // field on the business request that continues from the prewarm response. + // Only normalize false so a meaningful generate=true change remains visible. + if generate, ok := decoded["generate"].(bool); ok && !generate { + delete(decoded, "generate") + } return json.Marshal(decoded) } diff --git a/backend/internal/service/openai_ws_forwarder_success_test.go b/backend/internal/service/openai_ws_forwarder_success_test.go index 04bd356494..d8270dfe40 100644 --- a/backend/internal/service/openai_ws_forwarder_success_test.go +++ b/backend/internal/service/openai_ws_forwarder_success_test.go @@ -414,6 +414,8 @@ func TestOpenAIGatewayService_BuildOpenAIWSHeadersPreservesCodexIdentity(t *test "", "", "", + "", + "", ) require.NoError(t, err) diff --git a/backend/internal/service/openai_ws_forwarder_v2.go b/backend/internal/service/openai_ws_forwarder_v2.go index 62aee29c20..e0dbd58bd4 100644 --- a/backend/internal/service/openai_ws_forwarder_v2.go +++ b/backend/internal/service/openai_ws_forwarder_v2.go @@ -36,6 +36,7 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( if s == nil || account == nil { return nil, wrapOpenAIWSFallback("invalid_state", errors.New("service or account is nil")) } + responseModelObserver := &upstreamResponseModelObserver{} wsURL, err := s.buildOpenAIResponsesWSURL(account) if err != nil { @@ -136,7 +137,19 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( storeDisabledConnMode := s.openAIWSStoreDisabledConnMode() forceNewConnByPolicy := shouldForceNewConnOnStoreDisabled(storeDisabledConnMode, lastFailureReason) forceNewConn := forceNewConnByPolicy && storeDisabled && previousResponseID == "" && sessionHash != "" && preferredConnID == "" - wsHeaders, sessionResolution, buildHdrErr := s.buildOpenAIWSHeaders(ctx, c, account, token, decision, isCodexCLI, turnState, turnMetadata, promptCacheKey) + wsHeaders, sessionResolution, buildHdrErr := s.buildOpenAIWSHeaders( + ctx, + c, + account, + token, + decision, + isCodexCLI, + turnState, + turnMetadata, + promptCacheKey, + openAIWSPayloadString(payload, "model"), + openAIWSPayloadString(payload, "service_tier"), + ) if buildHdrErr != nil { return nil, fmt.Errorf("build ws headers: %w", buildHdrErr) } @@ -512,6 +525,7 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( if eventType == "" { continue } + responseModelObserver.ObserveOpenAI(message, eventType) eventCount++ if firstEventType == "" { firstEventType = eventType @@ -748,20 +762,22 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2( ) return &OpenAIForwardResult{ - RequestID: responseID, - Usage: *usage, - Model: originalModel, - UpstreamModel: mappedModel, - ImageCount: imageCounter.Count(), - ImageOutputSizes: imageCounter.Sizes(), - ServiceTier: extractOpenAIServiceTier(reqBody), - ReasoningEffort: extractOpenAIReasoningEffort(reqBody, mappedModel, originalModel), - Stream: reqStream, - OpenAIWSMode: true, - UpstreamTerminalEvent: upstreamTerminalEvent, - ResponseHeaders: lease.HandshakeHeaders(), - Duration: time.Since(startTime), - FirstTokenMs: firstTokenMs, + RequestID: responseID, + Usage: *usage, + Model: originalModel, + UpstreamModel: mappedModel, + UpstreamResponseModel: responseModelObserver.Model(), + UpstreamResponseModelConflict: responseModelObserver.Conflict(), + ImageCount: imageCounter.Count(), + ImageOutputSizes: imageCounter.Sizes(), + ServiceTier: extractOpenAIServiceTier(reqBody), + ReasoningEffort: extractOpenAIReasoningEffort(reqBody, mappedModel, originalModel), + Stream: reqStream, + OpenAIWSMode: true, + UpstreamTerminalEvent: upstreamTerminalEvent, + ResponseHeaders: lease.HandshakeHeaders(), + Duration: time.Since(startTime), + FirstTokenMs: firstTokenMs, }, nil } diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index b2c7c67f64..ac6a24e78a 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -180,6 +180,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( if writeClientMessage == nil { return nil, errors.New("client websocket writer is nil") } + responseModelObserver := &upstreamResponseModelObserver{} body, err := prepareOpenAIWSHTTPBridgeBody(payload) if err != nil { @@ -207,7 +208,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( releaseUpstreamCtx() return nil, fmt.Errorf("apply grok Free function-tool cache route: %w", err) } - upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token, grokCacheIdentity, s.cfg) + upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token, grokCacheIdentity, s.cfg, s.settingService) } else { upstreamReq, err = s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token) } @@ -299,18 +300,20 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( resultWithUsage := func() *OpenAIForwardResult { imageCount := imageCounter.Count() result := &OpenAIForwardResult{ - RequestID: responseID, - Usage: usage, - Model: originalModel, - UpstreamModel: mappedModel, - ServiceTier: extractOpenAIServiceTierFromBody(body), - ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(body, mappedModel, originalModel), body, mappedModel), - Stream: reqStream, - OpenAIWSMode: true, - UpstreamTerminalEvent: upstreamTerminalEvent, - ResponseHeaders: cloneHeader(resp.Header), - Duration: time.Since(turnStart), - FirstTokenMs: firstTokenMs, + RequestID: responseID, + Usage: usage, + Model: originalModel, + UpstreamModel: mappedModel, + UpstreamResponseModel: responseModelObserver.Model(), + UpstreamResponseModelConflict: responseModelObserver.Conflict(), + ServiceTier: extractOpenAIServiceTierFromBody(body), + ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(body, mappedModel, originalModel), body, mappedModel), + Stream: reqStream, + OpenAIWSMode: true, + UpstreamTerminalEvent: upstreamTerminalEvent, + ResponseHeaders: cloneHeader(resp.Header), + Duration: time.Since(turnStart), + FirstTokenMs: firstTokenMs, } if replayInput := replayCollector.Items(); len(replayInput) > 0 { result.wsReplayInput = replayInput @@ -355,6 +358,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( upstreamMessage = normalized } eventType, eventResponseID, _ := parseOpenAIWSEventEnvelope(upstreamMessage) + responseModelObserver.ObserveOpenAI(upstreamMessage, eventType) if responseID == "" && eventResponseID != "" { responseID = eventResponseID } @@ -421,8 +425,18 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( upstreamEventErr = errors.New(errMessage) } + // 客户端写出副本改写容量降载码:Codex 对 error/response.failed 中的 + // server_is_overloaded / slow_down 判致命并终止会话,改写后走客户端内置 + // 重试。账号状态与终止事件判定(下方 handleOpenAIWSTerminalTransientFailure) + // 仍使用未改写的 upstreamMessage。 + clientMessage := upstreamMessage + if eventType == "error" || eventType == "response.failed" { + if rewritten, changed := sanitizeOpenAICapacityShedErrorCodeForClient(clientMessage); changed { + clientMessage = rewritten + } + } if !clientDisconnected { - if err := writeClientMessage(upstreamMessage); err != nil { + if err := writeClientMessage(clientMessage); err != nil { if isOpenAIWSClientDisconnectError(err) { clientDisconnected = true closeStatus, closeReason := summarizeOpenAIWSReadCloseError(err) diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 802146cd3f..1e352f012e 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -220,6 +220,69 @@ func TestProxyOpenAIWSHTTPBridgeTurnSSEErrorFailoverSafety(t *testing.T) { } } +// 桥接转发 error / response.failed 给 WS 客户端前必须把容量降载码改写为可重试 +// 的 server_error:Codex 对 server_is_overloaded/slow_down 判致命并终止会话。 +// 账号状态判定使用改写前的原始事件,不受影响。 +func TestProxyOpenAIWSHTTPBridgeTurnRewritesCapacityShedCodeForClient(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + turn int + body string + wantErr bool + }{ + { + name: "turn2_error_frame", + turn: 2, + body: "data: {\"type\":\"error\",\"error\":{\"type\":\"service_unavailable_error\",\"code\":\"server_is_overloaded\",\"message\":\"Our servers are currently overloaded. Please try again later.\"}}\n\n", + wantErr: true, + }, + { + // response.failed 不走 error 事件分支:即便 turn 1 也会被当终止事件 + // 原样转发(不 failover),因此改写必须在这里同样生效。 + name: "turn1_bare_response_failed", + turn: 1, + body: "data: {\"type\":\"response.failed\",\"response\":{\"id\":\"resp_shed\",\"status\":\"failed\",\"error\":{\"code\":\"server_is_overloaded\",\"message\":\"Our servers are currently overloaded. Please try again later.\"}}}\n\n", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(tt.body)), + }} + svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ID: 11, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1} + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) + var writes [][]byte + + _, err := svc.proxyOpenAIWSHTTPBridgeTurn( + context.Background(), c, account, "sk-test", payload, len(payload), + "gpt-5", "", "", "", "", tt.turn, + func(message []byte) error { + writes = append(writes, append([]byte(nil), message...)) + return nil + }, + ) + + if tt.wantErr { + require.Error(t, err) + } else { + require.NoError(t, err) + } + require.Len(t, writes, 1) + require.Contains(t, string(writes[0]), `"code":"server_error"`) + require.NotContains(t, string(writes[0]), "server_is_overloaded") + require.Contains(t, string(writes[0]), "Our servers are currently overloaded") + }) + } +} + func TestProxyOpenAIWSHTTPBridgeTurnRequiresTerminalEvent(t *testing.T) { gin.SetMode(gin.TestMode) @@ -671,7 +734,7 @@ func TestProxyResponsesWebSocketFromClientForGrokUsesXAIHTTPBridgeAndPreservesMa require.Len(t, upstream.bodies, 3) require.Equal(t, xai.DefaultCLIBaseURL+"/responses", upstream.lastReq.URL.String()) require.Equal(t, "Bearer access-token", upstream.lastReq.Header.Get("Authorization")) - require.Equal(t, "sub2api-grok/1.0", upstream.lastReq.Header.Get("User-Agent")) + require.Equal(t, xai.CLIUserAgent(xai.CLIClientVersion), upstream.lastReq.Header.Get("User-Agent")) require.Equal(t, grokCLIVersion, upstream.lastReq.Header.Get("X-Grok-Client-Version")) require.Equal(t, "grok-4.5", gjson.GetBytes(upstream.bodies[0], "model").String()) require.Equal(t, "grok-4.3", gjson.GetBytes(upstream.bodies[1], "model").String()) diff --git a/backend/internal/service/openai_ws_pool.go b/backend/internal/service/openai_ws_pool.go index 8579cc58d6..048031d640 100644 --- a/backend/internal/service/openai_ws_pool.go +++ b/backend/internal/service/openai_ws_pool.go @@ -76,6 +76,10 @@ type openAIWSAcquireRequest struct { ForcePreferredConn bool } +type openAIWSHandshakeCompatibilityKey struct { + betaFeatures string +} + type openAIWSConnLease struct { pool *openAIWSConnPool accountID int64 @@ -240,8 +244,9 @@ type openAIWSConn struct { id string ws openAIWSClientConn - handshakeHeaders http.Header - betaFeatures string + handshakeHeaders http.Header + handshakeCompatibility openAIWSHandshakeCompatibilityKey + routingAffinity string leaseCh chan struct{} closedCh chan struct{} @@ -305,6 +310,14 @@ func (c *openAIWSConn) acquire(ctx context.Context) error { case <-c.closedCh: return errOpenAIWSConnClosed case <-c.leaseCh: + // A cancellation and a lease delivery can become ready together. Once + // the semaphore token has been consumed, check the context again and + // return it before reporting cancellation so a canceled waiter cannot + // strand a pooled connection. + if err := ctx.Err(); err != nil { + c.release() + return err + } select { case <-c.closedCh: c.release() @@ -525,8 +538,12 @@ func (c *openAIWSConn) handshakeHeader(name string) string { return strings.TrimSpace(c.handshakeHeaders.Get(strings.TrimSpace(name))) } -func (c *openAIWSConn) matchesBetaFeatures(betaFeatures string) bool { - return c != nil && c.betaFeatures == betaFeatures +func (c *openAIWSConn) matchesHandshakeCompatibility(compatibility openAIWSHandshakeCompatibilityKey) bool { + return c != nil && c.handshakeCompatibility == compatibility +} + +func (c *openAIWSConn) matchesRoutingAffinity(routingAffinity string) bool { + return c != nil && c.routingAffinity == routingAffinity } func (c *openAIWSConn) isPrewarmed() bool { @@ -838,7 +855,8 @@ func (p *openAIWSConnPool) acquire(ctx context.Context, req openAIWSAcquireReque retryAcquire: accountID := req.Account.ID - betaFeatures := normalizeOpenAIWSBetaFeatures(req.Headers) + compatibility := normalizeOpenAIWSHandshakeCompatibility(req.Headers) + routingAffinity := normalizeOpenAIWSRoutingAffinity(req.Headers) effectiveMaxConns := p.effectiveMaxConnsByAccount(req.Account) if effectiveMaxConns <= 0 { return nil, errOpenAIWSConnQueueFull @@ -846,7 +864,7 @@ retryAcquire: var evicted []*openAIWSConn ap := p.getOrCreateAccountPool(accountID) ap.mu.Lock() - ap.lastAcquire = cloneOpenAIWSAcquireRequestPtr(&req) + acquireGeneration := ap.generation now := time.Now() if ap.lastCleanupAt.IsZero() || now.Sub(ap.lastCleanupAt) >= openAIWSAcquireCleanupInterval { evicted = p.cleanupAccountLocked(ap, now, effectiveMaxConns) @@ -866,7 +884,7 @@ retryAcquire: return nil, errOpenAIWSPreferredConnUnavailable } preferredConn, ok := ap.conns[preferredConnID] - if !ok || !preferredConn.matchesBetaFeatures(betaFeatures) { + if !ok || !preferredConn.matchesHandshakeCompatibility(compatibility) { p.recordConnPickDuration(time.Since(pickStartedAt)) ap.mu.Unlock() closeOpenAIWSConns(evicted) @@ -895,6 +913,7 @@ retryAcquire: reused: true, } p.metrics.acquireReuseTotal.Add(1) + p.recordLastSuccessfulAcquire(accountID, acquireGeneration, req) p.ensureTargetIdleAsync(accountID) return lease, nil } @@ -942,12 +961,13 @@ retryAcquire: reused: true, } p.metrics.acquireReuseTotal.Add(1) + p.recordLastSuccessfulAcquire(accountID, acquireGeneration, req) p.ensureTargetIdleAsync(accountID) return lease, nil } if preferredConnID != "" { - if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) && conn.tryAcquire() { + if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesHandshakeCompatibility(compatibility) && conn.tryAcquire() { connPick := time.Since(pickStartedAt) p.recordConnPickDuration(connPick) ap.mu.Unlock() @@ -964,12 +984,16 @@ retryAcquire: } lease := &openAIWSConnLease{pool: p, accountID: accountID, conn: conn, connPick: connPick, reused: true} p.metrics.acquireReuseTotal.Add(1) + p.recordLastSuccessfulAcquire(accountID, acquireGeneration, req) p.ensureTargetIdleAsync(accountID) return lease, nil } } - best := p.pickLeastBusyConnLocked(ap, "", betaFeatures) + // A routing hint is advisory at WebSocket dial time. Prefer a pooled + // connection whose handshake used the same hint, but do not make that + // preference a continuation compatibility requirement. + best := p.pickLeastBusyConnWithRoutingAffinityLocked(ap, compatibility, routingAffinity) if best != nil && best.tryAcquire() { connPick := time.Since(pickStartedAt) p.recordConnPickDuration(connPick) @@ -987,11 +1011,12 @@ retryAcquire: } lease := &openAIWSConnLease{pool: p, accountID: accountID, conn: best, connPick: connPick, reused: true} p.metrics.acquireReuseTotal.Add(1) + p.recordLastSuccessfulAcquire(accountID, acquireGeneration, req) p.ensureTargetIdleAsync(accountID) return lease, nil } for _, conn := range ap.conns { - if conn == nil || conn == best || !conn.matchesBetaFeatures(betaFeatures) { + if conn == nil || conn == best || !conn.matchesHandshakeCompatibility(compatibility) || !conn.matchesRoutingAffinity(routingAffinity) { continue } if conn.tryAcquire() { @@ -1011,6 +1036,7 @@ retryAcquire: } lease := &openAIWSConnLease{pool: p, accountID: accountID, conn: conn, connPick: connPick, reused: true} p.metrics.acquireReuseTotal.Add(1) + p.recordLastSuccessfulAcquire(accountID, acquireGeneration, req) p.ensureTargetIdleAsync(accountID) return lease, nil } @@ -1018,12 +1044,18 @@ retryAcquire: } if !req.ForceNewConn && len(ap.conns)+ap.creating >= effectiveMaxConns { - compatible := p.pickLeastBusyConnLocked(ap, "", betaFeatures) - if idle := p.pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap, betaFeatures); idle != nil { + affine := p.pickLeastBusyConnWithRoutingAffinityLocked(ap, compatibility, routingAffinity) + if idle := p.pickOldestIdleConnWithoutHandshakeCompatibilityOrRoutingAffinityLocked(ap, compatibility, routingAffinity); idle != nil { delete(ap.conns, idle.id) evicted = append(evicted, idle) p.metrics.scaleDownTotal.Add(1) - } else if compatible == nil { + } else if affine == nil { + compatible := p.pickLeastBusyConnLocked(ap, "", compatibility) + if compatible != nil { + // Capacity is full and every compatible connection is busy. The + // hint remains soft here: queue on a compatible connection below. + goto acquireAtCapacity + } hasConnection := false for _, conn := range ap.conns { if conn != nil { @@ -1068,6 +1100,17 @@ retryAcquire: ap = p.getOrCreateAccountPool(accountID) ap.mu.Lock() ap.creating-- + if ap.generation != acquireGeneration { + ap.signalChangedLocked() + ap.mu.Unlock() + if conn != nil { + conn.close() + } + if retry < 1 { + return p.acquire(ctx, req, retry+1) + } + return nil, errOpenAIWSConnClosed + } if dialErr != nil { ap.prewarmFails++ ap.prewarmFailAt = time.Now() @@ -1075,20 +1118,26 @@ retryAcquire: ap.mu.Unlock() return nil, dialErr } + // Claim the freshly dialed connection before publishing it. Otherwise a + // topology waiter awakened below can take the free semaphore first and + // make the caller that paid for the dial queue behind it. + if !conn.tryAcquire() { + ap.signalChangedLocked() + ap.mu.Unlock() + conn.close() + return nil, errOpenAIWSConnClosed + } ap.conns[conn.id] = conn ap.prewarmFails = 0 ap.prewarmFailAt = time.Time{} + // Wake acquires that observed creating>0 with no compatible connection. + // Without this signal they can remain asleep until the new lease is + // released, even though the pool topology already changed. + ap.signalChangedLocked() ap.mu.Unlock() p.metrics.acquireCreateTotal.Add(1) - - if !conn.tryAcquire() { - if err := conn.acquire(ctx); err != nil { - conn.close() - p.evictConn(accountID, conn.id) - return nil, err - } - } lease := &openAIWSConnLease{pool: p, accountID: accountID, conn: conn, connPick: connPick} + p.recordLastSuccessfulAcquire(accountID, acquireGeneration, req) p.ensureTargetIdleAsync(accountID) return lease, nil } @@ -1100,7 +1149,8 @@ retryAcquire: return nil, errOpenAIWSConnQueueFull } - target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID, betaFeatures) +acquireAtCapacity: + target := p.pickLeastBusyConnLocked(ap, req.PreferredConnID, compatibility) connPick := time.Since(pickStartedAt) p.recordConnPickDuration(connPick) if target == nil { @@ -1142,6 +1192,7 @@ retryAcquire: p.metrics.acquireQueueWaitMs.Add(queueWait.Milliseconds()) lease := &openAIWSConnLease{pool: p, accountID: accountID, conn: target, queueWait: queueWait, connPick: connPick, reused: true} p.metrics.acquireReuseTotal.Add(1) + p.recordLastSuccessfulAcquire(accountID, acquireGeneration, req) p.ensureTargetIdleAsync(accountID) return lease, nil } @@ -1157,6 +1208,23 @@ func (p *openAIWSConnPool) recordConnPickDuration(duration time.Duration) { p.metrics.connPickMs.Add(duration.Milliseconds()) } +func (p *openAIWSConnPool) recordLastSuccessfulAcquire(accountID int64, generation uint64, req openAIWSAcquireRequest) { + if p == nil || accountID <= 0 { + return + } + ap, ok := p.getAccountPool(accountID) + if !ok || ap == nil { + return + } + ap.mu.Lock() + if ap.generation != generation { + ap.mu.Unlock() + return + } + ap.lastAcquire = cloneOpenAIWSAcquireRequestPtr(&req) + ap.mu.Unlock() +} + func (p *openAIWSConnPool) pickOldestIdleConnLocked(ap *openAIWSAccountPool) *openAIWSConn { if ap == nil || len(ap.conns) == 0 { return nil @@ -1173,13 +1241,19 @@ func (p *openAIWSConnPool) pickOldestIdleConnLocked(ap *openAIWSAccountPool) *op return oldest } -func (p *openAIWSConnPool) pickOldestIdleConnWithDifferentBetaFeaturesLocked(ap *openAIWSAccountPool, betaFeatures string) *openAIWSConn { +func (p *openAIWSConnPool) pickOldestIdleConnWithoutHandshakeCompatibilityOrRoutingAffinityLocked( + ap *openAIWSAccountPool, + compatibility openAIWSHandshakeCompatibilityKey, + routingAffinity string, +) *openAIWSConn { if ap == nil || len(ap.conns) == 0 { return nil } var oldest *openAIWSConn for _, conn := range ap.conns { - if conn == nil || conn.matchesBetaFeatures(betaFeatures) || conn.isLeased() || conn.waiters.Load() > 0 || p.isConnPinnedLocked(ap, conn.id) { + if conn == nil || + (conn.matchesHandshakeCompatibility(compatibility) && conn.matchesRoutingAffinity(routingAffinity)) || + conn.isLeased() || conn.waiters.Load() > 0 || p.isConnPinnedLocked(ap, conn.id) { continue } if oldest == nil || conn.lastUsedAt().Before(oldest.lastUsedAt()) { @@ -1330,13 +1404,17 @@ func (p *openAIWSConnPool) cleanupAccountLocked(ap *openAIWSAccountPool, now tim return evicted } -func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, preferredConnID, betaFeatures string) *openAIWSConn { +func (p *openAIWSConnPool) pickLeastBusyConnLocked( + ap *openAIWSAccountPool, + preferredConnID string, + compatibility openAIWSHandshakeCompatibilityKey, +) *openAIWSConn { if ap == nil || len(ap.conns) == 0 { return nil } preferredConnID = stringsTrim(preferredConnID) if preferredConnID != "" { - if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesBetaFeatures(betaFeatures) { + if conn, ok := ap.conns[preferredConnID]; ok && conn.matchesHandshakeCompatibility(compatibility) { return conn } } @@ -1344,7 +1422,37 @@ func (p *openAIWSConnPool) pickLeastBusyConnLocked(ap *openAIWSAccountPool, pref var bestWaiters int32 var bestLastUsed time.Time for _, conn := range ap.conns { - if conn == nil || !conn.matchesBetaFeatures(betaFeatures) { + if conn == nil || !conn.matchesHandshakeCompatibility(compatibility) { + continue + } + waiters := conn.waiters.Load() + lastUsed := conn.lastUsedAt() + if best == nil || + waiters < bestWaiters || + (waiters == bestWaiters && lastUsed.Before(bestLastUsed)) { + best = conn + bestWaiters = waiters + bestLastUsed = lastUsed + } + } + return best +} + +func (p *openAIWSConnPool) pickLeastBusyConnWithRoutingAffinityLocked( + ap *openAIWSAccountPool, + compatibility openAIWSHandshakeCompatibilityKey, + routingAffinity string, +) *openAIWSConn { + if ap == nil || len(ap.conns) == 0 { + return nil + } + var best *openAIWSConn + var bestWaiters int32 + var bestLastUsed time.Time + for _, conn := range ap.conns { + if conn == nil || + !conn.matchesHandshakeCompatibility(compatibility) || + !conn.matchesRoutingAffinity(routingAffinity) { continue } waiters := conn.waiters.Load() @@ -1488,12 +1596,20 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ if len(generations) > 0 { generation = generations[0] } + staleTarget := false defer func() { if ap, ok := p.getAccountPool(accountID); ok && ap != nil { ap.mu.Lock() ap.prewarmActive = false + ap.signalChangedLocked() ap.mu.Unlock() } + if staleTarget { + // A newer acquire arrived while the old dial was in flight. Re-run + // target selection only after clearing prewarmActive so the latest + // beta/hint target can fill the idle budget. + p.ensureTargetIdleAsync(accountID) + } }() for i := 0; i < total; i++ { @@ -1524,6 +1640,13 @@ func (p *openAIWSConnPool) prewarmConns(accountID int64, req openAIWSAcquireRequ conn.close() continue } + if !sameOpenAIWSPrewarmTarget(req, *ap.lastAcquire) { + staleTarget = true + ap.signalChangedLocked() + ap.mu.Unlock() + conn.close() + continue + } if len(ap.conns) >= p.effectiveMaxConnsByAccount(req.Account) { ap.signalChangedLocked() ap.mu.Unlock() @@ -1563,6 +1686,7 @@ func (p *openAIWSConnPool) ClearAccount(accountID int64) { ap.prewarmUntil = time.Time{} ap.prewarmFails = 0 ap.prewarmFailAt = time.Time{} + ap.signalChangedLocked() ap.mu.Unlock() closeOpenAIWSConns(conns) } @@ -1676,7 +1800,8 @@ func (p *openAIWSConnPool) dialConn(ctx context.Context, req openAIWSAcquireRequ } id := p.nextConnID(req.Account.ID) pooledConn := newOpenAIWSConn(id, req.Account.ID, conn, handshakeHeaders) - pooledConn.betaFeatures = normalizeOpenAIWSBetaFeatures(req.Headers) + pooledConn.handshakeCompatibility = normalizeOpenAIWSHandshakeCompatibility(req.Headers) + pooledConn.routingAffinity = normalizeOpenAIWSRoutingAffinity(req.Headers) return pooledConn, nil } @@ -1855,6 +1980,12 @@ func cloneOpenAIWSAcquireRequestPtr(req *openAIWSAcquireRequest) *openAIWSAcquir return &copied } +func sameOpenAIWSPrewarmTarget(a, b openAIWSAcquireRequest) bool { + return stringsTrim(a.WSURL) == stringsTrim(b.WSURL) && + stringsTrim(a.ProxyURL) == stringsTrim(b.ProxyURL) && + normalizeOpenAIWSHandshakeCompatibility(a.Headers) == normalizeOpenAIWSHandshakeCompatibility(b.Headers) +} + func normalizeOpenAIWSBetaFeatures(headers http.Header) string { features := make(map[string]struct{}) for name, values := range headers { @@ -1880,6 +2011,39 @@ func normalizeOpenAIWSBetaFeatures(headers http.Header) string { return strings.Join(normalized, ",") } +func normalizeOpenAIWSHandshakeCompatibility(headers http.Header) openAIWSHandshakeCompatibilityKey { + return openAIWSHandshakeCompatibilityKey{ + betaFeatures: normalizeOpenAIWSBetaFeatures(headers), + } +} + +func normalizeOpenAIWSRoutingAffinity(headers http.Header) string { + canonicalName := http.CanonicalHeaderKey(openAICodexRoutingHintHeader) + if values, ok := headers[canonicalName]; ok { + for _, value := range values { + if value = strings.TrimSpace(value); value != "" { + return value + } + } + } + + variantNames := make([]string, 0) + for name := range headers { + if name != canonicalName && strings.EqualFold(strings.TrimSpace(name), openAICodexRoutingHintHeader) { + variantNames = append(variantNames, name) + } + } + sort.Strings(variantNames) + for _, name := range variantNames { + for _, value := range headers[name] { + if value = strings.TrimSpace(value); value != "" { + return value + } + } + } + return "" +} + func cloneHeader(src http.Header) http.Header { if src == nil { return nil diff --git a/backend/internal/service/openai_ws_pool_test.go b/backend/internal/service/openai_ws_pool_test.go index 8d339359ee..f04b81d40d 100644 --- a/backend/internal/service/openai_ws_pool_test.go +++ b/backend/internal/service/openai_ws_pool_test.go @@ -61,6 +61,18 @@ func TestOpenAIWSConnPool_AcquireCleanupInterval(t *testing.T) { require.Less(t, openAIWSAcquireCleanupInterval, openAIWSBackgroundSweepTicker) } +func TestNormalizeOpenAIWSRoutingAffinityPrefersCanonicalAndSortsVariants(t *testing.T) { + headers := http.Header{ + "X-CODEX-ROUTING-HINT": []string{" variant-uppercase "}, + "X-Codex-Routing-Hint": []string{" ", " canonical "}, + } + + require.Equal(t, "canonical", normalizeOpenAIWSRoutingAffinity(headers)) + + delete(headers, "X-Codex-Routing-Hint") + require.Equal(t, "variant-uppercase", normalizeOpenAIWSRoutingAffinity(headers)) +} + func TestOpenAIWSConnLease_WriteJSONAndGuards(t *testing.T) { conn := newOpenAIWSConn("lease_write", 1, &openAIWSFakeConn{}, nil) lease := &openAIWSConnLease{conn: conn} @@ -310,6 +322,219 @@ func TestOpenAIWSConnPool_AcquireQueueWaitMetrics(t *testing.T) { require.GreaterOrEqual(t, metrics.ConnPickTotal, int64(1)) } +func TestOpenAIWSConnPool_DialSuccessWakesTopologyWaiterAndCanceledWaiterDoesNotLoseLease(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + cfg.Gateway.OpenAIWS.QueueLimitPerConn = 4 + + pool := newOpenAIWSConnPool(cfg) + dialer := newOpenAIWSFirstDialBlockingCaptureDialer() + pool.setClientDialerForTest(dialer) + account := &Account{ID: 991, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + req := openAIWSAcquireRequest{Account: account, WSURL: "wss://example.com/v1/responses"} + + type result struct { + lease *openAIWSConnLease + err error + } + firstCh := make(chan result, 1) + go func() { + lease, err := pool.Acquire(context.Background(), req) + firstCh <- result{lease: lease, err: err} + }() + <-dialer.firstStarted + + waitCtx, cancelWait := context.WithCancel(context.Background()) + secondCh := make(chan result, 1) + waitReq := req + waitReq.Headers = http.Header{openAICodexRoutingHintHeader: {"model=gpt-5.6-codex;tier=priority"}} + go func() { + lease, err := pool.Acquire(waitCtx, waitReq) + secondCh <- result{lease: lease, err: err} + }() + + close(dialer.releaseFirst) + first := <-firstCh + require.NoError(t, first.err) + require.NotNil(t, first.lease) + + // The second acquire initially waits on the account topology channel while + // the first dial is in flight. Dial success must wake it immediately so it + // can queue on the newly-created (still leased) connection. + require.Eventually(t, func() bool { + ap, ok := pool.getAccountPool(account.ID) + if !ok || ap == nil { + return false + } + ap.mu.Lock() + defer ap.mu.Unlock() + for _, conn := range ap.conns { + if conn != nil && conn.waiters.Load() == 1 { + return true + } + } + return false + }, time.Second, 5*time.Millisecond) + + cancelWait() + second := <-secondCh + require.ErrorIs(t, second.err, context.Canceled) + require.Nil(t, second.lease) + ap, ok := pool.getAccountPool(account.ID) + require.True(t, ok) + ap.mu.Lock() + require.NotNil(t, ap.lastAcquire) + require.Empty(t, normalizeOpenAIWSRoutingAffinity(ap.lastAcquire.Headers), "a canceled acquire must not replace the successful prewarm target") + ap.mu.Unlock() + first.lease.Release() + + third, err := pool.Acquire(context.Background(), req) + require.NoError(t, err) + require.True(t, third.Reused(), "a canceled waiter must not consume the released semaphore token") + require.Equal(t, first.lease.ConnID(), third.ConnID()) + third.Release() +} + +func TestOpenAIWSConnPool_PrewarmHintChangeDoesNotInvalidateHealthyDial(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 1 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + cfg.Gateway.OpenAIWS.DialTimeoutSeconds = 2 + + pool := newOpenAIWSConnPool(cfg) + dialer := newOpenAIWSFirstDialBlockingCaptureDialer() + pool.setClientDialerForTest(dialer) + account := &Account{ID: 992, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + oldHeaders := make(http.Header) + oldHeaders.Set(openAICodexRoutingHintHeader, "model=gpt-5.6-codex") + newHeaders := make(http.Header) + newHeaders.Set(openAICodexRoutingHintHeader, "model=gpt-5.6-codex;tier=priority") + ap := pool.getOrCreateAccountPool(account.ID) + ap.mu.Lock() + ap.lastAcquire = &openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + Headers: oldHeaders, + } + ap.mu.Unlock() + + pool.ensureTargetIdleAsync(account.ID) + <-dialer.firstStarted + + // Simulate a newer priority target arriving while the old model-only + // prewarm dial is in flight. Routing hints are advisory, so this alone must + // not discard an otherwise compatible connection. + ap.mu.Lock() + ap.lastAcquire = &openAIWSAcquireRequest{ + Account: account, + WSURL: "wss://example.com/v1/responses", + Headers: newHeaders, + } + ap.mu.Unlock() + close(dialer.releaseFirst) + + require.Eventually(t, func() bool { + ap.mu.Lock() + defer ap.mu.Unlock() + if ap.prewarmActive || len(ap.conns) != 1 { + return false + } + for _, conn := range ap.conns { + return conn != nil && conn.routingAffinity == "model=gpt-5.6-codex" + } + return false + }, 2*time.Second, 10*time.Millisecond) + require.Equal(t, 1, dialer.DialCount(), "routing-hint-only changes must not turn advisory metadata into hard reconnects") +} + +func TestOpenAIWSConnPool_ClearAccountWakesIncompatibleTopologyWaiter(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + + pool := newOpenAIWSConnPool(cfg) + dialer := &openAIWSCountingDialer{} + pool.setClientDialerForTest(dialer) + account := &Account{ID: 993, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + baseReq := openAIWSAcquireRequest{Account: account, WSURL: "wss://example.com/v1/responses"} + betaAReq := baseReq + betaAReq.Headers = http.Header{"X-Codex-Beta-Features": {"feature_a"}} + betaBReq := baseReq + betaBReq.Headers = http.Header{"X-Codex-Beta-Features": {"feature_b"}} + + busy, err := pool.Acquire(context.Background(), betaAReq) + require.NoError(t, err) + require.Equal(t, 1, dialer.DialCount()) + + type result struct { + lease *openAIWSConnLease + err error + } + resultCh := make(chan result, 1) + waitCtx, cancelWait := context.WithTimeout(context.Background(), time.Second) + defer cancelWait() + go func() { + lease, acquireErr := pool.Acquire(waitCtx, betaBReq) + resultCh <- result{lease: lease, err: acquireErr} + }() + + require.Never(t, func() bool { return dialer.DialCount() > 1 }, 50*time.Millisecond, 5*time.Millisecond) + pool.ClearAccount(account.ID) + + resultB := <-resultCh + require.NoError(t, resultB.err) + require.NotNil(t, resultB.lease) + require.False(t, resultB.lease.Reused()) + require.Equal(t, 2, dialer.DialCount(), "ClearAccount must wake the waiter to redial immediately") + resultB.lease.Release() + busy.Release() +} + +func TestOpenAIWSConnPool_ClearAccountDoesNotReviveInFlightDialGeneration(t *testing.T) { + cfg := &config.Config{} + cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 1 + cfg.Gateway.OpenAIWS.MinIdlePerAccount = 0 + cfg.Gateway.OpenAIWS.MaxIdlePerAccount = 1 + + pool := newOpenAIWSConnPool(cfg) + dialer := newOpenAIWSFirstDialBlockingCaptureDialer() + pool.setClientDialerForTest(dialer) + account := &Account{ID: 994, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + req := openAIWSAcquireRequest{Account: account, WSURL: "wss://example.com/v1/responses"} + + type result struct { + lease *openAIWSConnLease + err error + } + resultCh := make(chan result, 1) + go func() { + lease, err := pool.Acquire(context.Background(), req) + resultCh <- result{lease: lease, err: err} + }() + <-dialer.firstStarted + + pool.ClearAccount(account.ID) + close(dialer.releaseFirst) + got := <-resultCh + require.NoError(t, got.err) + require.NotNil(t, got.lease) + require.Equal(t, 2, dialer.DialCount(), "the pre-clear dial must be discarded and retried in the new generation") + require.True(t, strings.HasSuffix(got.lease.ConnID(), "_2")) + + ap, ok := pool.getAccountPool(account.ID) + require.True(t, ok) + ap.mu.Lock() + require.Equal(t, uint64(1), ap.generation) + require.Len(t, ap.conns, 1) + require.NotNil(t, ap.lastAcquire, "only the post-clear successful acquire may restore the prewarm target") + ap.mu.Unlock() + got.lease.Release() +} + func TestOpenAIWSConnPool_ForceNewConnSkipsReuse(t *testing.T) { cfg := &config.Config{} cfg.Gateway.OpenAIWS.MaxConnsPerAccount = 2 @@ -1468,6 +1693,21 @@ func TestOpenAIWSConn_AdditionalGuardBranches(t *testing.T) { closeOpenAIWSConns([]*openAIWSConn{nil, connOK}) } +func TestOpenAIWSConnPool_CanceledWaiterReturnsDeliveredLease(t *testing.T) { + conn := newOpenAIWSConn("cancelled_delivery", 1, &openAIWSFakeConn{}, nil) + + // Both branches of acquire's select are ready. Before the post-delivery + // cancellation check this intermittently returned nil after consuming the + // only lease token, which made the next pool acquire block forever. + for range 64 { + ctx, cancel := context.WithCancel(context.Background()) + cancel() + require.ErrorIs(t, conn.acquire(ctx), context.Canceled) + require.True(t, conn.tryAcquire(), "a canceled waiter must return a delivered lease token") + conn.release() + } +} + func TestOpenAIWSConnLease_MarkBrokenEvictsConn(t *testing.T) { pool := newOpenAIWSConnPool(&config.Config{}) accountID := int64(5001) @@ -1602,6 +1842,21 @@ type openAIWSCountingDialer struct { dialCount int } +type openAIWSFirstDialBlockingCaptureDialer struct { + mu sync.Mutex + dialCount int + headers []http.Header + firstStarted chan struct{} + releaseFirst chan struct{} +} + +func newOpenAIWSFirstDialBlockingCaptureDialer() *openAIWSFirstDialBlockingCaptureDialer { + return &openAIWSFirstDialBlockingCaptureDialer{ + firstStarted: make(chan struct{}), + releaseFirst: make(chan struct{}), + } +} + type openAIWSAlwaysFailDialer struct { mu sync.Mutex dialCount int @@ -1671,6 +1926,36 @@ func (d *openAIWSCountingDialer) Dial( return &openAIWSFakeConn{}, 0, nil, nil } +func (d *openAIWSFirstDialBlockingCaptureDialer) Dial( + ctx context.Context, + wsURL string, + headers http.Header, + proxyURL string, +) (openAIWSClientConn, int, http.Header, error) { + _ = wsURL + _ = proxyURL + d.mu.Lock() + d.dialCount++ + dialNumber := d.dialCount + d.headers = append(d.headers, cloneHeader(headers)) + d.mu.Unlock() + if dialNumber == 1 { + close(d.firstStarted) + select { + case <-ctx.Done(): + return nil, 0, nil, ctx.Err() + case <-d.releaseFirst: + } + } + return &openAIWSFakeConn{}, 0, nil, nil +} + +func (d *openAIWSFirstDialBlockingCaptureDialer) DialCount() int { + d.mu.Lock() + defer d.mu.Unlock() + return d.dialCount +} + func (d *openAIWSCountingDialer) DialCount() int { d.mu.Lock() defer d.mu.Unlock() diff --git a/backend/internal/service/openai_ws_state_store_test.go b/backend/internal/service/openai_ws_state_store_test.go index 235d42331d..9c98a1ea85 100644 --- a/backend/internal/service/openai_ws_state_store_test.go +++ b/backend/internal/service/openai_ws_state_store_test.go @@ -193,6 +193,20 @@ func (c *openAIWSStateStoreTimeoutProbeCache) DeleteSessionAccountID(ctx context return nil } +func (c *openAIWSStateStoreTimeoutProbeCache) SetGrokVideoPendingBilling(_ context.Context, _ string, _ []byte, _ time.Duration) error { + return nil +} +func (c *openAIWSStateStoreTimeoutProbeCache) GetGrokVideoPendingBilling(_ context.Context, _ string) ([]byte, error) { + return nil, nil +} +func (c *openAIWSStateStoreTimeoutProbeCache) ClaimGrokVideoBilled(_ context.Context, _ string, _ time.Duration) (bool, error) { + return true, nil +} + +func (c *openAIWSStateStoreTimeoutProbeCache) ReleaseGrokVideoBilled(_ context.Context, _ string) error { + return nil +} + func TestOpenAIWSStateStore_RedisOpsUseShortTimeout(t *testing.T) { probe := &openAIWSStateStoreTimeoutProbeCache{} store := NewOpenAIWSStateStore(probe) diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay.go b/backend/internal/service/openai_ws_v2/passthrough_relay.go index b6936fe994..b3fe6d3084 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay.go @@ -31,6 +31,8 @@ type Usage struct { type RelayResult struct { RequestModel string + ResponseModel string + ResponseModelConflict bool Usage Usage RequestID string TerminalEventType string @@ -42,12 +44,14 @@ type RelayResult struct { } type RelayTurnResult struct { - RequestModel string - Usage Usage - RequestID string - TerminalEventType string - Duration time.Duration - FirstTokenMs *int + RequestModel string + ResponseModel string + ResponseModelConflict bool + Usage Usage + RequestID string + TerminalEventType string + Duration time.Duration + FirstTokenMs *int } type RelayExit struct { @@ -90,6 +94,8 @@ type relayState struct { requestModelMu sync.RWMutex requestModel string lastResponseID string + lastResponseModel string + responseConflict bool terminalEventType string firstTokenMs *int turnTimingByID map[string]*relayTurnTiming @@ -104,17 +110,22 @@ type relayExitSignal struct { } type observedUpstreamEvent struct { - terminal bool - eventType string - responseID string - usage Usage - duration time.Duration - firstToken *int + terminal bool + eventType string + responseID string + usage Usage + responseModel string + responseConflict bool + duration time.Duration + firstToken *int } type relayTurnTiming struct { - startAt time.Time - firstTokenMs *int + startAt time.Time + firstTokenMs *int + firstResponseModel string + terminalResponseModel string + responseModelConflict bool } func Relay( @@ -684,15 +695,19 @@ func observeUpstreamMessage( responseID: responseID, usage: parsedUsage, } + var turnTiming *relayTurnTiming if responseID != "" { - turnTiming := openAIWSRelayGetOrInitTurnTiming(state, responseID, now) + turnTiming = openAIWSRelayGetOrInitTurnTiming(state, responseID, now) if turnTiming != nil && turnTiming.firstTokenMs == nil && isTokenEvent(eventType) { ms := int(now.Sub(turnTiming.startAt).Milliseconds()) if ms >= 0 { turnTiming.firstTokenMs = &ms } } + } else { + turnTiming = state.activeTurn } + observeRelayTurnResponseModel(turnTiming, firstRelayResponseModel(message), isTerminalEvent(eventType)) if !isTerminalEvent(eventType) { return observed } @@ -701,6 +716,10 @@ func observeUpstreamMessage( if responseID != "" { state.lastResponseID = responseID if turnTiming, ok := openAIWSRelayDeleteTurnTiming(state, responseID); ok { + observed.responseModel = relayTurnResponseModel(&turnTiming) + observed.responseConflict = turnTiming.responseModelConflict + state.lastResponseModel = observed.responseModel + state.responseConflict = observed.responseConflict duration := now.Sub(turnTiming.startAt) if duration < 0 { duration = 0 @@ -729,15 +748,64 @@ func emitTurnComplete( requestModel = state.currentRequestModel() } onTurnComplete(RelayTurnResult{ - RequestModel: requestModel, - Usage: observed.usage, - RequestID: responseID, - TerminalEventType: observed.eventType, - Duration: observed.duration, - FirstTokenMs: openAIWSRelayCloneIntPtr(observed.firstToken), + RequestModel: requestModel, + ResponseModel: observed.responseModel, + ResponseModelConflict: observed.responseConflict, + Usage: observed.usage, + RequestID: responseID, + TerminalEventType: observed.eventType, + Duration: observed.duration, + FirstTokenMs: openAIWSRelayCloneIntPtr(observed.firstToken), }) } +func firstRelayResponseModel(message []byte) string { + if len(message) == 0 { + return "" + } + values := gjson.GetManyBytes(message, "response.model", "model") + for _, value := range values { + if value.Type != gjson.String { + continue + } + if model := strings.TrimSpace(value.String()); model != "" { + return model + } + } + return "" +} + +func observeRelayTurnResponseModel(turn *relayTurnTiming, model string, terminal bool) { + if turn == nil { + return + } + model = strings.TrimSpace(model) + if model == "" { + return + } + current := relayTurnResponseModel(turn) + if current != "" && !strings.EqualFold(current, model) { + turn.responseModelConflict = true + } + if terminal { + turn.terminalResponseModel = model + return + } + if turn.firstResponseModel == "" { + turn.firstResponseModel = model + } +} + +func relayTurnResponseModel(turn *relayTurnTiming) string { + if turn == nil { + return "" + } + if turn.terminalResponseModel != "" { + return turn.terminalResponseModel + } + return turn.firstResponseModel +} + func openAIWSRelayGetOrInitTurnTiming(state *relayState, responseID string, now time.Time) *relayTurnTiming { if state == nil { return nil @@ -888,6 +956,8 @@ func enrichResult(result *RelayResult, state *relayState, duration time.Duration return } result.RequestModel = state.currentRequestModel() + result.ResponseModel = state.lastResponseModel + result.ResponseModelConflict = state.responseConflict result.Usage = state.usage result.RequestID = state.lastResponseID result.TerminalEventType = state.terminalEventType diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go index 6036be7451..ed403e3448 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay_internal_test.go @@ -457,6 +457,60 @@ func TestRelayTurnTimingHelpersCoverage(t *testing.T) { require.False(t, ok) } +func TestObserveUpstreamMessage_ResponseModelIsTurnLocalAndTerminalWins(t *testing.T) { + t.Parallel() + + state := &relayState{requestModel: "gpt-5.6-sol"} + startAt := time.Unix(0, 0) + now := startAt + nowFn := func() time.Time { + now = now.Add(5 * time.Millisecond) + return now + } + + created := observeUpstreamMessage( + state, + []byte(`{"type":"response.created","response":{"id":"resp_1","model":"gpt-5.5"}}`), + startAt, + nowFn, + nil, + ) + require.False(t, created.terminal) + + completed := observeUpstreamMessage( + state, + []byte(`{"type":"response.completed","response":{"id":"resp_1","model":"gpt-5.4","usage":{"input_tokens":1,"output_tokens":2}}}`), + startAt, + nowFn, + nil, + ) + require.True(t, completed.terminal) + require.Equal(t, "gpt-5.4", completed.responseModel) + require.True(t, completed.responseConflict) + + var firstTurn RelayTurnResult + emitTurnComplete(func(turn RelayTurnResult) { firstTurn = turn }, state, completed) + require.Equal(t, "gpt-5.4", firstTurn.ResponseModel) + require.True(t, firstTurn.ResponseModelConflict) + + observeUpstreamMessage( + state, + []byte(`{"type":"response.created","response":{"id":"resp_2","model":"gpt-5.3"}}`), + startAt, + nowFn, + nil, + ) + second := observeUpstreamMessage( + state, + []byte(`{"type":"response.completed","response":{"id":"resp_2","model":"GPT-5.3","usage":{"input_tokens":3,"output_tokens":4}}}`), + startAt, + nowFn, + nil, + ) + require.Equal(t, "GPT-5.3", second.responseModel) + require.False(t, second.responseConflict, "the previous turn must not contaminate this turn") +} + func TestObserveUpstreamMessage_ResponseIDFallbackPolicy(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter.go b/backend/internal/service/openai_ws_v2_passthrough_adapter.go index ae74027bca..775e26980c 100644 --- a/backend/internal/service/openai_ws_v2_passthrough_adapter.go +++ b/backend/internal/service/openai_ws_v2_passthrough_adapter.go @@ -800,7 +800,19 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( turnState = strings.TrimSpace(c.GetHeader(openAIWSTurnStateHeader)) turnMetadata = strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)) } - headers, _, buildHdrErr := s.buildOpenAIWSHeaders(ctx, c, account, token, wsDecision, isCodexCLI, turnState, turnMetadata, promptCacheKey) + headers, _, buildHdrErr := s.buildOpenAIWSHeaders( + ctx, + c, + account, + token, + wsDecision, + isCodexCLI, + turnState, + turnMetadata, + promptCacheKey, + gjson.GetBytes(firstClientMessage, "model").String(), + gjson.GetBytes(firstClientMessage, "service_tier").String(), + ) if buildHdrErr != nil { return fmt.Errorf("build ws headers: %w", buildHdrErr) } @@ -1103,16 +1115,18 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( CacheReadInputTokens: turn.Usage.CacheReadInputTokens, ImageOutputTokens: turn.Usage.ImageOutputTokens, }, - Model: turnRequestModel, - UpstreamModel: openAIWSDifferentModel(turnRequestModel, turnUpstreamModel), - ServiceTier: usageMeta.serviceTier.Load(), - ReasoningEffort: usageMeta.reasoningEffort.Load(), - Stream: true, - OpenAIWSMode: true, - UpstreamTerminalEvent: normalizeOpenAIWSTerminalEvent(turn.TerminalEventType), - ResponseHeaders: cloneHeader(handshakeHeaders), - Duration: turn.Duration, - FirstTokenMs: turn.FirstTokenMs, + Model: turnRequestModel, + UpstreamModel: openAIWSDifferentModel(turnRequestModel, turnUpstreamModel), + UpstreamResponseModel: turn.ResponseModel, + UpstreamResponseModelConflict: turn.ResponseModelConflict, + ServiceTier: usageMeta.serviceTier.Load(), + ReasoningEffort: usageMeta.reasoningEffort.Load(), + Stream: true, + OpenAIWSMode: true, + UpstreamTerminalEvent: normalizeOpenAIWSTerminalEvent(turn.TerminalEventType), + ResponseHeaders: cloneHeader(handshakeHeaders), + Duration: turn.Duration, + FirstTokenMs: turn.FirstTokenMs, } logOpenAIWSV2Passthrough( "relay_turn_completed account_id=%d turn=%d request_id=%s terminal_event=%s turn_requested_model=%s turn_upstream_model=%s duration_ms=%d first_token_ms=%d input_tokens=%d output_tokens=%d cache_read_tokens=%d", @@ -1222,16 +1236,18 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( CacheReadInputTokens: relayResult.Usage.CacheReadInputTokens, ImageOutputTokens: relayResult.Usage.ImageOutputTokens, }, - Model: resultRequestModel, - UpstreamModel: openAIWSDifferentModel(resultRequestModel, resultUpstreamModel), - ServiceTier: usageMeta.serviceTier.Load(), - ReasoningEffort: usageMeta.reasoningEffort.Load(), - Stream: true, - OpenAIWSMode: true, - UpstreamTerminalEvent: normalizeOpenAIWSTerminalEvent(relayResult.TerminalEventType), - ResponseHeaders: cloneHeader(handshakeHeaders), - Duration: relayResult.Duration, - FirstTokenMs: relayResult.FirstTokenMs, + Model: resultRequestModel, + UpstreamModel: openAIWSDifferentModel(resultRequestModel, resultUpstreamModel), + UpstreamResponseModel: relayResult.ResponseModel, + UpstreamResponseModelConflict: relayResult.ResponseModelConflict, + ServiceTier: usageMeta.serviceTier.Load(), + ReasoningEffort: usageMeta.reasoningEffort.Load(), + Stream: true, + OpenAIWSMode: true, + UpstreamTerminalEvent: normalizeOpenAIWSTerminalEvent(relayResult.TerminalEventType), + ResponseHeaders: cloneHeader(handshakeHeaders), + Duration: relayResult.Duration, + FirstTokenMs: relayResult.FirstTokenMs, } turnCount := int(completedTurns.Load()) diff --git a/backend/internal/service/ops_system_log_sink.go b/backend/internal/service/ops_system_log_sink.go index 1cbac237c4..7a247a290d 100644 --- a/backend/internal/service/ops_system_log_sink.go +++ b/backend/internal/service/ops_system_log_sink.go @@ -34,6 +34,10 @@ type OpsSystemLogSink struct { batchSize int flushInterval time.Duration + // 连续写入失败后的退避参数。构造后只读,测试可在 Start 前覆盖。 + flushBackoff time.Duration + flushBackoffMax time.Duration + ctx context.Context cancel context.CancelFunc wg sync.WaitGroup @@ -48,22 +52,54 @@ type OpsSystemLogSink struct { const maxSystemLogHostLength = 255 +const ( + // 首次写入失败后暂停落库的时长,之后逐次翻倍到上限。 + defaultOpsSystemLogFlushBackoff = 2 * time.Second + // 退避上限。日志是尽力而为的观测数据,不值得为它无限期占用连接池。 + defaultOpsSystemLogFlushBackoffMax = 60 * time.Second +) + func NewOpsSystemLogSink(opsRepo OpsRepository) *OpsSystemLogSink { ctx, cancel := context.WithCancel(context.Background()) rawHost, err := os.Hostname() s := &OpsSystemLogSink{ - opsRepo: opsRepo, - host: normalizeSystemLogHost(rawHost, err), - queue: make(chan *logger.LogEvent, 5000), - batchSize: 200, - flushInterval: time.Second, - ctx: ctx, - cancel: cancel, + opsRepo: opsRepo, + host: normalizeSystemLogHost(rawHost, err), + queue: make(chan *logger.LogEvent, 5000), + batchSize: 200, + flushInterval: time.Second, + flushBackoff: defaultOpsSystemLogFlushBackoff, + flushBackoffMax: defaultOpsSystemLogFlushBackoffMax, + ctx: ctx, + cancel: cancel, } s.lastError.Store("") return s } +// flushBackoffFor 返回第 failures 次连续失败后的退避时长(指数退避,封顶)。 +func (s *OpsSystemLogSink) flushBackoffFor(failures int) time.Duration { + base := s.flushBackoff + if base <= 0 { + base = defaultOpsSystemLogFlushBackoff + } + maxBackoff := s.flushBackoffMax + if maxBackoff <= 0 { + maxBackoff = defaultOpsSystemLogFlushBackoffMax + } + if maxBackoff < base { + maxBackoff = base + } + backoff := base + for i := 1; i < failures && backoff < maxBackoff; i++ { + backoff *= 2 + } + if backoff > maxBackoff { + backoff = maxBackoff + } + return backoff +} + func normalizeSystemLogHost(host string, err error) string { host = strings.TrimSpace(host) if err != nil || host == "" { @@ -146,20 +182,37 @@ func (s *OpsSystemLogSink) run() { defer ticker.Stop() batch := make([]*logger.LogEvent, 0, s.batchSize) + // 仅在本 goroutine 内读写,无需加锁。 + failures := 0 + var suppressedUntil time.Time flush := func(baseCtx context.Context) { if len(batch) == 0 { return } + now := time.Now() + if now.Before(suppressedUntil) { + // 退避窗口内直接丢弃本批:日志是尽力而为的观测数据,继续攒批只会把 + // 压力转移到内存,而每次重试都会再占用并取消一条池内连接。 + atomic.AddUint64(&s.droppedCount, uint64(len(batch))) + batch = batch[:0] + return + } started := time.Now() inserted, err := s.flushBatch(baseCtx, batch) delay := time.Since(started) if err != nil { + failures++ + backoff := s.flushBackoffFor(failures) + suppressedUntil = time.Now().Add(backoff) atomic.AddUint64(&s.writeFailed, uint64(len(batch))) s.lastError.Store(err.Error()) - _, _ = fmt.Fprintf(os.Stderr, "time=%s level=WARN msg=\"ops system log sink flush failed\" err=%v batch=%d\n", - time.Now().Format(time.RFC3339Nano), err, len(batch), + // 每个退避窗口至多一条,避免数据库故障期间刷屏。 + _, _ = fmt.Fprintf(os.Stderr, "time=%s level=WARN msg=\"ops system log sink flush failed\" err=%v batch=%d failures=%d backoff=%s\n", + time.Now().Format(time.RFC3339Nano), err, len(batch), failures, backoff, ) } else { + failures = 0 + suppressedUntil = time.Time{} atomic.AddUint64(&s.writtenCount, uint64(inserted)) atomic.AddUint64(&s.totalDelayNs, uint64(delay.Nanoseconds())) s.lastError.Store("") diff --git a/backend/internal/service/ops_system_log_sink_backoff_test.go b/backend/internal/service/ops_system_log_sink_backoff_test.go new file mode 100644 index 0000000000..3a29cb6d0d --- /dev/null +++ b/backend/internal/service/ops_system_log_sink_backoff_test.go @@ -0,0 +1,225 @@ +package service + +import ( + "context" + "errors" + "sync/atomic" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" +) + +func opsSystemLogBackoffEvent() *logger.LogEvent { + return &logger.LogEvent{ + Time: time.Now().UTC(), + Level: "warn", + Component: "app", + Message: "boom", + Fields: map[string]any{}, + } +} + +// pumpOpsSystemLogEvents 持续投递日志事件,模拟故障期间业务侧不断产生 WARN/ERROR。 +func pumpOpsSystemLogEvents(t *testing.T, sink *OpsSystemLogSink) { + t.Helper() + stop := make(chan struct{}) + done := make(chan struct{}) + go func() { + defer close(done) + for { + select { + case <-stop: + return + default: + } + sink.WriteLogEvent(opsSystemLogBackoffEvent()) + time.Sleep(2 * time.Millisecond) + } + }() + t.Cleanup(func() { + close(stop) + <-done + }) +} + +func waitForOpsSystemLogCondition(t *testing.T, timeout time.Duration, cond func() bool) bool { + t.Helper() + deadline := time.Now().Add(timeout) + for time.Now().Before(deadline) { + if cond() { + return true + } + time.Sleep(5 * time.Millisecond) + } + return cond() +} + +func TestOpsSystemLogSinkFlushBackoffFor(t *testing.T) { + sink := &OpsSystemLogSink{flushBackoff: time.Second, flushBackoffMax: 8 * time.Second} + + cases := []struct { + name string + failures int + want time.Duration + }{ + {"first_failure", 1, time.Second}, + {"second_failure", 2, 2 * time.Second}, + {"third_failure", 3, 4 * time.Second}, + {"fourth_failure", 4, 8 * time.Second}, + {"capped", 9, 8 * time.Second}, + {"large_streak_does_not_overflow", 1000, 8 * time.Second}, + {"zero_streak_uses_base", 0, time.Second}, + {"negative_streak_uses_base", -1, time.Second}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if got := sink.flushBackoffFor(tc.failures); got != tc.want { + t.Fatalf("flushBackoffFor(%d) = %s, want %s", tc.failures, got, tc.want) + } + }) + } +} + +// 未显式配置时回落到默认值;max 小于 base 时以 base 为准,不产生比基准还短的退避。 +func TestOpsSystemLogSinkFlushBackoffForFallbacks(t *testing.T) { + zero := &OpsSystemLogSink{} + if got := zero.flushBackoffFor(1); got != defaultOpsSystemLogFlushBackoff { + t.Fatalf("zero-value base = %s, want %s", got, defaultOpsSystemLogFlushBackoff) + } + if got := zero.flushBackoffFor(100); got != defaultOpsSystemLogFlushBackoffMax { + t.Fatalf("zero-value cap = %s, want %s", got, defaultOpsSystemLogFlushBackoffMax) + } + + inverted := &OpsSystemLogSink{flushBackoff: 5 * time.Second, flushBackoffMax: time.Second} + if got := inverted.flushBackoffFor(3); got != 5*time.Second { + t.Fatalf("inverted bounds = %s, want 5s", got) + } +} + +// issue #5265:写入失败后如果按 flushInterval 继续每秒重试,每一轮都会占用并取消 +// 一条池内连接(远程 PG 上 COPY 取消会让连接协议失步而被销毁),小连接池会被日志 +// 通道长期占满,业务侧最终报 Billing 503。失败后必须退避。 +func TestOpsSystemLogSinkSuppressesRetriesDuringBackoff(t *testing.T) { + var calls int64 + repo := &opsRepoMock{ + BatchInsertSystemLogsFn: func(_ context.Context, _ []*OpsInsertSystemLogInput) (int64, error) { + atomic.AddInt64(&calls, 1) + return 0, errors.New("db unavailable") + }, + } + + sink := NewOpsSystemLogSink(repo) + sink.batchSize = 1 + sink.flushInterval = 5 * time.Millisecond + sink.flushBackoff = 800 * time.Millisecond + sink.flushBackoffMax = 800 * time.Millisecond + sink.Start() + defer sink.Stop() + pumpOpsSystemLogEvents(t, sink) + + if !waitForOpsSystemLogCondition(t, 2*time.Second, func() bool { return atomic.LoadInt64(&calls) >= 1 }) { + t.Fatalf("first flush never happened") + } + + // 退避窗口内不得再打上游:修复前 5ms 的 flushInterval 会在这 300ms 里打出上百次。 + time.Sleep(300 * time.Millisecond) + if got := atomic.LoadInt64(&calls); got != 1 { + t.Fatalf("upstream calls during backoff = %d, want 1", got) + } + + // 被抑制的批次记为 dropped,而不是 write_failed —— 没有尝试过就不算写失败。 + health := sink.Health() + if health.DroppedCount == 0 { + t.Fatalf("dropped_count should grow while flushing is suppressed") + } + if health.WriteFailed == 0 || health.LastError == "" { + t.Fatalf("failed flush should still surface in health: %+v", health) + } +} + +// 退避到期后必须自动恢复,不能变成永久停写。 +func TestOpsSystemLogSinkResumesAfterBackoffWindow(t *testing.T) { + var calls int64 + repo := &opsRepoMock{ + BatchInsertSystemLogsFn: func(_ context.Context, _ []*OpsInsertSystemLogInput) (int64, error) { + atomic.AddInt64(&calls, 1) + return 0, errors.New("db unavailable") + }, + } + + sink := NewOpsSystemLogSink(repo) + sink.batchSize = 1 + sink.flushInterval = 5 * time.Millisecond + sink.flushBackoff = 150 * time.Millisecond + sink.flushBackoffMax = 150 * time.Millisecond + sink.Start() + defer sink.Stop() + pumpOpsSystemLogEvents(t, sink) + + if !waitForOpsSystemLogCondition(t, 3*time.Second, func() bool { return atomic.LoadInt64(&calls) >= 3 }) { + t.Fatalf("sink did not resume flushing after backoff, calls=%d", atomic.LoadInt64(&calls)) + } +} + +// 一次成功必须清空失败计数与抑制窗口,否则短暂抖动后写入速率无法恢复。 +func TestOpsSystemLogSinkSuccessClearsSuppression(t *testing.T) { + var calls int64 + repo := &opsRepoMock{ + BatchInsertSystemLogsFn: func(_ context.Context, inputs []*OpsInsertSystemLogInput) (int64, error) { + if atomic.AddInt64(&calls, 1) == 1 { + return 0, errors.New("db unavailable") + } + return int64(len(inputs)), nil + }, + } + + sink := NewOpsSystemLogSink(repo) + sink.batchSize = 1 + sink.flushInterval = 5 * time.Millisecond + sink.flushBackoff = 150 * time.Millisecond + sink.flushBackoffMax = 150 * time.Millisecond + sink.Start() + defer sink.Stop() + pumpOpsSystemLogEvents(t, sink) + + // 第 2 次调用(恢复后的首次成功)之后不应再有抑制:调用次数需要快速爬升。 + if !waitForOpsSystemLogCondition(t, 3*time.Second, func() bool { return atomic.LoadInt64(&calls) >= 2 }) { + t.Fatalf("sink never retried after the first failure") + } + recovered := atomic.LoadInt64(&calls) + time.Sleep(200 * time.Millisecond) + if got := atomic.LoadInt64(&calls) - recovered; got < 3 { + t.Fatalf("calls after recovery = %d in 200ms, want >=3 (suppression not cleared)", got) + } + if sink.Health().WrittenCount == 0 { + t.Fatalf("written_count should grow after recovery") + } +} + +// 健康路径不受影响:一直成功就永远不进入退避。 +func TestOpsSystemLogSinkHealthyPathNeverSuppressed(t *testing.T) { + var calls int64 + repo := &opsRepoMock{ + BatchInsertSystemLogsFn: func(_ context.Context, inputs []*OpsInsertSystemLogInput) (int64, error) { + atomic.AddInt64(&calls, 1) + return int64(len(inputs)), nil + }, + } + + sink := NewOpsSystemLogSink(repo) + sink.batchSize = 1 + sink.flushInterval = 5 * time.Millisecond + sink.flushBackoff = time.Hour + sink.flushBackoffMax = time.Hour + sink.Start() + defer sink.Stop() + pumpOpsSystemLogEvents(t, sink) + + if !waitForOpsSystemLogCondition(t, 3*time.Second, func() bool { return atomic.LoadInt64(&calls) >= 10 }) { + t.Fatalf("healthy sink should keep flushing, calls=%d", atomic.LoadInt64(&calls)) + } + if got := sink.Health().DroppedCount; got != 0 { + t.Fatalf("healthy sink dropped_count = %d, want 0", got) + } +} diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 266f806e83..d6f4286be1 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -138,6 +138,97 @@ func (s *RateLimitService) notifyAccountSchedulingBlockCleared(accountID int64) s.runtimeBlocker.ClearAccountSchedulingBlock(accountID) } +// ApplyAccountSchedulingThreshold evaluates admin-configured per-platform +// utilization thresholds and, when breached, parks the account as temp- +// unschedulable until the winning window resets. Returns true when the account +// is blocked (either newly or already paused for the same threshold reason). +func (s *RateLimitService) ApplyAccountSchedulingThreshold(ctx context.Context, account *Account) bool { + if s == nil || s.settingService == nil || s.accountRepo == nil || account == nil || account.ID <= 0 { + return false + } + if !account.IsActive() || !account.Schedulable { + return false + } + + now := time.Now().UTC() + thresholds := s.settingService.GetAccountSchedulingThresholds(ctx) + decision := EvaluateAccountSchedulingThreshold(account, thresholds, now) + if !decision.ShouldPause || decision.Until == nil || !decision.Until.After(now) { + return false + } + + reason := BuildDetailedAccountSchedulingThresholdReason(AccountSchedulingThresholdReasonInput{ + Platform: decision.Platform, + Window: decision.Window, + Scope: decision.Scope, + ThresholdPercent: decision.ThresholdPercent, + UsedPercent: decision.UsedPercent, + Until: *decision.Until, + Now: now, + }) + + if accountHasSameSchedulingThresholdPause(account, *decision.Until, reason) { + return true + } + if !account.IsSchedulable() { + return false + } + + account.TempUnschedulableUntil = cloneTimePtr(decision.Until) + account.TempUnschedulableReason = reason + s.notifyAccountSchedulingBlocked(account, *decision.Until, "account_scheduling_threshold") + + if err := s.accountRepo.SetTempUnschedulable(ctx, account.ID, *decision.Until, reason); err != nil { + slog.Warn("account_scheduling_threshold_set_temp_unsched_failed", + "account_id", account.ID, + "platform", decision.Platform, + "window", decision.Window, + "scope", decision.Scope, + "threshold_percent", decision.ThresholdPercent, + "used_percent", decision.UsedPercent, + "until", decision.Until.UTC(), + "error", err) + } else if s.tempUnschedCache != nil { + if state := tempUnschedStateFromStoredReason(reason, decision.Until.Unix()); state != nil { + if err := s.tempUnschedCache.SetTempUnsched(ctx, account.ID, state); err != nil { + slog.Warn("account_scheduling_threshold_cache_set_failed", "account_id", account.ID, "error", err) + } + } + } + + slog.Info("account_scheduling_threshold_temp_unschedulable", + "account_id", account.ID, + "platform", decision.Platform, + "window", decision.Window, + "scope", decision.Scope, + "threshold_percent", decision.ThresholdPercent, + "used_percent", decision.UsedPercent, + "until", decision.Until.UTC()) + return true +} + +func accountHasSameSchedulingThresholdPause(account *Account, until time.Time, reason string) bool { + if account == nil || account.TempUnschedulableUntil == nil { + return false + } + if account.TempUnschedulableUntil.UTC().Unix() != until.UTC().Unix() { + return false + } + + existing, ok := parseTempUnschedReasonPayload(account.TempUnschedulableReason) + if !ok || existing.Source != AccountSchedulingThresholdReasonSource { + return false + } + next, ok := parseTempUnschedReasonPayload(reason) + if !ok || next.Source != AccountSchedulingThresholdReasonSource { + return false + } + + existing.TriggeredAtUnix = 0 + next.TriggeredAtUnix = 0 + return existing == next +} + // ErrorPolicyResult 表示错误策略检查的结果 type ErrorPolicyResult int @@ -2055,6 +2146,9 @@ func (s *RateLimitService) HandleUpstreamModelNotFound(ctx context.Context, acco if modelKey == "" { return false } + if shouldSkipCodexPlanGatedImageModelCooldown(ctx, reason, requestedModel, modelKey) { + return true + } resetAt := time.Now().Add(cooldown) if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, modelKey, resetAt, reason); err != nil { slog.Warn("upstream_model_not_found_set_model_rate_limit_failed", "account_id", account.ID, "model", modelKey, "reason", reason, "error", err) @@ -2064,6 +2158,29 @@ func (s *RateLimitService) HandleUpstreamModelNotFound(ctx context.Context, acco return true } +// shouldSkipCodexPlanGatedImageModelCooldown 判断这次 Codex plan-gated 400 是否 +// 属于"图片模型被文本端点拒绝"。 +// +// 这类错误是确定性的端点错配,不是账号能力缺失:同一账号通过 /v1/images/* 依然 +// 能出图。在这里写 per-model 冷却,会让一次用错端点的请求把整个号池对正确的生图 +// 端点下线(#4828)。 +// +// 但请求本身就从 /v1/images/* 入站时不适用——那种情况下被拒说明账号确实不具备 +// 该模型能力,冷却是必要的刹车:没有它,每个生图请求都会完整走一遍号池,对上游 +// 形成无上界的 400 放大。 +// +// 请求模型与最终冷却键都要判:冷却键走的是 account.GetMappedModel,账号可能把 +// 文本别名映射到 gpt-image-*,只判请求模型会漏掉这种形态。 +func shouldSkipCodexPlanGatedImageModelCooldown(ctx context.Context, reason, requestedModel, modelKey string) bool { + if reason != upstreamCodexPlanGatedModelReason { + return false + } + if OpenAIImagesEndpointFromContext(ctx) { + return false + } + return IsGPTImageGenerationModel(requestedModel) || IsGPTImageGenerationModel(modelKey) +} + func modelRateLimitKeyForUpstreamModelNotFound(ctx context.Context, account *Account, requestedModel string) string { modelKey := strings.TrimSpace(requestedModel) if account == nil || modelKey == "" { diff --git a/backend/internal/service/ratelimit_service_model_not_found_test.go b/backend/internal/service/ratelimit_service_model_not_found_test.go index 55becb28a5..5e2fa0f71d 100644 --- a/backend/internal/service/ratelimit_service_model_not_found_test.go +++ b/backend/internal/service/ratelimit_service_model_not_found_test.go @@ -382,6 +382,49 @@ func TestRateLimitService_HandleUpstreamError_CodexPlanGatedModelIgnoresAPIKeyAc require.Empty(t, repo.modelRateLimitCalls) } +func TestRateLimitService_HandleUpstreamError_CodexPlanGatedImageModelSkipsCooldown(t *testing.T) { + for _, model := range []string{"gpt-image-1", "gpt-image-1.5", "gpt-image-2"} { + t.Run(model, func(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := openAICodexPlanGatedOAuthAccount() + + handled := svc.HandleUpstreamError( + context.Background(), + account, + http.StatusBadRequest, + http.Header{}, + []byte(`{"detail":"The '`+model+`' model is not supported when using Codex with a ChatGPT account."}`), + model, + ) + + require.True(t, handled, "attempt should still fail over") + require.Empty(t, repo.modelRateLimitCalls, + "image models must not be cooled down: the account still serves them over /v1/images/*") + require.Zero(t, repo.tempCalls) + }) + } +} + +func TestRateLimitService_HandleUpstreamError_CodexPlanGatedTextModelStillCoolsDown(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := openAICodexPlanGatedOAuthAccount() + + handled := svc.HandleUpstreamError( + context.Background(), + account, + http.StatusBadRequest, + http.Header{}, + []byte(`{"detail":"The 'gpt-5.6-sol' model is not supported when using Codex with a ChatGPT account."}`), + "gpt-5.6-sol", + ) + + require.True(t, handled) + require.Len(t, repo.modelRateLimitCalls, 1, "non-image plan-gated models keep the existing cooldown") + require.Equal(t, upstreamCodexPlanGatedModelReason, repo.modelRateLimitCalls[0].reason) +} + func openAICodexPlanGatedOAuthAccount() *Account { return &Account{ ID: 202, @@ -392,3 +435,93 @@ func openAICodexPlanGatedOAuthAccount() *Account { Credentials: map[string]any{}, } } + +// 请求本身就走 /v1/images/* 时必须保留冷却。 +// +// OAuth 账号的 /v1/images/* 上游同样是 Codex Responses(openai_images_responses.go +// → handleOpenAIImagesErrorResponse → handleOpenAIAccountUpstreamError → +// HandleUpstreamModelNotFound),所以这条路径也会命中 plan-gated 分支。账号确实 +// 不具备生图能力时,冷却是唯一的刹车:调度层靠 model_rate_limits 跳过该账号后 +// 快速 503;一旦跳过冷却,每个请求都会完整走一遍号池,对上游形成无上界的 400 放大。 +func TestRateLimitService_HandleUpstreamError_CodexPlanGatedImageModelKeepsCooldownOnImagesEndpoint(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := openAICodexPlanGatedOAuthAccount() + + handled := svc.HandleUpstreamError( + WithOpenAIImagesEndpoint(context.Background()), + account, + http.StatusBadRequest, + http.Header{}, + []byte(`{"detail":"The 'gpt-image-2' model is not supported when using Codex with a ChatGPT account."}`), + "gpt-image-2", + ) + + require.True(t, handled) + require.Len(t, repo.modelRateLimitCalls, 1, + "/v1/images/* 上的 plan-gated 拒绝是真实的能力缺失,必须保留冷却刹车") + require.Equal(t, "gpt-image-2", repo.modelRateLimitCalls[0].scope) + require.Equal(t, upstreamCodexPlanGatedModelReason, repo.modelRateLimitCalls[0].reason) +} + +// 仅 WithOpenAIImageGenerationIntent(/v1/responses 因模型名自动置位)不算专用生图 +// 端点,仍按"用错端点"处理。 +func TestRateLimitService_HandleUpstreamError_CodexPlanGatedImageModelSkipsCooldownOnIntentOnly(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := openAICodexPlanGatedOAuthAccount() + + handled := svc.HandleUpstreamError( + WithOpenAIImageGenerationIntent(context.Background()), + account, + http.StatusBadRequest, + http.Header{}, + []byte(`{"detail":"The 'gpt-image-2' model is not supported when using Codex with a ChatGPT account."}`), + "gpt-image-2", + ) + + require.True(t, handled) + require.Empty(t, repo.modelRateLimitCalls) +} + +// 守卫口径必须与冷却键一致:冷却键走 account.GetMappedModel,账号可以把文本别名 +// 映射到 gpt-image-*,只判请求模型会漏掉这种形态,原 bug 原样复现。 +func TestRateLimitService_HandleUpstreamError_CodexPlanGatedImageModelSkipsCooldownViaModelMapping(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := openAICodexPlanGatedOAuthAccount() + account.Credentials["model_mapping"] = map[string]any{"my-draw-alias": "gpt-image-2"} + + handled := svc.HandleUpstreamError( + context.Background(), + account, + http.StatusBadRequest, + http.Header{}, + []byte(`{"detail":"The 'gpt-image-2' model is not supported when using Codex with a ChatGPT account."}`), + "my-draw-alias", + ) + + require.True(t, handled) + require.Empty(t, repo.modelRateLimitCalls, + "映射后的上游模型是图片模型,冷却键会写到 gpt-image-2 上,守卫必须一并识别") +} + +// 404 model-not-found 分支不受守卫影响:即使是图片模型也照常冷却。 +func TestRateLimitService_HandleUpstreamError_ModelNotFoundImageModelStillCoolsDown(t *testing.T) { + repo := &modelNotFoundAccountRepoStub{} + svc := &RateLimitService{accountRepo: repo} + account := openAICodexPlanGatedOAuthAccount() + + handled := svc.HandleUpstreamError( + context.Background(), + account, + http.StatusNotFound, + http.Header{}, + []byte(`{"error":{"message":"The model 'gpt-image-2' does not exist","code":"model_not_found"}}`), + "gpt-image-2", + ) + + require.True(t, handled) + require.Len(t, repo.modelRateLimitCalls, 1, "守卫只作用于 codex plan-gated 分支") + require.Equal(t, upstreamModelNotFoundReason, repo.modelRateLimitCalls[0].reason) +} diff --git a/backend/internal/service/ratelimit_service_scheduling_threshold_test.go b/backend/internal/service/ratelimit_service_scheduling_threshold_test.go new file mode 100644 index 0000000000..5069c4dec3 --- /dev/null +++ b/backend/internal/service/ratelimit_service_scheduling_threshold_test.go @@ -0,0 +1,166 @@ +//go:build unit + +package service + +import ( + "context" + "encoding/json" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +func TestRateLimitService_ApplyAccountSchedulingThreshold_SetsTempUnschedulable(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":80}` + + accountRepo := &rateLimitAccountRepoStub{} + rl := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rl.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + + until := time.Now().UTC().Add(6 * time.Hour) + account := &Account{ + ID: 1001, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + Extra: map[string]any{ + "codex_7d_used_percent": 91.5, + "codex_7d_reset_at": until.Format(time.RFC3339), + }, + } + + blocked := rl.ApplyAccountSchedulingThreshold(context.Background(), account) + + require.True(t, blocked) + require.Equal(t, 1, accountRepo.tempCalls) + require.NotNil(t, account.TempUnschedulableUntil) + require.WithinDuration(t, until, *account.TempUnschedulableUntil, time.Second) + require.True(t, IsAccountSchedulingThresholdReason(accountRepo.lastTempReason)) + + var payload map[string]any + require.NoError(t, json.Unmarshal([]byte(accountRepo.lastTempReason), &payload)) + require.Equal(t, PlatformOpenAI, payload["platform"]) + require.Equal(t, "7d", payload["window"]) + require.Equal(t, float64(80), payload["threshold_percent"]) + require.Equal(t, float64(91.5), payload["used_percent"]) + require.Contains(t, payload["error_message"], "91.5% used >= 80%") +} + +func TestRateLimitService_ApplyAccountSchedulingThreshold_UsesAccountOverrideInReason(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":90}` + + accountRepo := &rateLimitAccountRepoStub{} + rl := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rl.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + + until := time.Now().UTC().Add(6 * time.Hour) + account := &Account{ + ID: 1003, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{ + "account_scheduling_threshold": 80, + }, + Extra: map[string]any{ + "codex_7d_used_percent": 85.5, + "codex_7d_reset_at": until.Format(time.RFC3339), + }, + } + + blocked := rl.ApplyAccountSchedulingThreshold(context.Background(), account) + + require.True(t, blocked) + require.Equal(t, 1, accountRepo.tempCalls) + + var payload map[string]any + require.NoError(t, json.Unmarshal([]byte(accountRepo.lastTempReason), &payload)) + require.Equal(t, float64(80), payload["threshold_percent"]) + require.Equal(t, float64(85.5), payload["used_percent"]) + require.Contains(t, payload["error_message"], "85.5% used >= 80%") +} + +func TestRateLimitService_ApplyAccountSchedulingThreshold_SkipsDuplicateTempUnschedulable(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":80}` + + accountRepo := &rateLimitAccountRepoStub{} + rl := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rl.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + + until := time.Now().UTC().Add(6 * time.Hour).Truncate(time.Second) + existingReason := BuildDetailedAccountSchedulingThresholdReason(AccountSchedulingThresholdReasonInput{ + Platform: PlatformOpenAI, + Window: "7d", + ThresholdPercent: 80, + UsedPercent: 91.5, + Until: until, + Now: until.Add(-time.Hour), + }) + account := &Account{ + ID: 1002, + Platform: PlatformOpenAI, + Status: StatusActive, + Schedulable: true, + TempUnschedulableUntil: &until, + TempUnschedulableReason: existingReason, + Extra: map[string]any{ + "codex_7d_used_percent": 91.5, + "codex_7d_reset_at": until.Format(time.RFC3339), + }, + } + + blocked := rl.ApplyAccountSchedulingThreshold(context.Background(), account) + + require.True(t, blocked) + require.Equal(t, 0, accountRepo.tempCalls) + require.Equal(t, existingReason, account.TempUnschedulableReason) + require.NotNil(t, account.TempUnschedulableUntil) + require.True(t, until.Equal(*account.TempUnschedulableUntil)) +} + +func TestRateLimitService_ApplyAccountSchedulingThreshold_UnsupportedPlatformDoesNotBlock(t *testing.T) { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + + settingsRepo := newMockSettingRepo() + settingsRepo.data[SettingKeyAccountSchedulingThresholds] = `{"openai":80}` + + accountRepo := &rateLimitAccountRepoStub{} + rl := NewRateLimitService(accountRepo, nil, &config.Config{}, nil, nil) + rl.SetSettingService(NewSettingService(settingsRepo, &config.Config{})) + + account := &Account{ + ID: 2002, + Platform: PlatformKiro, + Status: StatusActive, + Schedulable: true, + Credentials: map[string]any{ + "account_scheduling_threshold": 1, + }, + Extra: map[string]any{ + "kiro_sched_utilization": 99.0, + "kiro_sched_reset_at": time.Now().UTC().Add(24 * time.Hour).Format(time.RFC3339), + }, + } + + blocked := rl.ApplyAccountSchedulingThreshold(context.Background(), account) + + require.False(t, blocked) + require.Equal(t, 0, accountRepo.tempCalls) + require.Nil(t, account.TempUnschedulableUntil) + require.Empty(t, account.TempUnschedulableReason) +} diff --git a/backend/internal/service/setting_features.go b/backend/internal/service/setting_features.go index 6fba70948f..cfb83f0b5f 100644 --- a/backend/internal/service/setting_features.go +++ b/backend/internal/service/setting_features.go @@ -11,6 +11,7 @@ import ( "math" "strconv" "strings" + "time" ) // IsRegistrationEnabled 检查是否开放注册 @@ -1078,6 +1079,62 @@ func (s *SettingService) GetDefaultPlatformQuotas(ctx context.Context) (map[stri return out, nil // 补齐全部允许 platform key,保持与旧实现一致的下游契约 } +// GetAccountSchedulingThresholds returns per-platform auto-pause thresholds (1..100). +// 100 disables the threshold for that platform. Hot-path cached with singleflight. +func (s *SettingService) GetAccountSchedulingThresholds(ctx context.Context) map[string]int { + if s == nil || s.settingRepo == nil { + return defaultAccountSchedulingThresholds() + } + if cached, ok := accountSchedulingThresholdsCache.Load().(*cachedAccountSchedulingThresholds); ok { + if cached != nil && len(cached.thresholds) > 0 && time.Now().UnixNano() < cached.expiresAt { + return cloneAccountSchedulingThresholds(cached.thresholds) + } + } + + result, err, _ := accountSchedulingThresholdsSF.Do(SettingKeyAccountSchedulingThresholds, func() (any, error) { + if cached, ok := accountSchedulingThresholdsCache.Load().(*cachedAccountSchedulingThresholds); ok { + if cached != nil && len(cached.thresholds) > 0 && time.Now().UnixNano() < cached.expiresAt { + return cloneAccountSchedulingThresholds(cached.thresholds), nil + } + } + + thresholds := defaultAccountSchedulingThresholds() + dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), accountSchedulingThresholdsDBTimeout) + defer cancel() + + raw, err := s.settingRepo.GetValue(dbCtx, SettingKeyAccountSchedulingThresholds) + if err != nil { + slog.Warn("failed to get account scheduling thresholds, falling back to defaults", "error", err) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{ + thresholds: cloneAccountSchedulingThresholds(thresholds), + expiresAt: time.Now().Add(accountSchedulingThresholdsErrorTTL).UnixNano(), + }) + return cloneAccountSchedulingThresholds(thresholds), nil + } + + if trimmed := strings.TrimSpace(raw); trimmed != "" { + if parsed, err := parseAccountSchedulingThresholdsSetting(trimmed); err != nil { + slog.Warn("failed to parse account scheduling thresholds, falling back to defaults", "error", err) + } else { + thresholds = parsed + } + } + + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{ + thresholds: cloneAccountSchedulingThresholds(thresholds), + expiresAt: time.Now().Add(accountSchedulingThresholdsCacheTTL).UnixNano(), + }) + return cloneAccountSchedulingThresholds(thresholds), nil + }) + if err != nil { + return defaultAccountSchedulingThresholds() + } + if thresholds, ok := result.(map[string]int); ok { + return cloneAccountSchedulingThresholds(thresholds) + } + return defaultAccountSchedulingThresholds() +} + // GetAuthSourcePlatformQuotas 读取指定 auth source 的 platform quota 覆盖(仅返回有配置的平台,override 语义)。 func (s *SettingService) GetAuthSourcePlatformQuotas(ctx context.Context, source string) map[string]*DefaultPlatformQuotaSetting { out := map[string]*DefaultPlatformQuotaSetting{} diff --git a/backend/internal/service/setting_gateway_runtime.go b/backend/internal/service/setting_gateway_runtime.go index d2d614b34e..2f13278e3b 100644 --- a/backend/internal/service/setting_gateway_runtime.go +++ b/backend/internal/service/setting_gateway_runtime.go @@ -72,6 +72,19 @@ const gatewayForwardingCacheTTL = 60 * time.Second const gatewayForwardingErrorTTL = 5 * time.Second const gatewayForwardingDBTimeout = 5 * time.Second +// cachedAccountSchedulingThresholds 缓存平台自动停调阈值(进程内缓存,60s TTL) +type cachedAccountSchedulingThresholds struct { + thresholds map[string]int + expiresAt int64 // unix nano +} + +var accountSchedulingThresholdsCache atomic.Value // *cachedAccountSchedulingThresholds +var accountSchedulingThresholdsSF singleflight.Group + +const accountSchedulingThresholdsCacheTTL = 60 * time.Second +const accountSchedulingThresholdsErrorTTL = 5 * time.Second +const accountSchedulingThresholdsDBTimeout = 5 * time.Second + // cachedAntigravityUserAgentVersion 缓存 Antigravity UA 版本号(进程内缓存,60s TTL) type cachedAntigravityUserAgentVersion struct { version string diff --git a/backend/internal/service/setting_parse.go b/backend/internal/service/setting_parse.go index 92df463a39..7701e50e83 100644 --- a/backend/internal/service/setting_parse.go +++ b/backend/internal/service/setting_parse.go @@ -15,6 +15,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/openai" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" ) // InitializeDefaultSettings 初始化默认设置 @@ -189,6 +190,11 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error { SettingKeyChannelMonitorDefaultIntervalSeconds: "60", SettingKeyChannelMonitorHideThroughput: "true", + // Grok: safe defaults — no cross-vendor model rewrite unless operators enable it. + SettingKeyGrokDefaultTextModel: "grok-4.5", + SettingKeyGrokCrossClientModelMapEnabled: "true", + SettingKeyGrokDefaultBaseURLMode: GrokDefaultBaseURLModeCLI, + // Available channels feature (default disabled; opt-in) SettingKeyAvailableChannelsEnabled: "false", @@ -791,6 +797,16 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin // 公开读取路径给出同一个值,否则管理端看到“未隐藏”而用户端实际已隐藏。 result.ChannelMonitorHideThroughput = !isFalseSettingValue(settings[SettingKeyChannelMonitorHideThroughput]) + // Grok default mapping policy + result.GrokDefaultTextModel = strings.TrimSpace(settings[SettingKeyGrokDefaultTextModel]) + if result.GrokDefaultTextModel == "" { + result.GrokDefaultTextModel = "grok-4.5" + } + // Default true (missing/empty → enabled) so Claude/Codex→Grok mapping keeps working. + // Operators can set false to disable silent cross-client rewrite. + result.GrokCrossClientModelMapEnabled = !isFalseSettingValue(settings[SettingKeyGrokCrossClientModelMapEnabled]) + result.GrokDefaultBaseURLMode = normalizeGrokDefaultBaseURLMode(settings[SettingKeyGrokDefaultBaseURLMode]) + // Available channels feature (default: disabled; strict true) result.AvailableChannelsEnabled = settings[SettingKeyAvailableChannelsEnabled] == "true" @@ -935,9 +951,23 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin result.DefaultPlatformQuotas = parsed } } + result.AccountSchedulingThresholds = defaultAccountSchedulingThresholds() + if raw := strings.TrimSpace(settings[SettingKeyAccountSchedulingThresholds]); raw != "" { + if thresholds, err := parseAccountSchedulingThresholdsSetting(raw); err != nil { + slog.Warn("[Setting] parseSettings: unmarshal account_scheduling_thresholds failed", "error", err) + } else { + result.AccountSchedulingThresholds = thresholds + } + } result.AllowUserViewErrorRequests = settings[SettingKeyAllowUserViewErrorRequests] == "true" // default false + // Publish Grok default model_mapping options for accounts with empty mapping. + xai.SetRuntimeModelMappingOptions(xai.ModelMappingOptions{ + DefaultText: result.GrokDefaultTextModel, + EnableCrossClientMap: result.GrokCrossClientModelMapEnabled, + }) + return result } diff --git a/backend/internal/service/setting_service.go b/backend/internal/service/setting_service.go index 5d8baf4c10..f9091e8a27 100644 --- a/backend/internal/service/setting_service.go +++ b/backend/internal/service/setting_service.go @@ -5,14 +5,84 @@ import ( "encoding/json" "errors" "fmt" + "strings" "sync/atomic" "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "golang.org/x/sync/singleflight" "sync" ) +const ( + GrokDefaultBaseURLModeAPI = "api" + GrokDefaultBaseURLModeUSEast1 = "us-east-1" + GrokDefaultBaseURLModeUSWest2 = "us-west-2" + GrokDefaultBaseURLModeEUWest1 = "eu-west-1" + GrokDefaultBaseURLModeCLI = "cli" +) + +func normalizeGrokDefaultBaseURLMode(mode string) string { + switch strings.ToLower(strings.TrimSpace(mode)) { + case GrokDefaultBaseURLModeAPI: + return GrokDefaultBaseURLModeAPI + case GrokDefaultBaseURLModeUSEast1: + return GrokDefaultBaseURLModeUSEast1 + case GrokDefaultBaseURLModeUSWest2: + return GrokDefaultBaseURLModeUSWest2 + case GrokDefaultBaseURLModeEUWest1: + return GrokDefaultBaseURLModeEUWest1 + case GrokDefaultBaseURLModeCLI: + return GrokDefaultBaseURLModeCLI + default: + return GrokDefaultBaseURLModeCLI + } +} + +func GrokBaseURLForMode(mode string) string { + switch normalizeGrokDefaultBaseURLMode(mode) { + case GrokDefaultBaseURLModeAPI: + return xai.DefaultBaseURL + case GrokDefaultBaseURLModeUSEast1: + return xai.DefaultUSEast1BaseURL + case GrokDefaultBaseURLModeUSWest2: + return xai.DefaultUSWest2BaseURL + case GrokDefaultBaseURLModeEUWest1: + return xai.DefaultEUWest1BaseURL + default: + return xai.DefaultCLIBaseURL + } +} + +func (s *SettingService) GetGrokDefaultBaseURLMode(ctx context.Context) string { + if s == nil || s.settingRepo == nil { + return GrokDefaultBaseURLModeCLI + } + dbCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), gatewayForwardingDBTimeout) + defer cancel() + raw, err := s.settingRepo.GetValue(dbCtx, SettingKeyGrokDefaultBaseURLMode) + if err != nil { + return GrokDefaultBaseURLModeCLI + } + return normalizeGrokDefaultBaseURLMode(raw) +} + +func (s *SettingService) GetGrokDefaultBaseURL(ctx context.Context) string { + return GrokBaseURLForMode(s.GetGrokDefaultBaseURLMode(ctx)) +} + +func (s *SettingService) ResolveGrokBaseURL(ctx context.Context, account *Account) string { + def := xai.DefaultCLIBaseURL + if s != nil { + def = s.GetGrokDefaultBaseURL(ctx) + } + if account == nil { + return def + } + return account.GetGrokBaseURLOr(def) +} + var ( ErrRegistrationDisabled = infraerrors.Forbidden("REGISTRATION_DISABLED", "registration is currently disabled") ErrSettingNotFound = infraerrors.NotFound("SETTING_NOT_FOUND", "setting not found") diff --git a/backend/internal/service/setting_service_platform_threshold_test.go b/backend/internal/service/setting_service_platform_threshold_test.go new file mode 100644 index 0000000000..7b95cb8f78 --- /dev/null +++ b/backend/internal/service/setting_service_platform_threshold_test.go @@ -0,0 +1,151 @@ +//go:build unit + +package service + +import ( + "context" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/stretchr/testify/require" +) + +func newSettingServiceForPlatformThresholdTest(seed map[string]string) *SettingService { + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + repo := newMockSettingRepo() + for k, v := range seed { + repo.data[k] = v + } + return NewSettingService(repo, &config.Config{}) +} + +func TestPlatformSchedulingThresholds_RoundTrip_DefaultsAndStoredValues(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(nil) + + got := svc.parseSettings(map[string]string{}) + require.Equal(t, map[string]int{ + PlatformOpenAI: 100, + PlatformAnthropic: 100, + PlatformGrok: 100, + }, got.AccountSchedulingThresholds) + + got = svc.parseSettings(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":91,"grok":77,"gemini":85,"kiro":99}`, + }) + require.Equal(t, 91, got.AccountSchedulingThresholds[PlatformOpenAI]) + require.Equal(t, 100, got.AccountSchedulingThresholds[PlatformAnthropic]) + require.Equal(t, 77, got.AccountSchedulingThresholds[PlatformGrok]) + require.NotContains(t, got.AccountSchedulingThresholds, PlatformGemini) + require.NotContains(t, got.AccountSchedulingThresholds, "kiro") +} + +func TestBuildSystemSettingsUpdates_PersistsAccountSchedulingThresholds(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(nil) + + updates, err := svc.buildSystemSettingsUpdates(context.Background(), &SystemSettings{ + AccountSchedulingThresholds: map[string]int{ + PlatformOpenAI: 91, + PlatformAnthropic: 88, + PlatformGrok: 77, + }, + }) + require.NoError(t, err) + require.JSONEq(t, `{"openai":91,"anthropic":88,"grok":77}`, updates[SettingKeyAccountSchedulingThresholds]) +} + +func TestValidateAndNormalizeAccountSchedulingThresholds_FillsMissingPlatforms(t *testing.T) { + normalized, err := validateAndNormalizeAccountSchedulingThresholds(map[string]int{ + PlatformOpenAI: 91, + }) + require.NoError(t, err) + require.Equal(t, 91, normalized[PlatformOpenAI]) + require.Equal(t, 100, normalized[PlatformAnthropic]) + require.Equal(t, 100, normalized[PlatformGrok]) + require.NotContains(t, normalized, PlatformGemini) + require.NotContains(t, normalized, "kiro") + require.NotContains(t, normalized, PlatformAntigravity) +} + +func TestValidateAndNormalizeAccountSchedulingThresholds_RejectsUnsupportedPlatforms(t *testing.T) { + _, err := validateAndNormalizeAccountSchedulingThresholds(map[string]int{ + PlatformGemini: 85, + }) + require.Error(t, err) +} + +func TestUpdateSettings_StoresAccountSchedulingThresholds(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(nil) + + err := svc.UpdateSettings(context.Background(), &SystemSettings{ + AccountSchedulingThresholds: map[string]int{ + PlatformOpenAI: 92, + PlatformAnthropic: 89, + PlatformGrok: 76, + }, + }) + require.NoError(t, err) + + got := svc.parseSettings(map[string]string{ + SettingKeyAccountSchedulingThresholds: svc.settingRepo.(*mockSettingRepo).data[SettingKeyAccountSchedulingThresholds], + }) + require.Equal(t, 92, got.AccountSchedulingThresholds[PlatformOpenAI]) + require.Equal(t, 89, got.AccountSchedulingThresholds[PlatformAnthropic]) + require.Equal(t, 76, got.AccountSchedulingThresholds[PlatformGrok]) + require.NotContains(t, got.AccountSchedulingThresholds, "kiro") +} + +func TestGetAccountSchedulingThresholds_ReadsStoredValue(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":93,"grok":88,"kiro":87}`, + }) + + got := svc.GetAccountSchedulingThresholds(context.Background()) + + require.Equal(t, 93, got[PlatformOpenAI]) + require.Equal(t, 100, got[PlatformAnthropic]) + require.Equal(t, 88, got[PlatformGrok]) + require.NotContains(t, got, "kiro") +} + +func TestUpdateSettings_OmittedAccountSchedulingThresholdsDoesNotCacheDefaults(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":85,"grok":88,"kiro":87}`, + }) + + err := svc.UpdateSettings(context.Background(), &SystemSettings{ + FrontendURL: "https://example.test", + }) + require.NoError(t, err) + + got := svc.GetAccountSchedulingThresholds(context.Background()) + require.Equal(t, 85, got[PlatformOpenAI]) + require.Equal(t, 88, got[PlatformGrok]) + require.NotContains(t, got, "kiro") +} + +func TestAccountSchedulingThresholds_InvalidStoredValueUsesSameDefaultsInSettingsAndCache(t *testing.T) { + svc := newSettingServiceForPlatformThresholdTest(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":0,"grok":88,"kiro":87}`, + }) + + settings := svc.parseSettings(map[string]string{ + SettingKeyAccountSchedulingThresholds: `{"openai":0,"grok":88,"kiro":87}`, + }) + cached := svc.GetAccountSchedulingThresholds(context.Background()) + + require.Equal(t, settings.AccountSchedulingThresholds, cached) + require.Equal(t, 100, cached[PlatformOpenAI]) + require.Equal(t, 88, cached[PlatformGrok]) + require.NotContains(t, cached, "kiro") +} + +func TestGetAccountSchedulingThresholds_NilRepoReturnsDefaults(t *testing.T) { + svc := &SettingService{} + got := svc.GetAccountSchedulingThresholds(context.Background()) + require.Equal(t, map[string]int{ + PlatformOpenAI: 100, + PlatformAnthropic: 100, + PlatformGrok: 100, + }, got) +} diff --git a/backend/internal/service/setting_update.go b/backend/internal/service/setting_update.go index 6f38a4d366..1d40f20ca8 100644 --- a/backend/internal/service/setting_update.go +++ b/backend/internal/service/setting_update.go @@ -417,6 +417,15 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting } updates[SettingKeyChannelMonitorHideThroughput] = strconv.FormatBool(settings.ChannelMonitorHideThroughput) + // Grok model mapping policy + if v := strings.TrimSpace(settings.GrokDefaultTextModel); v != "" { + updates[SettingKeyGrokDefaultTextModel] = v + } else { + updates[SettingKeyGrokDefaultTextModel] = "grok-4.5" + } + updates[SettingKeyGrokCrossClientModelMapEnabled] = strconv.FormatBool(settings.GrokCrossClientModelMapEnabled) + updates[SettingKeyGrokDefaultBaseURLMode] = normalizeGrokDefaultBaseURLMode(settings.GrokDefaultBaseURLMode) + // Available channels feature switch updates[SettingKeyAvailableChannelsEnabled] = strconv.FormatBool(settings.AvailableChannelsEnabled) @@ -513,12 +522,88 @@ func (s *SettingService) buildSystemSettingsUpdates(ctx context.Context, setting } updates[SettingKeyDefaultPlatformQuotas] = string(blob) } + if settings.AccountSchedulingThresholds != nil { + normalized, err := validateAndNormalizeAccountSchedulingThresholds(settings.AccountSchedulingThresholds) + if err != nil { + return nil, err + } + blob, err := json.Marshal(normalized) + if err != nil { + return nil, fmt.Errorf("marshal account scheduling thresholds: %w", err) + } + updates[SettingKeyAccountSchedulingThresholds] = string(blob) + } updates[SettingKeyAllowUserViewErrorRequests] = strconv.FormatBool(settings.AllowUserViewErrorRequests) return updates, nil } +func defaultAccountSchedulingThresholds() map[string]int { + return map[string]int{ + PlatformOpenAI: 100, + PlatformAnthropic: 100, + PlatformGrok: 100, + } +} + +func validateAndNormalizeAccountSchedulingThresholds(input map[string]int) (map[string]int, error) { + normalized := defaultAccountSchedulingThresholds() + for platform, value := range input { + allowed := false + for _, item := range AllowedSchedulingThresholdPlatforms { + if item == platform { + allowed = true + break + } + } + if !allowed { + return nil, infraerrors.BadRequest("INVALID_ACCOUNT_SCHEDULING_THRESHOLDS", fmt.Sprintf("unknown platform %q", platform)) + } + if value < 1 || value > 100 { + return nil, infraerrors.BadRequest("INVALID_ACCOUNT_SCHEDULING_THRESHOLDS", "platform scheduling threshold must be between 1 and 100") + } + normalized[platform] = value + } + return normalized, nil +} + +func parseAccountSchedulingThresholdsSetting(raw string) (map[string]int, error) { + thresholds := defaultAccountSchedulingThresholds() + raw = strings.TrimSpace(raw) + if raw == "" { + return thresholds, nil + } + parsed := map[string]int{} + if err := json.Unmarshal([]byte(raw), &parsed); err != nil { + return thresholds, err + } + for _, platform := range AllowedSchedulingThresholdPlatforms { + if value, ok := parsed[platform]; ok { + thresholds[platform] = boundedIntOrDefault(value, 1, 100, 100) + } + } + return thresholds, nil +} + +func boundedIntOrDefault(value, minValue, maxValue, defaultValue int) int { + if value < minValue || value > maxValue { + return defaultValue + } + return value +} + +func cloneAccountSchedulingThresholds(input map[string]int) map[string]int { + if len(input) == 0 { + return defaultAccountSchedulingThresholds() + } + cloned := make(map[string]int, len(input)) + for key, value := range input { + cloned[key] = value + } + return cloned +} + // validateDefaultPlatformQuotaMap 校验 platform quota map 的合法性: // 平台名须在 AllowedQuotaPlatforms 白名单内,每个非 nil 上限须 finite 且 >= 0。 // 系统层和 auth-source 层共用此 helper。 @@ -674,6 +759,20 @@ func (s *SettingService) refreshCachedSettings(settings *SystemSettings) { expiresAt: 0, }) } + accountSchedulingThresholdsSF.Forget(SettingKeyAccountSchedulingThresholds) + if settings.AccountSchedulingThresholds != nil { + normalizedThresholds, err := validateAndNormalizeAccountSchedulingThresholds(settings.AccountSchedulingThresholds) + if err != nil { + normalizedThresholds = defaultAccountSchedulingThresholds() + } + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{ + thresholds: cloneAccountSchedulingThresholds(normalizedThresholds), + expiresAt: time.Now().Add(accountSchedulingThresholdsCacheTTL).UnixNano(), + }) + } else { + // Partial/omitted payload: clear cache so the next hot-path read reloads from DB. + accountSchedulingThresholdsCache.Store(&cachedAccountSchedulingThresholds{}) + } if s.cfg != nil { s.cfg.SetForwardedClientIPSettings(settings.APIKeyACLTrustForwardedIP, settings.ForwardedClientIPHeaders) } diff --git a/backend/internal/service/settings_view.go b/backend/internal/service/settings_view.go index 28fa75fce6..3ccdd02398 100644 --- a/backend/internal/service/settings_view.go +++ b/backend/internal/service/settings_view.go @@ -201,6 +201,11 @@ type SystemSettings struct { ChannelMonitorDefaultIntervalSeconds int `json:"channel_monitor_default_interval_seconds"` ChannelMonitorHideThroughput bool `json:"channel_monitor_hide_throughput"` + // Grok model mapping policy (admin settings; empty mapping falls back to these). + GrokDefaultTextModel string `json:"grok_default_text_model"` + GrokCrossClientModelMapEnabled bool `json:"grok_cross_client_model_map_enabled"` + GrokDefaultBaseURLMode string `json:"grok_default_base_url_mode"` + // Available Channels feature (user-facing aggregate view) AvailableChannelsEnabled bool `json:"available_channels_enabled"` @@ -294,6 +299,9 @@ type SystemSettings struct { // 系统全局默认平台配额(key = platform,nil/缺省 = 不限制) DefaultPlatformQuotas map[string]*DefaultPlatformQuotaSetting `json:"default_platform_quotas"` + // 系统全局账号自动停调阈值(key = platform,100 = disabled) + AccountSchedulingThresholds map[string]int `json:"account_scheduling_thresholds"` + // 允许终端用户在用量页查看自己的失败请求 AllowUserViewErrorRequests bool } @@ -369,6 +377,11 @@ type PublicSettings struct { ChannelMonitorDefaultIntervalSeconds int `json:"channel_monitor_default_interval_seconds"` ChannelMonitorHideThroughput bool `json:"channel_monitor_hide_throughput"` + // Grok model mapping policy (admin settings). + GrokDefaultTextModel string `json:"grok_default_text_model"` + GrokCrossClientModelMapEnabled bool `json:"grok_cross_client_model_map_enabled"` + GrokDefaultBaseURLMode string `json:"grok_default_base_url_mode"` + // Available Channels feature (user-facing aggregate view) AvailableChannelsEnabled bool `json:"available_channels_enabled"` diff --git a/backend/internal/service/subscription_assign_idempotency_test.go b/backend/internal/service/subscription_assign_idempotency_test.go index d6dde5e944..ed6c4e0aae 100644 --- a/backend/internal/service/subscription_assign_idempotency_test.go +++ b/backend/internal/service/subscription_assign_idempotency_test.go @@ -9,6 +9,7 @@ import ( dbent "github.com/Wei-Shaw/sub2api/ent" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" "github.com/dgraph-io/ristretto" "github.com/stretchr/testify/require" ) @@ -157,10 +158,10 @@ func (userSubRepoNoop) UpdateStatus(context.Context, int64, string) error { func (userSubRepoNoop) UpdateNotes(context.Context, int64, string) error { panic("unexpected UpdateNotes call") } -func (userSubRepoNoop) ActivateWindows(context.Context, int64, time.Time) error { +func (userSubRepoNoop) ActivateWindows(context.Context, int64, time.Time, time.Time) error { panic("unexpected ActivateWindows call") } -func (userSubRepoNoop) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { +func (userSubRepoNoop) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time, time.Time) error { panic("unexpected ResetUsageWindows call") } func (userSubRepoNoop) ResetDailyUsage(context.Context, int64, *time.Time, time.Time) error { @@ -424,7 +425,7 @@ func TestAssignSubscriptionRenewsExpiredSemanticMatch(t *testing.T) { require.False(t, sub.StartsAt.Before(before)) require.False(t, sub.StartsAt.After(after)) require.Equal(t, sub.StartsAt.AddDate(0, 0, 30), sub.ExpiresAt) - require.Equal(t, sub.StartsAt, *sub.DailyWindowStart) + require.Equal(t, timezone.StartOfDay(sub.StartsAt), *sub.DailyWindowStart, "续期后日窗口应锚定当天 0 点") require.Equal(t, sub.StartsAt, *sub.WeeklyWindowStart) require.Equal(t, sub.StartsAt, *sub.MonthlyWindowStart) require.Zero(t, sub.DailyUsageUSD) diff --git a/backend/internal/service/subscription_daily_midnight_reset_test.go b/backend/internal/service/subscription_daily_midnight_reset_test.go new file mode 100644 index 0000000000..5f16847511 --- /dev/null +++ b/backend/internal/service/subscription_daily_midnight_reset_test.go @@ -0,0 +1,162 @@ +package service + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" + "github.com/stretchr/testify/require" +) + +// dailyMidnightResetRepo 记录 ResetDailyUsage 收到的新窗口起点。 +type dailyMidnightResetRepo struct { + userSubRepoNoop + + resetCalled bool + newWindowStart time.Time +} + +func (r *dailyMidnightResetRepo) ResetDailyUsage(_ context.Context, _ int64, _ *time.Time, newWindowStart time.Time) error { + r.resetCalled = true + r.newWindowStart = newWindowStart + return nil +} + +// midnightTestBase 返回一个固定日期在配置时区的 0 点,测试用时刻均从它推导, +// 保证断言在任意本地时区下都与生产逻辑同构。 +func midnightTestBase() time.Time { + return timezone.StartOfDay(time.Date(2026, 8, 6, 12, 0, 0, 0, timezone.Location())) +} + +func newMidnightTestSub(dailyWindowStart time.Time, base time.Time) *UserSubscription { + start := dailyWindowStart + return &UserSubscription{ + ID: 1, + UserID: 10, + GroupID: 20, + StartsAt: base.AddDate(0, 0, -3), + ExpiresAt: base.AddDate(0, 0, 30), + DailyUsageUSD: 43.34, + DailyWindowStart: &start, + } +} + +// 手动重置(或任何原因)留下的非 0 点锚点,跨 0 点后必须重置, +// 且新窗口起点是当天 0 点,而不是锚点+24h 的滚动时刻。 +func TestCheckAndResetWindows_DailyResetsAtMidnightNotRollingAnchor(t *testing.T) { + base := midnightTestBase() + manualResetAt := base.Add(16*time.Hour + 49*time.Minute) // 昨日 16:49 手动重置 + now := base.AddDate(0, 0, 1).Add(5 * time.Minute) // 次日 00:05 + + repo := &dailyMidnightResetRepo{} + svc := NewSubscriptionService(groupRepoNoop{}, repo, nil, nil, nil) + svc.now = func() time.Time { return now } + sub := newMidnightTestSub(manualResetAt, base) + + require.NoError(t, svc.CheckAndResetWindows(context.Background(), sub)) + + require.True(t, repo.resetCalled, "跨 0 点后日窗口必须重置") + require.Equal(t, base.AddDate(0, 0, 1), repo.newWindowStart, "新窗口起点应为当天 0 点") + require.Zero(t, sub.DailyUsageUSD) + require.Equal(t, base.AddDate(0, 0, 1), *sub.DailyWindowStart) +} + +// 同一日历日内(即使已过锚点+若干小时)不得重置:手动重置当天剩余时间继续累计用量。 +func TestCheckAndResetWindows_DailyNoResetWithinSameCalendarDay(t *testing.T) { + base := midnightTestBase() + manualResetAt := base.Add(16*time.Hour + 49*time.Minute) + now := base.Add(23*time.Hour + 59*time.Minute) + + repo := &dailyMidnightResetRepo{} + svc := NewSubscriptionService(groupRepoNoop{}, repo, nil, nil, nil) + svc.now = func() time.Time { return now } + sub := newMidnightTestSub(manualResetAt, base) + + require.NoError(t, svc.CheckAndResetWindows(context.Background(), sub)) + + require.False(t, repo.resetCalled, "同一日历日内不应重置日窗口") + require.Equal(t, 43.34, sub.DailyUsageUSD) +} + +// 0.1.170/171 期间产生的滚动锚点(如多日前的 17:18)自愈:下一次维护即拉回今天 0 点。 +func TestCheckAndResetWindows_LegacyRollingAnchorHealsToMidnight(t *testing.T) { + base := midnightTestBase() + staleAnchor := base.AddDate(0, 0, -3).Add(17*time.Hour + 18*time.Minute) + now := base.Add(10 * time.Hour) + + repo := &dailyMidnightResetRepo{} + svc := NewSubscriptionService(groupRepoNoop{}, repo, nil, nil, nil) + svc.now = func() time.Time { return now } + sub := newMidnightTestSub(staleAnchor, base) + + require.NoError(t, svc.CheckAndResetWindows(context.Background(), sub)) + + require.True(t, repo.resetCalled) + require.Equal(t, base, repo.newWindowStart, "滚动锚点应被拉回今天 0 点,而非按 17:18 步进") +} + +// 手动重置写入的 0 点锚点不会改变刷新节奏:次日 0 点照常需要重置。 +func TestNeedsDailyReset_MidnightScheduleSurvivesManualReset(t *testing.T) { + base := midnightTestBase() + sub := newMidnightTestSub(base, base) // 手动重置后锚点=当天 0 点 + + require.False(t, sub.NeedsDailyResetAt(base.Add(23*time.Hour+54*time.Minute)), "当天 23:54 不应重置") + require.True(t, sub.NeedsDailyResetAt(base.AddDate(0, 0, 1).Add(time.Minute)), "次日 00:01 应重置") +} + +// 多日订阅的日窗口展示的下次刷新时间固定为次日 0 点(截图中「X 小时后重置」的数据源)。 +func TestDailyResetTime_NextMidnightForMultiDaySubscription(t *testing.T) { + base := midnightTestBase() + + // 0 点锚点 → 次日 0 点 + sub := newMidnightTestSub(base, base) + resetAt := sub.DailyResetTime() + require.NotNil(t, resetAt) + require.Equal(t, base.AddDate(0, 0, 1), *resetAt) + + // 非 0 点滚动锚点 → 仍是其所在日的次日 0 点,而非锚点+24h + rolling := newMidnightTestSub(base.Add(16*time.Hour+49*time.Minute), base) + resetAt = rolling.DailyResetTime() + require.NotNil(t, resetAt) + require.Equal(t, base.AddDate(0, 0, 1), *resetAt) +} + +// 截图场景:跨 0 点后(后台尚未收到请求推进窗口)列表展示即应清零日用量。 +func TestNormalizeExpiredWindows_DailyUsageClearsAfterMidnight(t *testing.T) { + base := midnightTestBase() + manualResetAt := base.Add(16*time.Hour + 49*time.Minute) + now := base.AddDate(0, 0, 1).Add(time.Minute) // 次日 00:01 + + subs := []UserSubscription{*newMidnightTestSub(manualResetAt, base)} + normalizeExpiredWindowsAt(subs, now) + + require.Zero(t, subs[0].DailyUsageUSD, "跨 0 点后展示的日用量应清零") + require.Nil(t, subs[0].DailyWindowStart) +} + +// 日卡(一次性日额度)不受 0 点语义影响:跨 0 点不重置。 +func TestCheckAndResetWindows_OneTimeDailyCardStillExemptFromMidnightReset(t *testing.T) { + base := midnightTestBase() + startsAt := base.Add(17 * time.Hour) + anchor := base + now := base.AddDate(0, 0, 1).Add(2 * time.Hour) + + repo := &dailyMidnightResetRepo{} + svc := NewSubscriptionService(groupRepoNoop{}, repo, nil, nil, nil) + svc.now = func() time.Time { return now } + sub := &UserSubscription{ + ID: 1, + UserID: 10, + GroupID: 20, + StartsAt: startsAt, + ExpiresAt: startsAt.AddDate(0, 0, 1), + DailyUsageUSD: 10, + DailyWindowStart: &anchor, + } + + require.NoError(t, svc.CheckAndResetWindows(context.Background(), sub)) + + require.False(t, repo.resetCalled, "日卡为一次性配额,跨 0 点不应重置") + require.Equal(t, 10.0, sub.DailyUsageUSD) +} diff --git a/backend/internal/service/subscription_expiry_service_test.go b/backend/internal/service/subscription_expiry_service_test.go index bb4372fa46..74b82e56c3 100644 --- a/backend/internal/service/subscription_expiry_service_test.go +++ b/backend/internal/service/subscription_expiry_service_test.go @@ -87,11 +87,11 @@ func (r *subscriptionExpiryRepoStub) UpdateNotes(context.Context, int64, string) return nil } -func (r *subscriptionExpiryRepoStub) ActivateWindows(context.Context, int64, time.Time) error { +func (r *subscriptionExpiryRepoStub) ActivateWindows(context.Context, int64, time.Time, time.Time) error { return nil } -func (r *subscriptionExpiryRepoStub) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time) error { +func (r *subscriptionExpiryRepoStub) ResetUsageWindows(context.Context, int64, bool, bool, bool, time.Time, time.Time) error { return nil } diff --git a/backend/internal/service/subscription_monthly_window_test.go b/backend/internal/service/subscription_monthly_window_test.go index 6cd4410c7b..3e71ca17ef 100644 --- a/backend/internal/service/subscription_monthly_window_test.go +++ b/backend/internal/service/subscription_monthly_window_test.go @@ -8,12 +8,14 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" "github.com/stretchr/testify/require" ) type activateWindowUserSubRepo struct { userSubRepoNoop - windowStart time.Time + dailyStart time.Time + periodicStart time.Time } type monthlyResetUserSubRepo struct { @@ -28,8 +30,9 @@ func (r *monthlyResetUserSubRepo) ResetMonthlyUsage(_ context.Context, _ int64, return nil } -func (r *activateWindowUserSubRepo) ActivateWindows(_ context.Context, _ int64, start time.Time) error { - r.windowStart = start +func (r *activateWindowUserSubRepo) ActivateWindows(_ context.Context, _ int64, dailyStart, periodicStart time.Time) error { + r.dailyStart = dailyStart + r.periodicStart = periodicStart return nil } @@ -47,8 +50,9 @@ func TestDelayedFirstUseAnchorsMonthlyWindowAtActivation(t *testing.T) { require.NoError(t, svc.CheckAndActivateWindow(context.Background(), sub)) - require.Equal(t, activatedAt, repo.windowStart) - monthlyWindowStart := repo.windowStart + require.Equal(t, activatedAt, repo.periodicStart) + require.Equal(t, timezone.StartOfDay(activatedAt), repo.dailyStart) + monthlyWindowStart := repo.periodicStart resetAt, ok := sub.automaticWindowStartAt(&monthlyWindowStart, 30*24*time.Hour, activatedAt.Add(30*24*time.Hour)) require.True(t, ok) require.Equal(t, activatedAt.Add(30*24*time.Hour), resetAt) diff --git a/backend/internal/service/subscription_reset_quota_test.go b/backend/internal/service/subscription_reset_quota_test.go index df16db1d56..43ea5b0d10 100644 --- a/backend/internal/service/subscription_reset_quota_test.go +++ b/backend/internal/service/subscription_reset_quota_test.go @@ -8,6 +8,7 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" "github.com/stretchr/testify/require" ) @@ -24,7 +25,8 @@ type resetQuotaUserSubRepoStub struct { resetDailyErr error resetWeeklyErr error resetMonthlyErr error - windowStart time.Time + dailyStart time.Time + periodicStart time.Time } func (r *resetQuotaUserSubRepoStub) GetByID(_ context.Context, id int64) (*UserSubscription, error) { @@ -35,11 +37,12 @@ func (r *resetQuotaUserSubRepoStub) GetByID(_ context.Context, id int64) (*UserS return &cp, nil } -func (r *resetQuotaUserSubRepoStub) ResetUsageWindows(_ context.Context, _ int64, resetDaily, resetWeekly, resetMonthly bool, windowStart time.Time) error { +func (r *resetQuotaUserSubRepoStub) ResetUsageWindows(_ context.Context, _ int64, resetDaily, resetWeekly, resetMonthly bool, dailyStart, periodicStart time.Time) error { r.resetDailyCalled = resetDaily r.resetWeeklyCalled = resetWeekly r.resetMonthlyCalled = resetMonthly - r.windowStart = windowStart + r.dailyStart = dailyStart + r.periodicStart = periodicStart if resetDaily && r.resetDailyErr != nil { return r.resetDailyErr } @@ -54,15 +57,15 @@ func (r *resetQuotaUserSubRepoStub) ResetUsageWindows(_ context.Context, _ int64 } if resetDaily { r.sub.DailyUsageUSD = 0 - r.sub.DailyWindowStart = &windowStart + r.sub.DailyWindowStart = &dailyStart } if resetWeekly { r.sub.WeeklyUsageUSD = 0 - r.sub.WeeklyWindowStart = &windowStart + r.sub.WeeklyWindowStart = &periodicStart } if resetMonthly { r.sub.MonthlyUsageUSD = 0 - r.sub.MonthlyWindowStart = &windowStart + r.sub.MonthlyWindowStart = &periodicStart } return nil } @@ -105,8 +108,10 @@ func TestAdminResetQuota_ResetBoth(t *testing.T) { require.True(t, stub.resetDailyCalled, "应调用 ResetDailyUsage") require.True(t, stub.resetWeeklyCalled, "应调用 ResetWeeklyUsage") require.False(t, stub.resetMonthlyCalled, "不应调用 ResetMonthlyUsage") - require.Equal(t, resetAt, stub.windowStart) - require.Equal(t, resetAt, *result.DailyWindowStart) + // 手动重置后日窗口锚定当天 0 点(保持 0 点刷新节奏),周窗口锚定重置时刻。 + require.Equal(t, timezone.StartOfDay(resetAt), stub.dailyStart) + require.Equal(t, resetAt, stub.periodicStart) + require.Equal(t, timezone.StartOfDay(resetAt), *result.DailyWindowStart) require.Equal(t, resetAt, *result.WeeklyWindowStart) } diff --git a/backend/internal/service/subscription_service.go b/backend/internal/service/subscription_service.go index 47b18c582e..7d716b7a8e 100644 --- a/backend/internal/service/subscription_service.go +++ b/backend/internal/service/subscription_service.go @@ -13,6 +13,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/config" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" "github.com/dgraph-io/ristretto" "golang.org/x/sync/singleflight" ) @@ -380,13 +381,15 @@ func (s *SubscriptionService) withSubscriptionUpdateTx(ctx context.Context, fn f func renewedSubscriptionTerm(existingSub *UserSubscription, notes string, startsAt, expiresAt time.Time) *UserSubscription { renewed := *existingSub - windowStart := startsAt + // 日窗口按日历日对齐(0 点刷新);周/月窗口按订阅期限对齐(锚点为新周期起点)。 + dailyWindowStart := timezone.StartOfDay(startsAt) + periodicWindowStart := startsAt renewed.StartsAt = startsAt renewed.ExpiresAt = expiresAt renewed.Status = SubscriptionStatusActive - renewed.DailyWindowStart = &windowStart - renewed.WeeklyWindowStart = &windowStart - renewed.MonthlyWindowStart = &windowStart + renewed.DailyWindowStart = &dailyWindowStart + renewed.WeeklyWindowStart = &periodicWindowStart + renewed.MonthlyWindowStart = &periodicWindowStart renewed.DailyUsageUSD = 0 renewed.WeeklyUsageUSD = 0 renewed.MonthlyUsageUSD = 0 @@ -858,7 +861,9 @@ func (s *SubscriptionService) checkAndActivateWindowAt(ctx context.Context, sub return nil } - return s.userSubRepo.ActivateWindows(ctx, sub.ID, now) + // 日窗口锚定当天 0 点(日历日语义);周/月窗口锚定首次使用时刻(期限对齐语义, + // 锚点不得早于 StartsAt,否则最后一个不完整周期会重复发放额度,见 issue #5051)。 + return s.userSubRepo.ActivateWindows(ctx, sub.ID, timezone.StartOfDay(now), now) } // AdminResetQuota manually resets the daily, weekly, and/or monthly usage windows. @@ -870,8 +875,10 @@ func (s *SubscriptionService) AdminResetQuota(ctx context.Context, subscriptionI if err != nil { return nil, err } - windowStart := s.now() - if err := s.userSubRepo.ResetUsageWindows(ctx, sub.ID, resetDaily, resetWeekly, resetMonthly, windowStart); err != nil { + now := s.now() + // 日窗口锚点取当天 0 点:手动重置只清空用量,不改变“每天 0 点刷新”的节奏。 + // 周/月窗口保持锚定重置时刻(期限对齐滚动窗口语义)。 + if err := s.userSubRepo.ResetUsageWindows(ctx, sub.ID, resetDaily, resetWeekly, resetMonthly, timezone.StartOfDay(now), now); err != nil { return nil, err } // Invalidate L1 ristretto cache. Ristretto's Del() is asynchronous by design, @@ -890,8 +897,8 @@ func (s *SubscriptionService) CheckAndResetWindows(ctx context.Context, sub *Use now := s.now() needsInvalidateCache := false - // 日窗口重置(24小时) - if windowStart, ok := sub.automaticWindowStartAt(sub.DailyWindowStart, 24*time.Hour, now); !sub.HasOneTimeDailyQuota() && ok { + // 日窗口重置(每天 0 点刷新,按日历日对齐) + if windowStart, ok := sub.automaticDailyWindowStartAt(now); ok { expectedWindowStart := sub.DailyWindowStart if err := s.userSubRepo.ResetDailyUsage(ctx, sub.ID, expectedWindowStart, windowStart); err != nil { return err diff --git a/backend/internal/service/temp_unsched.go b/backend/internal/service/temp_unsched.go index 3871b72b58..06ae726df9 100644 --- a/backend/internal/service/temp_unsched.go +++ b/backend/internal/service/temp_unsched.go @@ -7,12 +7,15 @@ import ( // TempUnschedState 临时不可调度状态 type TempUnschedState struct { - UntilUnix int64 `json:"until_unix"` // 解除时间(Unix 时间戳) - TriggeredAtUnix int64 `json:"triggered_at_unix"` // 触发时间(Unix 时间戳) - StatusCode int `json:"status_code"` // 触发的错误码 - MatchedKeyword string `json:"matched_keyword"` // 匹配的关键词 - RuleIndex int `json:"rule_index"` // 触发的规则索引 - ErrorMessage string `json:"error_message"` // 错误消息 + UntilUnix int64 `json:"until_unix"` // 解除时间(Unix 时间戳) + TriggeredAtUnix int64 `json:"triggered_at_unix"` // 触发时间(Unix 时间戳) + StatusCode int `json:"status_code"` // 触发的错误码 + MatchedKeyword string `json:"matched_keyword"` // 匹配的关键词 + RuleIndex int `json:"rule_index"` // 触发的规则索引 + ErrorMessage string `json:"error_message"` // 错误消息 + TriggerCount int64 `json:"trigger_count,omitempty"` // 本次触发累计命中次数 + TriggerThreshold int `json:"trigger_threshold,omitempty"` // 触发阈值 + TriggerWindowMinutes int `json:"trigger_window_minutes,omitempty"` // 计数窗口(分钟) } // TempUnschedCache 临时不可调度缓存接口 diff --git a/backend/internal/service/token_refresh_service.go b/backend/internal/service/token_refresh_service.go index 31d2e19909..bbf1e8655e 100644 --- a/backend/internal/service/token_refresh_service.go +++ b/backend/internal/service/token_refresh_service.go @@ -1200,6 +1200,10 @@ func (s *TokenRefreshService) postRefreshActions(ctx context.Context, account *A s.ensureOpenAIPrivacy(ctx, account) // Antigravity OAuth: 刷新成功后,检查是否已设置 privacy_mode,未设置则调用 setUserSettings s.ensureAntigravityPrivacy(ctx, account) + // Grok: clear soft reauth flag after a successful credential refresh. + if account != nil && account.Platform == PlatformGrok && accountGrokNeedsReauth(account) { + clearGrokNeedsReauthExtra(ctx, s.accountRepo, account.ID) + } } func (s *TokenRefreshService) postRefreshStateSyncWithCleanup(parent context.Context, account *Account) { diff --git a/backend/internal/service/upstream_models.go b/backend/internal/service/upstream_models.go index 393d57870c..9a4e4311b0 100644 --- a/backend/internal/service/upstream_models.go +++ b/backend/internal/service/upstream_models.go @@ -193,7 +193,11 @@ func (s *AccountTestService) buildGrokUpstreamModelsRequest(ctx context.Context, if err != nil { return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err) } - validatedBaseURL, err := validator(account.GetGrokBaseURL()) + baseURL := account.GetGrokBaseURL() + if s.settingService != nil { + baseURL = s.settingService.ResolveGrokBaseURL(ctx, account) + } + validatedBaseURL, err := validator(baseURL) if err != nil { return nil, newUpstreamModelSyncConfigError("Invalid Grok base URL", err) } diff --git a/backend/internal/service/upstream_models_test.go b/backend/internal/service/upstream_models_test.go index 1c8920e71c..031d09b574 100644 --- a/backend/internal/service/upstream_models_test.go +++ b/backend/internal/service/upstream_models_test.go @@ -260,7 +260,7 @@ func TestBuildUpstreamModelsRequestSupportsGrokOAuth(t *testing.T) { require.Equal(t, "Bearer oauth-access-token", req.Header.Get("Authorization")) require.Equal(t, grokCLIVersion, req.Header.Get("X-Grok-Client-Version")) require.Equal(t, "interactive", req.Header.Get("X-Grok-Client-Mode")) - require.Equal(t, grokUpstreamUserAgent, req.Header.Get("User-Agent")) + require.Equal(t, defaultGrokUpstreamUserAgent(), req.Header.Get("User-Agent")) require.Equal(t, "grok-user-id", req.Header.Get("X-UserID")) require.Equal(t, "grok-user@example.com", req.Header.Get("X-Email")) require.NotContains(t, req.Header.Get("Authorization"), "oauth-refresh-token") diff --git a/backend/internal/service/upstream_response_model.go b/backend/internal/service/upstream_response_model.go new file mode 100644 index 0000000000..2452c382d0 --- /dev/null +++ b/backend/internal/service/upstream_response_model.go @@ -0,0 +1,179 @@ +package service + +import ( + "strings" + + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" +) + +const ( + upstreamResponseModelObserverContextKey = "upstream_response_model_observer" + upstreamResponseModelMaxLength = 200 +) + +// upstreamResponseModelObserver tracks one forwarding attempt (or one WS turn). +// A terminal declaration wins over an earlier declaration; otherwise the first +// declaration is retained. Conflicts are diagnostic only and never affect the +// forwarding or billing path. +type upstreamResponseModelObserver struct { + first string + terminal string + conflict bool +} + +func (o *upstreamResponseModelObserver) Observe(model string, terminal bool) { + model = normalizeObservedUpstreamResponseModel(model) + if model == "" { + return + } + current := o.Model() + if current != "" && !strings.EqualFold(current, model) { + o.conflict = true + } + if terminal { + o.terminal = model + return + } + if o.first == "" { + o.first = model + } +} + +func normalizeObservedUpstreamResponseModel(model string) string { + model = strings.TrimSpace(model) + if model == "" { + return "" + } + runes := []rune(model) + if len(runes) > upstreamResponseModelMaxLength { + model = string(runes[:upstreamResponseModelMaxLength]) + } + return model +} + +func (o *upstreamResponseModelObserver) ObserveOpenAI(payload []byte, eventType string) { + if len(payload) == 0 || !gjson.ValidBytes(payload) { + return + } + model := firstTrimmedGJSONModel( + gjson.GetBytes(payload, "response.model"), + gjson.GetBytes(payload, "model"), + ) + o.Observe(model, isUpstreamResponseModelTerminalEvent(eventType)) +} + +func (o *upstreamResponseModelObserver) ObserveAnthropic(payload []byte) { + if len(payload) == 0 || !gjson.ValidBytes(payload) { + return + } + model := firstTrimmedGJSONModel( + gjson.GetBytes(payload, "message.model"), + gjson.GetBytes(payload, "model"), + ) + o.Observe(model, false) +} + +func (o *upstreamResponseModelObserver) ObserveGemini(payload []byte) { + if len(payload) == 0 || !gjson.ValidBytes(payload) { + return + } + model := firstTrimmedGJSONModel( + gjson.GetBytes(payload, "modelVersion"), + gjson.GetBytes(payload, "response.modelVersion"), + ) + // Gemini streaming has no universal terminal event carrying modelVersion; + // treating each declaration as terminal retains the latest chunk. + o.Observe(model, true) +} + +func (o *upstreamResponseModelObserver) Model() string { + if o == nil { + return "" + } + if o.terminal != "" { + return o.terminal + } + return o.first +} + +func (o *upstreamResponseModelObserver) Conflict() bool { + return o != nil && o.conflict +} + +func beginUpstreamResponseModelObservation(c *gin.Context) *upstreamResponseModelObserver { + observer := &upstreamResponseModelObserver{} + if c != nil { + c.Set(upstreamResponseModelObserverContextKey, observer) + } + return observer +} + +func upstreamResponseModelObserverFromContext(c *gin.Context) *upstreamResponseModelObserver { + if c == nil { + return nil + } + value, ok := c.Get(upstreamResponseModelObserverContextKey) + if !ok { + return nil + } + observer, _ := value.(*upstreamResponseModelObserver) + return observer +} + +func observedUpstreamResponseModel(c *gin.Context) string { + return upstreamResponseModelObserverFromContext(c).Model() +} + +func observedUpstreamResponseModelConflict(c *gin.Context) bool { + return upstreamResponseModelObserverFromContext(c).Conflict() +} + +func observeOpenAISSEBody(observer *upstreamResponseModelObserver, body string) { + if observer == nil || strings.TrimSpace(body) == "" { + return + } + forEachOpenAISSEDataPayload(body, func(payload []byte) { + eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String()) + observer.ObserveOpenAI(payload, eventType) + }) +} + +func firstTrimmedGJSONModel(values ...gjson.Result) string { + for _, value := range values { + if !value.Exists() || value.Type != gjson.String { + continue + } + if model := strings.TrimSpace(value.String()); model != "" { + return model + } + } + return "" +} + +func isUpstreamResponseModelTerminalEvent(eventType string) bool { + switch strings.TrimSpace(eventType) { + case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled": + return true + default: + return false + } +} + +func upstreamModelMismatch(sentModel, responseModel string) *bool { + responseModel = strings.TrimSpace(responseModel) + if responseModel == "" { + return nil + } + sentModel = strings.TrimSpace(sentModel) + mismatch := sentModel == "" || !strings.EqualFold(sentModel, responseModel) + return &mismatch +} + +func upstreamSentModel(requestedModel, upstreamModel string) string { + sentModel := strings.TrimSpace(upstreamModel) + if sentModel == "" { + sentModel = strings.TrimSpace(requestedModel) + } + return sentModel +} diff --git a/backend/internal/service/upstream_response_model_test.go b/backend/internal/service/upstream_response_model_test.go new file mode 100644 index 0000000000..37d382e851 --- /dev/null +++ b/backend/internal/service/upstream_response_model_test.go @@ -0,0 +1,75 @@ +package service + +import ( + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestUpstreamResponseModelObserverTerminalWinsAndRecordsConflict(t *testing.T) { + observer := &upstreamResponseModelObserver{} + + observer.ObserveOpenAI([]byte(`{"type":"response.created","response":{"model":"gpt-5.5"}}`), "response.created") + observer.ObserveOpenAI([]byte(`{"type":"response.completed","response":{"model":"gpt-5.4"}}`), "response.completed") + + require.Equal(t, "gpt-5.4", observer.Model()) + require.True(t, observer.Conflict()) +} + +func TestUpstreamResponseModelObserverSupportsAnthropicAndGeminiShapes(t *testing.T) { + t.Run("anthropic", func(t *testing.T) { + observer := &upstreamResponseModelObserver{} + observer.ObserveAnthropic([]byte(`{"type":"message_start","message":{"model":"claude-sonnet-4-20250514"}}`)) + require.Equal(t, "claude-sonnet-4-20250514", observer.Model()) + }) + + t.Run("gemini outer and nested", func(t *testing.T) { + observer := &upstreamResponseModelObserver{} + observer.ObserveGemini([]byte(`{"response":{"modelVersion":"gemini-2.5-pro"}}`)) + observer.ObserveGemini([]byte(`{"modelVersion":"gemini-2.5-pro-latest"}`)) + require.Equal(t, "gemini-2.5-pro-latest", observer.Model()) + require.True(t, observer.Conflict()) + }) +} + +func TestUpstreamResponseModelObservationAttemptReset(t *testing.T) { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(nil) + + first := beginUpstreamResponseModelObservation(c) + first.Observe("failed-attempt-model", false) + second := beginUpstreamResponseModelObservation(c) + second.Observe("successful-attempt-model", false) + + require.Equal(t, "successful-attempt-model", observedUpstreamResponseModel(c)) + require.False(t, observedUpstreamResponseModelConflict(c)) +} + +func TestUpstreamModelMismatchThreeStateAndCaseInsensitiveComparison(t *testing.T) { + require.Nil(t, upstreamModelMismatch("gpt-5.5", "")) + + matched := upstreamModelMismatch("gpt-5.5", "GPT-5.5") + require.NotNil(t, matched) + require.False(t, *matched) + + mismatched := upstreamModelMismatch("gpt-5.5", "gpt-5.4") + require.NotNil(t, mismatched) + require.True(t, *mismatched) +} + +func TestObserveOpenAISSEBodyIgnoresMalformedPayload(t *testing.T) { + observer := &upstreamResponseModelObserver{} + observeOpenAISSEBody(observer, "data: not-json\n\ndata: {\"type\":\"response.completed\",\"response\":{\"model\":\"gpt-5.4\"}}\n\n") + + require.Equal(t, "gpt-5.4", observer.Model()) + require.False(t, observer.Conflict()) +} + +func TestUpstreamResponseModelObserverBoundsUntrustedModelName(t *testing.T) { + observer := &upstreamResponseModelObserver{} + observer.Observe(" "+strings.Repeat("模", upstreamResponseModelMaxLength+1)+" ", false) + + require.Len(t, []rune(observer.Model()), upstreamResponseModelMaxLength) +} diff --git a/backend/internal/service/usage_log.go b/backend/internal/service/usage_log.go index b35d149ea8..7a41555ece 100644 --- a/backend/internal/service/usage_log.go +++ b/backend/internal/service/usage_log.go @@ -114,6 +114,12 @@ type UsageLog struct { // UpstreamModel is the actual model sent to the upstream provider after mapping. // Nil means no mapping was applied (requested model was used as-is). UpstreamModel *string + // UpstreamResponseModel is the model declared by the successful upstream + // response before client-facing model rewrites or protocol conversion. + UpstreamResponseModel *string + // UpstreamModelMismatch is nil when no upstream model was observed. Otherwise + // it compares UpstreamResponseModel with the actual model sent upstream. + UpstreamModelMismatch *bool // ChannelID 渠道 ID ChannelID *int64 // ModelMappingChain 模型映射链,如 "a→b→c" diff --git a/backend/internal/service/user_subscription.go b/backend/internal/service/user_subscription.go index 3a73554a4b..fb4d77f0f1 100644 --- a/backend/internal/service/user_subscription.go +++ b/backend/internal/service/user_subscription.go @@ -1,6 +1,10 @@ package service -import "time" +import ( + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" +) const subscriptionDayDuration = 24 * time.Hour @@ -75,13 +79,8 @@ func (s *UserSubscription) NeedsDailyReset() bool { } func (s *UserSubscription) NeedsDailyResetAt(now time.Time) bool { - if s.DailyWindowStart == nil { - return false - } - if s.HasOneTimeDailyQuota() { - return false - } - return !now.Before(s.DailyWindowStart.Add(24 * time.Hour)) + _, ok := s.automaticDailyWindowStartAt(now) + return ok } func (s *UserSubscription) NeedsWeeklyReset() bool { @@ -107,8 +106,26 @@ func (s *UserSubscription) NeedsMonthlyResetAt(now time.Time) bool { } func (s *UserSubscription) canAutomaticallyResetDailyAt(now time.Time) bool { - _, ok := s.automaticWindowStartAt(s.DailyWindowStart, 24*time.Hour, now) - return !s.HasOneTimeDailyQuota() && ok + _, ok := s.automaticDailyWindowStartAt(now) + return ok +} + +// automaticDailyWindowStartAt 计算日窗口按“配置时区日历日”对齐后的当前窗口起点。 +// 日额度固定在每天 0 点刷新(与周/月的期限对齐滚动窗口语义不同),因此只要持久化 +// 的窗口起点落在更早的日历日,就允许推进到今天 0 点。手动重置、激活等写入的任何 +// 非 0 点锚点都会在下一个 0 点被拉回日历日边界,不会永久漂移刷新时刻。 +func (s *UserSubscription) automaticDailyWindowStartAt(now time.Time) (time.Time, bool) { + if s.DailyWindowStart == nil { + return time.Time{}, false + } + if s.HasOneTimeDailyQuota() { + return time.Time{}, false + } + today := timezone.StartOfDay(now) + if !today.After(timezone.StartOfDay(*s.DailyWindowStart)) { + return time.Time{}, false + } + return today, true } func (s *UserSubscription) canAutomaticallyResetWeeklyAt(now time.Time) bool { @@ -121,6 +138,9 @@ func (s *UserSubscription) canAutomaticallyResetMonthlyAt(now time.Time) bool { return ok } +// automaticWindowStartAt 计算周/月窗口(期限对齐滚动窗口)的当前窗口起点。 +// 窗口从锚点按整数个 period 步进,且不越过订阅到期时间,避免最后一个不完整 +// 周期重复发放额度(issue #5051)。日窗口不走此函数,见 automaticDailyWindowStartAt。 func (s *UserSubscription) automaticWindowStartAt(previous *time.Time, period time.Duration, now time.Time) (time.Time, bool) { if previous == nil { return time.Time{}, false @@ -155,7 +175,8 @@ func (s *UserSubscription) DailyResetTime() *time.Time { t := s.ExpiresAt return &t } - t := s.DailyWindowStart.Add(24 * time.Hour) + // 日窗口按日历日对齐:下次刷新固定在窗口起点所在日的次日 0 点。 + t := timezone.StartOfDay(*s.DailyWindowStart).AddDate(0, 0, 1) return &t } diff --git a/backend/internal/service/user_subscription_daily_quota_test.go b/backend/internal/service/user_subscription_daily_quota_test.go index 6e62eb3539..43e9799333 100644 --- a/backend/internal/service/user_subscription_daily_quota_test.go +++ b/backend/internal/service/user_subscription_daily_quota_test.go @@ -6,6 +6,7 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/pkg/timezone" "github.com/stretchr/testify/require" ) @@ -58,7 +59,7 @@ func TestAssignOrExtendSubscription_ExpiredDailyCardStartsNewOneTimeQuota(t *tes require.True(t, renewed.StartsAt.After(oldStart), "重新购买过期订阅时应重置当前周期 StartsAt") require.False(t, renewed.ExpiresAt.After(renewed.StartsAt.AddDate(0, 0, 1))) require.NotNil(t, renewed.DailyWindowStart) - require.Equal(t, renewed.StartsAt, *renewed.DailyWindowStart) + require.Equal(t, timezone.StartOfDay(renewed.StartsAt), *renewed.DailyWindowStart, "续期后日窗口应锚定当天 0 点") require.Equal(t, 0.0, renewed.DailyUsageUSD) require.Equal(t, 0.0, renewed.WeeklyUsageUSD) require.Equal(t, 0.0, renewed.MonthlyUsageUSD) diff --git a/backend/internal/service/user_subscription_port.go b/backend/internal/service/user_subscription_port.go index 7ce79f4b27..a20e4a4f26 100644 --- a/backend/internal/service/user_subscription_port.go +++ b/backend/internal/service/user_subscription_port.go @@ -29,8 +29,13 @@ type UserSubscriptionRepository interface { UpdateStatus(ctx context.Context, subscriptionID int64, status string) error UpdateNotes(ctx context.Context, subscriptionID int64, notes string) error - ActivateWindows(ctx context.Context, id int64, start time.Time) error - ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, newWindowStart time.Time) error + // ActivateWindows 首次使用时激活用量窗口。日窗口按日历日对齐,锚点为当天 0 点 + // (dailyStart);周/月窗口为期限对齐滚动窗口,锚点为激活时刻(periodicStart)。 + // 仅当三个窗口均未激活时生效。 + ActivateWindows(ctx context.Context, id int64, dailyStart, periodicStart time.Time) error + // ResetUsageWindows 手动重置所选窗口的用量。日窗口锚点写入 dailyStart(当天 0 点, + // 保持 0 点刷新节奏不漂移);周/月窗口锚点写入 periodicStart(重置时刻)。 + ResetUsageWindows(ctx context.Context, id int64, resetDaily, resetWeekly, resetMonthly bool, dailyStart, periodicStart time.Time) error ResetDailyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error ResetWeeklyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error ResetMonthlyUsage(ctx context.Context, id int64, expectedWindowStart *time.Time, newWindowStart time.Time) error diff --git a/backend/internal/service/video_billing.go b/backend/internal/service/video_billing.go new file mode 100644 index 0000000000..66a5b950f6 --- /dev/null +++ b/backend/internal/service/video_billing.go @@ -0,0 +1,155 @@ +package service + +import ( + "log/slog" + "sort" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" +) + +// Canonical video price family keys used in groups.video_model_prices JSONB. +const ( + VideoPriceFamilyGrokImagineVideo = "grok-imagine-video" + VideoPriceFamilyGrokImagineVideo15 = "grok-imagine-video-1.5" +) + +// CanonicalGrokImagineVideoPriceFamily normalizes model aliases / preview / legacy +// IDs onto the price-family keys stored in video_model_prices. +func CanonicalGrokImagineVideoPriceFamily(model string) string { + if model == "" { + return "" + } + // Prefer shared xAI helper for known aliases. Keep future native Imagine + // models distinct so operators can assign them independent prices. + if c := xai.CanonicalImagineVideoModel(model); c != "" { + switch c { + case xai.DefaultImagineVideo15Model: + return VideoPriceFamilyGrokImagineVideo15 + case xai.DefaultImagineVideoModel: + return VideoPriceFamilyGrokImagineVideo + } + if strings.HasPrefix(c, "grok-imagine-video-") { + return c + } + } + m := strings.ToLower(strings.TrimSpace(model)) + for _, prefix := range []string{"xai/", "x-ai/", "grok/"} { + if strings.HasPrefix(m, prefix) { + m = strings.TrimPrefix(m, prefix) + break + } + } + switch { + case m == "grok-imagine-video-1.5" || m == "grok-imagine-video-1.5-preview" || + m == "grok-video-1.5" || strings.Contains(m, "video-1.5"): + return VideoPriceFamilyGrokImagineVideo15 + case m == "grok-imagine-video" || m == "grok-imagine-video-preview" || + m == "grok-video" || m == "grok-video-latest": + return VideoPriceFamilyGrokImagineVideo + default: + return "" + } +} + +// NormalizeVideoModelPrices cleans and canonicalizes a per-model resolution map. +// Keys become price families; tiers use 480p/720p/1080p. Negative prices dropped. +// +// Model keys are walked in sorted order rather than in Go map order: several +// aliases can canonicalize onto the same family, and an unordered walk would +// make the winning price for a conflicting tier vary between processes. +// Unrecognized tiers are dropped with a warning instead of silently collapsing +// into the 480p bucket. +func NormalizeVideoModelPrices(in map[string]map[string]float64) map[string]map[string]float64 { + if len(in) == 0 { + return nil + } + modelKeys := make([]string, 0, len(in)) + for modelKey := range in { + modelKeys = append(modelKeys, modelKey) + } + sort.Strings(modelKeys) + out := make(map[string]map[string]float64) + for _, modelKey := range modelKeys { + tierPrices := in[modelKey] + if len(tierPrices) == 0 { + continue + } + family := CanonicalGrokImagineVideoPriceFamily(modelKey) + if family == "" { + key := strings.ToLower(strings.TrimSpace(modelKey)) + switch key { + case VideoPriceFamilyGrokImagineVideo, VideoPriceFamilyGrokImagineVideo15: + family = key + default: + if key == "" { + continue + } + family = key + } + } + normalizedTiers := out[family] + if normalizedTiers == nil { + normalizedTiers = make(map[string]float64) + } + tierKeys := make([]string, 0, len(tierPrices)) + for tierKey := range tierPrices { + tierKeys = append(tierKeys, tierKey) + } + sort.Strings(tierKeys) + for _, tierKey := range tierKeys { + price := tierPrices[tierKey] + if price < 0 { + continue + } + tier, ok := LookupVideoBillingResolution(tierKey) + if !ok { + slog.Warn("video_model_prices_unknown_resolution_dropped", + "model_key", modelKey, + "family", family, + "resolution", tierKey) + continue + } + if existing, exists := normalizedTiers[tier]; exists && existing != price { + slog.Warn("video_model_prices_conflicting_tier_price", + "model_key", modelKey, + "family", family, + "resolution", tier, + "previous_price", existing, + "price", price) + } + normalizedTiers[tier] = price + } + if len(normalizedTiers) > 0 { + out[family] = normalizedTiers + } + } + if len(out) == 0 { + return nil + } + return out +} + +// LookupVideoModelPrice returns a per-second price from a model×resolution map, or nil. +func LookupVideoModelPrice(prices map[string]map[string]float64, model, resolution string) *float64 { + if len(prices) == 0 { + return nil + } + family := CanonicalGrokImagineVideoPriceFamily(model) + if family == "" { + family = strings.ToLower(strings.TrimSpace(model)) + } + if family == "" { + return nil + } + tierPrices, ok := prices[family] + if !ok || len(tierPrices) == 0 { + return nil + } + tier := NormalizeVideoBillingResolutionOrDefault(resolution) + if price, ok := tierPrices[tier]; ok { + p := price + return &p + } + return nil +} diff --git a/backend/internal/service/video_billing_resolution.go b/backend/internal/service/video_billing_resolution.go index cb713f6877..11782db851 100644 --- a/backend/internal/service/video_billing_resolution.go +++ b/backend/internal/service/video_billing_resolution.go @@ -31,15 +31,27 @@ func NormalizeVideoBillingDurationSecondsOrDefault(durationSeconds int) int { return durationSeconds } -func NormalizeVideoBillingResolutionOrDefault(resolution string) string { +// LookupVideoBillingResolution 归一化分辨率并报告是否为已知档位。 +// 配置解析路径必须用它而不是 OrDefault:把无法识别的档位(如 "4k"、拼错的 +// "1080i")静默折算成 480p,会让管理员配的高分辨率单价被挂到低分辨率档上。 +func LookupVideoBillingResolution(resolution string) (string, bool) { switch strings.ToLower(strings.TrimSpace(resolution)) { case "480", "480p", "sd": - return VideoBillingResolution480P + return VideoBillingResolution480P, true case "720", "720p", "hd": - return VideoBillingResolution720P + return VideoBillingResolution720P, true case "1080", "1080p", "full_hd", "full-hd", "fhd": - return VideoBillingResolution1080P + return VideoBillingResolution1080P, true default: - return VideoBillingResolution480P + return "", false } } + +// NormalizeVideoBillingResolutionOrDefault 用于运行时计费:上游回传的分辨率 +// 缺失或无法识别时按最低档兜底,保证请求仍可计费。 +func NormalizeVideoBillingResolutionOrDefault(resolution string) string { + if normalized, ok := LookupVideoBillingResolution(resolution); ok { + return normalized + } + return VideoBillingResolution480P +} diff --git a/backend/internal/service/video_billing_test.go b/backend/internal/service/video_billing_test.go new file mode 100644 index 0000000000..6cf2e936c8 --- /dev/null +++ b/backend/internal/service/video_billing_test.go @@ -0,0 +1,124 @@ +package service + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestCanonicalGrokImagineVideoPriceFamily(t *testing.T) { + t.Parallel() + require.Equal(t, VideoPriceFamilyGrokImagineVideo, CanonicalGrokImagineVideoPriceFamily("grok-imagine-video")) + require.Equal(t, VideoPriceFamilyGrokImagineVideo15, CanonicalGrokImagineVideoPriceFamily("grok-imagine-video-1.5")) + require.Equal(t, VideoPriceFamilyGrokImagineVideo15, CanonicalGrokImagineVideoPriceFamily("grok-imagine-video-1.5-preview")) + require.Equal(t, VideoPriceFamilyGrokImagineVideo15, CanonicalGrokImagineVideoPriceFamily("xai/grok-video-1.5")) + require.Equal(t, "grok-imagine-video-2", CanonicalGrokImagineVideoPriceFamily("grok-imagine-video-2")) + require.Equal(t, "grok-imagine-video-2", CanonicalGrokImagineVideoPriceFamily("xai/grok-imagine-video-2")) +} + +func TestNormalizeAndLookupVideoModelPrices(t *testing.T) { + t.Parallel() + raw := map[string]map[string]float64{ + "grok-imagine-video-1.5-preview": {"480p": 0.08, "720p": 0.14}, + "grok-imagine-video": {"480p": 0.05}, + "grok-imagine-video-2": {"1080p": 0.4}, + } + norm := NormalizeVideoModelPrices(raw) + require.NotNil(t, norm) + require.Contains(t, norm, VideoPriceFamilyGrokImagineVideo15) + require.Contains(t, norm, VideoPriceFamilyGrokImagineVideo) + require.Contains(t, norm, "grok-imagine-video-2") + + p15 := LookupVideoModelPrice(norm, "grok-imagine-video-1.5", "480p") + require.NotNil(t, p15) + require.InDelta(t, 0.08, *p15, 1e-9) + + pBase := LookupVideoModelPrice(norm, "grok-imagine-video", "480p") + require.NotNil(t, pBase) + require.InDelta(t, 0.05, *pBase, 1e-9) + // A missing model-specific tier must fall back to the flat tier price, + // rather than borrowing another model-specific resolution. + require.Nil(t, LookupVideoModelPrice(norm, "grok-imagine-video", "720p")) + + p2 := LookupVideoModelPrice(norm, "grok-imagine-video-2", "1080p") + require.NotNil(t, p2) + require.InDelta(t, 0.4, *p2, 1e-9) + + // Unmatched model → nil (caller falls back to flat columns / defaults). + require.Nil(t, LookupVideoModelPrice(norm, "unknown-model", "480p")) +} + +func TestNormalizeVideoModelPricesDropsUnknownResolutions(t *testing.T) { + t.Parallel() + // "4k" and "1080i" are not billable tiers. Collapsing them into 480p would + // charge a 480p request at the operator's high-resolution price. + norm := NormalizeVideoModelPrices(map[string]map[string]float64{ + "grok-imagine-video": {"480p": 0.05, "4k": 0.50, "1080i": 0.30}, + }) + require.NotNil(t, norm) + require.Equal(t, map[string]float64{VideoBillingResolution480P: 0.05}, norm[VideoPriceFamilyGrokImagineVideo]) + + // A model whose tiers are all unrecognized contributes no family at all. + require.Nil(t, NormalizeVideoModelPrices(map[string]map[string]float64{ + "grok-imagine-video": {"4k": 0.50}, + })) +} + +func TestNormalizeVideoModelPricesIsDeterministicAcrossAliasConflicts(t *testing.T) { + t.Parallel() + // Both keys canonicalize onto grok-imagine-video-1.5 and disagree on 480p. + // Whichever price wins, it must be the same one on every run — a Go map walk + // would let two processes bill the same request differently. + raw := map[string]map[string]float64{ + "grok-imagine-video-1.5": {"480p": 0.08}, + "grok-imagine-video-1.5-preview": {"480p": 0.11}, + "grok-video-1.5": {"480p": 0.09}, + } + first := NormalizeVideoModelPrices(raw) + require.NotNil(t, first) + for i := 0; i < 50; i++ { + require.Equal(t, first, NormalizeVideoModelPrices(raw), "run %d diverged", i) + } + + // Aliases within one model key normalize to the same tier deterministically too. + aliased := map[string]map[string]float64{ + "grok-imagine-video": {"720": 0.12, "720p": 0.13, "hd": 0.14}, + } + firstAliased := NormalizeVideoModelPrices(aliased) + require.NotNil(t, firstAliased) + for i := 0; i < 50; i++ { + require.Equal(t, firstAliased, NormalizeVideoModelPrices(aliased), "run %d diverged", i) + } +} + +func TestLookupVideoBillingResolutionReportsUnknownTiers(t *testing.T) { + t.Parallel() + for _, in := range []string{"480", "480p", "SD", "720", "hd", "1080", "full-hd", " fhd "} { + normalized, ok := LookupVideoBillingResolution(in) + require.True(t, ok, "input=%q", in) + require.NotEmpty(t, normalized) + } + for _, in := range []string{"", "4k", "1080i", "2160p", "potato"} { + normalized, ok := LookupVideoBillingResolution(in) + require.False(t, ok, "input=%q", in) + require.Empty(t, normalized) + } + // Runtime billing still needs a tier for unrecognized upstream values. + require.Equal(t, VideoBillingResolution480P, NormalizeVideoBillingResolutionOrDefault("4k")) + require.Equal(t, VideoBillingResolution1080P, NormalizeVideoBillingResolutionOrDefault("full_hd")) +} + +func TestVideoModelPriceMissingTierFallsBackToFlatTierPrice(t *testing.T) { + t.Parallel() + flat720P := 0.7 + service := &BillingService{} + + result := service.CalculateVideoCost("grok-imagine-video", "720p", 1, 1, &VideoPriceConfig{ + Price720P: &flat720P, + ModelPrices: map[string]map[string]float64{ + VideoPriceFamilyGrokImagineVideo: {VideoBillingResolution480P: 0.05}, + }, + }, 1) + + require.InDelta(t, flat720P, result.TotalCost, 1e-9) +} diff --git a/backend/internal/service/websearch_config.go b/backend/internal/service/websearch_config.go index bb33368d39..d4b0823adc 100644 --- a/backend/internal/service/websearch_config.go +++ b/backend/internal/service/websearch_config.go @@ -3,6 +3,7 @@ package service import ( "context" "encoding/json" + "errors" "fmt" "log/slog" "sync/atomic" @@ -106,6 +107,15 @@ func (s *SettingService) loadWebSearchConfigFromDB() (*WebSearchEmulationConfig, raw, err := s.settingRepo.GetValue(dbCtx, SettingKeyWebSearchEmulationConfig) if err != nil { + // Missing key is the normal first-boot state: return empty disabled config. + if errors.Is(err, ErrSettingNotFound) { + cfg := &WebSearchEmulationConfig{} + webSearchEmulationCache.Store(&cachedWebSearchEmulationConfig{ + config: cfg, + expiresAt: time.Now().Add(webSearchEmulationCacheTTL).UnixNano(), + }) + return cfg, nil + } webSearchEmulationCache.Store(&cachedWebSearchEmulationConfig{ config: &WebSearchEmulationConfig{}, expiresAt: time.Now().Add(webSearchEmulationErrorTTL).UnixNano(), diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 5b112a539a..93555730a7 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -11,11 +11,21 @@ import ( "github.com/Wei-Shaw/sub2api/internal/payment" "github.com/Wei-Shaw/sub2api/internal/pkg/antigravity" "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + "github.com/Wei-Shaw/sub2api/internal/pkg/xai" "github.com/google/wire" "github.com/redis/go-redis/v9" "go.uber.org/zap" ) +func ProvideGrokOAuthService(proxyRepo ProxyRepository, oauthClient GrokOAuthClient, cfg *config.Config, redisClient *redis.Client) *GrokOAuthService { + svc := NewGrokOAuthService(proxyRepo, oauthClient, cfg) + // wire.go is depguard-exempt for redis; construct the Redis session store here. + if redisClient != nil { + svc = svc.WithSessionStore(xai.NewRedisSessionStore(redisClient)) + } + return svc +} + // BuildInfo contains build information type BuildInfo struct { Version string @@ -220,6 +230,7 @@ func ProvideAccountTestService( cfg *config.Config, tlsFPProfileService *TLSFingerprintProfileService, openAIGatewayService *OpenAIGatewayService, + settingService *SettingService, ) *AccountTestService { service := NewAccountTestService( accountRepo, @@ -232,6 +243,7 @@ func ProvideAccountTestService( tlsFPProfileService, ) service.agentIdentityWS = openAIGatewayService + service.SetSettingService(settingService) return service } @@ -242,8 +254,11 @@ func ProvideGrokQuotaService( httpUpstream HTTPUpstream, cfg *config.Config, usageLogRepo UsageLogRepository, + settingService *SettingService, ) *GrokQuotaService { - return NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream, cfg, usageLogRepo) + service := NewGrokQuotaService(accountRepo, proxyRepo, tokenProvider, httpUpstream, cfg, usageLogRepo) + service.SetSettingService(settingService) + return service } // ProvideGeminiTokenProvider creates GeminiTokenProvider with OAuthRefreshAPI injection @@ -763,7 +778,7 @@ var ProviderSet = wire.NewSet( wire.Bind(new(AccountRuntimeBlocker), new(*OpenAIGatewayService)), NewOAuthService, ProvideOpenAIOAuthService, - NewGrokOAuthService, + ProvideGrokOAuthService, wire.Bind(new(GrokOAuthTokenService), new(*GrokOAuthService)), NewGeminiOAuthService, NewGeminiQuotaService, diff --git a/backend/internal/testutil/stubs.go b/backend/internal/testutil/stubs.go index a0bd4bc669..9890ec14a5 100644 --- a/backend/internal/testutil/stubs.go +++ b/backend/internal/testutil/stubs.go @@ -104,6 +104,20 @@ func (c StubGatewayCache) DeleteSessionAccountID(_ context.Context, _ int64, _ s return nil } +func (c StubGatewayCache) SetGrokVideoPendingBilling(_ context.Context, _ string, _ []byte, _ time.Duration) error { + return nil +} +func (c StubGatewayCache) GetGrokVideoPendingBilling(_ context.Context, _ string) ([]byte, error) { + return nil, nil +} +func (c StubGatewayCache) ClaimGrokVideoBilled(_ context.Context, _ string, _ time.Duration) (bool, error) { + return true, nil +} + +func (c StubGatewayCache) ReleaseGrokVideoBilled(_ context.Context, _ string) error { + return nil +} + // ============================================================ // StubSessionLimitCache — service.SessionLimitCache 的空实现 // ============================================================ diff --git a/backend/migrations/194_add_usage_log_upstream_response_model.sql b/backend/migrations/194_add_usage_log_upstream_response_model.sql new file mode 100644 index 0000000000..a5865aca17 --- /dev/null +++ b/backend/migrations/194_add_usage_log_upstream_response_model.sql @@ -0,0 +1,3 @@ +ALTER TABLE usage_logs + ADD COLUMN IF NOT EXISTS upstream_response_model VARCHAR(200), + ADD COLUMN IF NOT EXISTS upstream_model_mismatch BOOLEAN; diff --git a/backend/migrations/195_add_usage_log_upstream_model_mismatch_index_notx.sql b/backend/migrations/195_add_usage_log_upstream_model_mismatch_index_notx.sql new file mode 100644 index 0000000000..811ca8786c --- /dev/null +++ b/backend/migrations/195_add_usage_log_upstream_model_mismatch_index_notx.sql @@ -0,0 +1,3 @@ +CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_usage_logs_upstream_model_mismatch_created_at + ON usage_logs (created_at DESC, id DESC) + WHERE upstream_model_mismatch IS TRUE; diff --git a/backend/migrations/217_group_video_model_prices.sql b/backend/migrations/217_group_video_model_prices.sql new file mode 100644 index 0000000000..61080015d4 --- /dev/null +++ b/backend/migrations/217_group_video_model_prices.sql @@ -0,0 +1,8 @@ +-- Per-model-family video per-second prices for Grok Imagine. +-- Shape: {"grok-imagine-video":{"480p":0.05,"720p":0.07},"grok-imagine-video-1.5":{"480p":0.08,"720p":0.14,"1080p":0.25}} +-- Resolution order in billing: per-model map → legacy video_price_* columns → code defaults (model-aware). +ALTER TABLE groups + ADD COLUMN IF NOT EXISTS video_model_prices JSONB; + +COMMENT ON COLUMN groups.video_model_prices IS + '可选:按模型族×分辨率覆盖视频每秒单价 (USD/s)。key 为规范模型族 (grok-imagine-video / grok-imagine-video-1.5),value 为分辨率→单价映射;NULL/空表示不覆盖,回退到 video_price_* 列或官方默认'; diff --git a/backend/migrations/218_group_audio_voice_pricing.sql b/backend/migrations/218_group_audio_voice_pricing.sql new file mode 100644 index 0000000000..c1b8652265 --- /dev/null +++ b/backend/migrations/218_group_audio_voice_pricing.sql @@ -0,0 +1,5 @@ +-- Grok Voice 显式定价:realtime / TTS / STT。 +-- NULL = 使用代码默认单价;显式 0 = 免费;>0 = 分组覆盖价。 +ALTER TABLE groups ADD COLUMN IF NOT EXISTS audio_realtime_price_per_min DECIMAL(20,8); +ALTER TABLE groups ADD COLUMN IF NOT EXISTS audio_tts_price_per_million_chars DECIMAL(20,8); +ALTER TABLE groups ADD COLUMN IF NOT EXISTS audio_stt_price_per_hour DECIMAL(20,8); diff --git a/backend/migrations/219_group_search_price_per_1k.sql b/backend/migrations/219_group_search_price_per_1k.sql new file mode 100644 index 0000000000..54b6ef81f4 --- /dev/null +++ b/backend/migrations/219_group_search_price_per_1k.sql @@ -0,0 +1,3 @@ +-- Grok / 通用搜索工具显式定价(per 1000 calls,USD)。 +-- NULL = 使用代码默认 $10/1k;显式 0 = 免费;>0 = 分组覆盖价。 +ALTER TABLE groups ADD COLUMN IF NOT EXISTS search_price_per_1k DECIMAL(20,8); diff --git a/backend/migrations/220_clear_non_grok_video_generation_config.sql b/backend/migrations/220_clear_non_grok_video_generation_config.sql new file mode 100644 index 0000000000..05571d3709 --- /dev/null +++ b/backend/migrations/220_clear_non_grok_video_generation_config.sql @@ -0,0 +1,42 @@ +-- Videos are Grok/xAI-only. Clear stale video pricing from non-Grok groups. +-- Columns match migrations 170/217 (video_price_* / video_model_prices), not a +-- separate allow_video_generation flag which was never applied on this branch. + +-- Snapshot before clearing. The UPDATE below is irreversible, and an operator +-- who deliberately priced video on a non-Grok group would otherwise lose that +-- configuration with no way to recover it. CREATE TABLE IF NOT EXISTS ... AS +-- SELECT is a no-op on re-run, so this stays idempotent. +CREATE TABLE IF NOT EXISTS groups_video_price_backup_220 AS +SELECT id AS group_id, + platform, + video_price_480p, + video_price_720p, + video_price_1080p, + video_model_prices, + now() AS backed_up_at +FROM groups +WHERE platform IS DISTINCT FROM 'grok' + AND platform IS DISTINCT FROM 'composite' + AND ( + video_price_480p IS NOT NULL + OR video_price_720p IS NOT NULL + OR video_price_1080p IS NOT NULL + OR video_model_prices IS NOT NULL + ); + +COMMENT ON TABLE groups_video_price_backup_220 IS + '迁移 220 清空非 Grok/非 composite 分组视频价前的快照。composite 可能路由到 Grok 账号,予以保留。确认无需回滚后可安全 DROP;回滚方式:UPDATE groups g SET video_price_480p = b.video_price_480p, ... FROM groups_video_price_backup_220 b WHERE g.id = b.group_id'; + +UPDATE groups +SET video_price_480p = NULL, + video_price_720p = NULL, + video_price_1080p = NULL, + video_model_prices = NULL +WHERE platform IS DISTINCT FROM 'grok' + AND platform IS DISTINCT FROM 'composite' + AND ( + video_price_480p IS NOT NULL + OR video_price_720p IS NOT NULL + OR video_price_1080p IS NOT NULL + OR video_model_prices IS NOT NULL + ); diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml index a6ef8a9813..095c550993 100644 --- a/deploy/config.example.yaml +++ b/deploy/config.example.yaml @@ -440,6 +440,21 @@ gateway: failure_threshold: 2 window_seconds: 60 ttl_seconds: 600 + # Grok free-tier local soft gate (scheduler filter only; admin QueryQuota/import probe bypasses it). + # Enabled by default because free detection requires an explicit subscription_tier/plan_type of "free". + # Stats/query failures fail open so DB issues do not block all Grok traffic. + grok: + # Email/password OAuth is off by default and hidden in the admin UI. + # Setting true enables POST /admin/grok/oauth/password (password → SSO → Build OAuth). + # Prefer SSO cookie, browser OAuth, or refresh_token re-auth in production. + password_auth_enabled: false + free_quota_soft_gate_enabled: true + free_quota_token_limit: 500000 + free_quota_soft_gate_percent: 95 + free_quota_window_hours: 24 + # Stats cache for free-tier soft gate. Hot path never waits on DB: misses + # fail open and refresh in the background. Prefer >= 60s in production. + free_quota_stats_cache_seconds: 60 # HTTP upstream connection pool settings (HTTP/2 + multi-proxy scenario defaults) # HTTP 上游连接池配置(HTTP/2 + 多代理场景默认值) # Max idle connections across all hosts diff --git a/frontend/src/api/__tests__/admin.grok.spec.ts b/frontend/src/api/__tests__/admin.grok.spec.ts index 7560ce443f..87b0a7f7e6 100644 --- a/frontend/src/api/__tests__/admin.grok.spec.ts +++ b/frontend/src/api/__tests__/admin.grok.spec.ts @@ -8,7 +8,7 @@ vi.mock('@/api/client', () => ({ apiClient: { post }, })) -import { createFromSSO, getGrokSSOImportTimeout } from '@/api/admin/grok' +import { authorizePassword, createFromSSO, getGrokSSOImportTimeout } from '@/api/admin/grok' describe('admin Grok SSO import API', () => { beforeEach(() => { @@ -34,4 +34,20 @@ describe('admin Grok SSO import API', () => { { timeout: expectedTimeout }, ) }) + + it('preserves password whitespace and applies the authorization timeout', async () => { + post.mockResolvedValueOnce({ data: { access_token: 'access-token' } }) + + await authorizePassword(' user@example.com ---- password with spaces ', 7) + + expect(post).toHaveBeenCalledWith( + '/admin/grok/oauth/password', + { + email: 'user@example.com', + password: ' password with spaces ', + proxy_id: 7, + }, + { timeout: 120_000 }, + ) + }) }) diff --git a/frontend/src/api/__tests__/admin.system.rollback.spec.ts b/frontend/src/api/__tests__/admin.system.rollback.spec.ts index 15d0989ffb..09eab12b65 100644 --- a/frontend/src/api/__tests__/admin.system.rollback.spec.ts +++ b/frontend/src/api/__tests__/admin.system.rollback.spec.ts @@ -41,7 +41,11 @@ describe('admin system rollback API', () => { const result = await rollback('0.1.146') - expect(post).toHaveBeenCalledWith('/admin/system/rollback', { version: '0.1.146' }) + expect(post).toHaveBeenCalledWith( + '/admin/system/rollback', + { version: '0.1.146' }, + { timeout: 15 * 60 * 1000 } + ) expect(result.need_restart).toBe(true) }) @@ -50,6 +54,10 @@ describe('admin system rollback API', () => { await rollback() - expect(post).toHaveBeenCalledWith('/admin/system/rollback', undefined) + expect(post).toHaveBeenCalledWith( + '/admin/system/rollback', + undefined, + { timeout: 15 * 60 * 1000 } + ) }) }) diff --git a/frontend/src/api/admin/accounts.ts b/frontend/src/api/admin/accounts.ts index bd9c54f8e1..21ce9f092f 100644 --- a/frontend/src/api/admin/accounts.ts +++ b/frontend/src/api/admin/accounts.ts @@ -317,6 +317,19 @@ export async function getUsage(id: number, source?: 'passive' | 'active', force? return data } +export interface BatchAccountUsageResponse { + usage: Record + errors: Record +} + +export async function getBatchUsage(accountIds: number[], force?: boolean): Promise { + const { data } = await apiClient.post('/admin/accounts/usage/batch', { + account_ids: accountIds, + force: force === true + }) + return data +} + /** * Clear account rate limit status * @param id - Account ID @@ -986,6 +999,7 @@ export const accountsAPI = { getStats, clearError, getUsage, + getBatchUsage, getTodayStats, getBatchTodayStats, clearRateLimit, diff --git a/frontend/src/api/admin/dashboard.ts b/frontend/src/api/admin/dashboard.ts index ae20d33f8e..5117039f61 100644 --- a/frontend/src/api/admin/dashboard.ts +++ b/frontend/src/api/admin/dashboard.ts @@ -56,6 +56,7 @@ export interface TrendParams { request_type?: UsageRequestType stream?: boolean billing_type?: number | null + upstream_model_mismatch?: boolean } export interface TrendResponse { @@ -87,6 +88,7 @@ export interface ModelStatsParams { request_type?: UsageRequestType stream?: boolean billing_type?: number | null + upstream_model_mismatch?: boolean } export interface ModelStatsResponse { @@ -115,6 +117,7 @@ export interface GroupStatsParams { request_type?: UsageRequestType stream?: boolean billing_type?: number | null + upstream_model_mismatch?: boolean } export interface GroupStatsResponse { diff --git a/frontend/src/api/admin/grok.ts b/frontend/src/api/admin/grok.ts index f59c05c1d7..c1dc1cbaa2 100644 --- a/frontend/src/api/admin/grok.ts +++ b/frontend/src/api/admin/grok.ts @@ -19,6 +19,17 @@ export interface GrokAuthUrlRequest { redirect_uri?: string } +export interface GrokOAuthCapabilities { + password_auth_enabled: boolean +} + +const GROK_AUTHORIZATION_TIMEOUT_MS = 120_000 + +export async function getCapabilities(): Promise { + const { data } = await apiClient.get('/admin/grok/oauth/capabilities') + return data +} + export interface GrokExchangeCodeRequest { session_id: string state: string @@ -170,4 +181,48 @@ export async function createFromSSO(payload: GrokSSOToOAuthRequest): Promise { + const payload: Record = { sso_token: ssoToken } + if (proxyId) payload.proxy_id = proxyId + const { data } = await apiClient.post('/admin/grok/oauth/sso-token', payload, { + timeout: GROK_AUTHORIZATION_TIMEOUT_MS + }) + return data +} + +/** + * Password login → ephemeral SSO → Build OAuth. + * Password is only sent over the wire for this call; never persist it in credentials. + */ +export async function authorizePassword( + emailAndPassword: string, + proxyId?: number | null +): Promise { + // Format: email----password (password may contain dashes). + const sep = '----' + const idx = emailAndPassword.indexOf(sep) + const email = (idx >= 0 ? emailAndPassword.slice(0, idx) : emailAndPassword).trim() + const password = idx >= 0 ? emailAndPassword.slice(idx + sep.length) : '' + const payload: Record = { email, password } + if (proxyId) payload.proxy_id = proxyId + const { data } = await apiClient.post('/admin/grok/oauth/password', payload, { + timeout: GROK_AUTHORIZATION_TIMEOUT_MS + }) + return data +} + +export default { + generateAuthUrl, + getCapabilities, + exchangeCode, + refreshGrokToken, + queryQuota, + resetQuota, + createFromSSO, + validateSSOToken, + authorizePassword, +} diff --git a/frontend/src/api/admin/settings.ts b/frontend/src/api/admin/settings.ts index 1b5d347661..08629df300 100644 --- a/frontend/src/api/admin/settings.ts +++ b/frontend/src/api/admin/settings.ts @@ -32,6 +32,38 @@ export type DefaultPlatformQuotasMap = Partial + +export const SCHEDULING_THRESHOLD_PLATFORMS: SchedulingThresholdPlatformType[] = [ + "openai", + "anthropic", + "grok", +] + +export function normalizeAccountSchedulingThresholdsMap( + input?: Partial> | null, +): AccountSchedulingThresholdsMap { + const result = {} as AccountSchedulingThresholdsMap + for (const platform of SCHEDULING_THRESHOLD_PLATFORMS) { + const value = input?.[platform] + result[platform] = typeof value === "number" && Number.isFinite(value) + ? Math.min(100, Math.max(1, Math.trunc(value))) + : 100 + } + return result +} + +export function sanitizeAccountSchedulingThresholdsMap( + input?: Partial> | null, +): AccountSchedulingThresholdsMap { + return normalizeAccountSchedulingThresholdsMap(input) +} + /** 归一化为全 4 平台 × 3 窗口(缺失填 null),供模板非空绑定 */ export function normalizePlatformQuotasMap(input?: DefaultPlatformQuotasMap | null): DefaultPlatformQuotasMap { const result: DefaultPlatformQuotasMap = {} @@ -556,6 +588,12 @@ export interface SystemSettings { fallback_model_openai: string; fallback_model_gemini: string; fallback_model_antigravity: string; + grok_default_text_model: string; + grok_cross_client_model_map_enabled: boolean; + grok_default_base_url_mode: string; + + // Per-platform account auto-pause thresholds (100 = disabled) + account_scheduling_thresholds: AccountSchedulingThresholdsMap; // Identity patch configuration (Claude -> Gemini) enable_identity_patch: boolean; @@ -874,6 +912,10 @@ export interface UpdateSettingsRequest { fallback_model_openai?: string; fallback_model_gemini?: string; fallback_model_antigravity?: string; + grok_default_text_model?: string; + grok_cross_client_model_map_enabled?: boolean; + grok_default_base_url_mode?: string; + account_scheduling_thresholds?: AccountSchedulingThresholdsMap; enable_identity_patch?: boolean; identity_patch_prompt?: string; ops_monitoring_enabled?: boolean; diff --git a/frontend/src/api/admin/usage.ts b/frontend/src/api/admin/usage.ts index 83be033c11..1c740caabe 100644 --- a/frontend/src/api/admin/usage.ts +++ b/frontend/src/api/admin/usage.ts @@ -84,6 +84,7 @@ export interface AdminUsageQueryParams extends UsageQueryParams { user_id?: number exact_total?: boolean billing_mode?: string + upstream_model_mismatch?: boolean sort_by?: string sort_order?: 'asc' | 'desc' // 错误请求 tab 专属筛选(仅传给错误列表接口;共用同一 filters 对象) @@ -123,6 +124,7 @@ export async function getStats(params: { model?: string request_type?: UsageRequestType stream?: boolean + upstream_model_mismatch?: boolean period?: string start_date?: string end_date?: string diff --git a/frontend/src/components/account/AccountStatusIndicator.vue b/frontend/src/components/account/AccountStatusIndicator.vue index 9c1d25f635..4130ecd8ec 100644 --- a/frontend/src/components/account/AccountStatusIndicator.vue +++ b/frontend/src/components/account/AccountStatusIndicator.vue @@ -14,15 +14,19 @@ @@ -356,82 +354,72 @@
- {{ grokEntitlementLabel || t('admin.accounts.forbidden') }} + {{ usageInfo?.grok_entitlement_status || t('admin.accounts.forbidden') }}
-
- - {{ grokEntitlementLabel }} - -
-
-
- - {{ formatWindowRequests(grokLocalUsage) }} req - - - {{ formatWindowTokens(grokLocalUsage) }} - - - A ${{ formatWindowCost(grokLocalUsage) }} - + + + +
+ {{ usageErrorLabel }}
- - - -
{{ t('admin.accounts.usageWindow.grokRetryAfter', { time: grokRetryAfterLabel }) }}
-
- {{ grokQuotaUnknownLabel }} -
-
- {{ usageErrorLabel }} -
-
- {{ grokQuotaStatusLine }} -
- + +
+
+
-
+
-
-
@@ -629,11 +617,10 @@ import { ref, computed, onMounted, onBeforeUnmount, onUnmounted, watch } from 'vue' import { useI18n } from 'vue-i18n' import { adminAPI } from '@/api/admin' -import type { GrokQuotaProbeResult } from '@/api/admin/grok' import type { Account, AccountUsageInfo, GeminiCredentials, WindowStats } from '@/types' import { buildOpenAIUsageRefreshKey } from '@/utils/accountUsageRefresh' import { enqueueUsageRequest } from '@/utils/usageLoadQueue' -import { formatCompactNumber, formatRelativeTime } from '@/utils/format' +import { formatCompactNumber } from '@/utils/format' import UsageProgressBar from './UsageProgressBar.vue' import AccountQuotaInfo from './AccountQuotaInfo.vue' import OpenAIQuotaResetCell from './OpenAIQuotaResetCell.vue' @@ -652,16 +639,25 @@ const props = withDefaults( todayStats?: WindowStats | null todayStatsLoading?: boolean manualRefreshToken?: number + batchedUsage?: AccountUsageInfo | null + batchedUsageError?: string | null + batchedUsageLoading?: boolean + requestBatchedUsage?: ((account: Account, options?: { force?: boolean }) => void) | null }>(), { todayStats: null, todayStatsLoading: false, - manualRefreshToken: 0 + manualRefreshToken: 0, + batchedUsage: null, + batchedUsageError: null, + batchedUsageLoading: false, + requestBatchedUsage: null } ) const emit = defineEmits<{ 'account-updated': [account: Account] + 'usage-loaded': [usage: AccountUsageInfo] }>() const { t } = useI18n() @@ -674,6 +670,9 @@ const loading = ref(false) const activeQueryLoading = ref(false) const error = ref(null) const usageInfo = ref(null) +watch(usageInfo, (usage) => { + if (usage) emit('usage-loaded', usage) +}) const suppressOpenAIUsageRefreshUntil = ref(0) const rootRef = ref(null) const isDesktopViewport = ref( @@ -713,6 +712,8 @@ const shouldFetchUsage = computed(() => { return false }) +const isBatchManaged = computed(() => typeof props.requestBatchedUsage === 'function') + const showGeminiTodayStats = computed(() => { return props.account.platform === 'gemini' && props.account.type === 'service_account' }) @@ -1072,17 +1073,6 @@ interface GrokQuotaBarInfo { resetsAt: string | null } -const makeGrokQuotaBar = (quota?: { limit?: number | null; remaining?: number | null; reset_at?: string | null } | null): GrokQuotaBarInfo | null => { - if (!quota || quota.limit == null || quota.remaining == null || quota.limit <= 0) return null - const remaining = Math.min(quota.limit, Math.max(0, quota.remaining)) - return { - utilization: (remaining / quota.limit) * 100, - resetsAt: quota.reset_at || null - } -} - -const grokRequestQuotaBar = computed(() => makeGrokQuotaBar(usageInfo.value?.grok_request_quota)) -const grokTokenQuotaBar = computed(() => makeGrokQuotaBar(usageInfo.value?.grok_token_quota)) const grokBilling = computed(() => usageInfo.value?.grok_billing || null) const grokWeeklyBillingBar = computed((): GrokQuotaBarInfo | null => { const billing = grokBilling.value @@ -1094,6 +1084,63 @@ const grokWeeklyBillingBar = computed((): GrokQuotaBarInfo | null => { resetsAt: billing.period_end || null } }) +// Monthly used/limit % from billing probe (used_percent or derived from cents). +const grokMonthlyBillingBar = computed((): GrokQuotaBarInfo | null => { + const billing = grokBilling.value + if (!billing) return null + let utilization: number | null = null + if (billing.used_percent != null && Number.isFinite(billing.used_percent)) { + utilization = billing.used_percent + } else if ( + billing.monthly_limit_cents != null && + billing.monthly_limit_cents > 0 && + billing.used_cents != null + ) { + utilization = (billing.used_cents / billing.monthly_limit_cents) * 100 + } + if (utilization == null) return null + // Avoid duplicating the weekly bar when period_type is weekly-only without monthly. + if (billing.period_type?.toLowerCase() === 'weekly' && billing.monthly_limit_cents == null) { + return null + } + return { + utilization: Math.min(100, Math.max(0, utilization)), + resetsAt: billing.billing_period_end || billing.period_end || null + } +}) +const formatGrokMoney = (value?: number | null) => { + if (value == null || Number.isNaN(value)) return '0' + if (value >= 1000) return formatCompactNumber(value) + if (value >= 100) return value.toFixed(0) + if (value >= 10) return value.toFixed(1) + return value.toFixed(2) +} +// Prepaid money line for paid Grok: show when prepaid_balance is present. +// Monthly used/limit numbers are optional context; primary progress is the 30d bar. +const grokPrepaidMoneyLine = computed(() => { + const billing = grokBilling.value + if (!billing) return null + const prepaid = billing.prepaid_balance + // "只针对预付": only render when prepaid field exists (including $0.00). + if (prepaid == null || !Number.isFinite(prepaid)) return null + const used = + billing.monthly_used != null + ? billing.monthly_used + : billing.used_cents != null + ? billing.used_cents / 100 + : 0 + const limit = + billing.monthly_limit != null + ? billing.monthly_limit + : billing.monthly_limit_cents != null + ? billing.monthly_limit_cents / 100 + : 0 + return { + prepaid: formatGrokMoney(prepaid), + used: formatGrokMoney(used), + limit: formatGrokMoney(limit) + } +}) const grokPlanLabelIsFree = (value: string) => value.includes('free') || value.includes('basic') const grokPlanLabelIsPaid = (value: string) => { return value !== '' && !grokPlanLabelIsFree(value) && !value.includes('unknown') @@ -1119,14 +1166,6 @@ const grokIsFree = computed(() => { return billing != null }) const grokFreeQuotaUsage = computed(() => usageInfo.value?.grok_local_usage_24h || null) -const grokLocalUsage = computed(() => { - if (grokIsFree.value) return grokFreeQuotaUsage.value - return props.todayStats || - usageInfo.value?.grok_local_usage || - usageInfo.value?.grok_local_usage_7d || - usageInfo.value?.grok_local_usage_monthly || - null -}) const grokFreeTokenBar = computed(() => { if (!grokIsFree.value || !grokFreeQuotaUsage.value) return null const limit = usageInfo.value?.grok_free_token_limit @@ -1136,7 +1175,12 @@ const grokFreeTokenBar = computed(() => { }) const grokQuotaUnknown = computed(() => { if (props.account.platform !== 'grok') return false - if (grokBilling.value || grokFreeTokenBar.value || grokRequestQuotaBar.value || grokTokenQuotaBar.value) return false + if (grokIsFree.value) { + return !grokFreeTokenBar.value + } + if (grokWeeklyBillingBar.value || grokMonthlyBillingBar.value || grokPrepaidMoneyLine.value) { + return false + } return usageInfo.value?.grok_quota_snapshot_state !== 'observed' }) const grokQuotaUnknownLabel = computed(() => { @@ -1144,33 +1188,6 @@ const grokQuotaUnknownLabel = computed(() => { ? t('admin.accounts.usageWindow.grokNoHeaders') : t('admin.accounts.usageWindow.grokUnknown') }) -const grokQuotaStatusLine = computed(() => { - if (props.account.platform !== 'grok') return null - const parts: string[] = [] - const status = usageInfo.value?.grok_last_status_code - if (status) { - parts.push(t('admin.accounts.usageWindow.grokLastStatus', { status })) - } - if (usageInfo.value?.grok_last_quota_probe_at) { - parts.push( - t('admin.accounts.usageWindow.grokLastProbe', { - time: formatRelativeTime(usageInfo.value.grok_last_quota_probe_at) - }) - ) - } - if (usageInfo.value?.grok_last_headers_seen_at) { - parts.push( - t('admin.accounts.usageWindow.grokLastHeadersSeen', { - time: formatRelativeTime(usageInfo.value.grok_last_headers_seen_at) - }) - ) - } - return parts.length > 0 ? parts.join(' | ') : null -}) -const grokEntitlementLabel = computed(() => { - const status = (usageInfo.value?.grok_entitlement_status || '').trim() - return status || null -}) const grokRetryAfterLabel = computed(() => { const seconds = usageInfo.value?.grok_retry_after_seconds if (seconds == null || seconds <= 0) return null @@ -1179,11 +1196,6 @@ const grokRetryAfterLabel = computed(() => { return `${minutes}m` }) -const formatWindowRequests = (stats: WindowStats) => formatCompactNumber(stats.requests, { allowBillions: false }) -const formatWindowTokens = (stats: WindowStats) => formatCompactNumber(stats.tokens) -const formatWindowCost = (stats: WindowStats) => stats.cost.toFixed(2) -const formatWindowUserCost = (stats: WindowStats) => (stats.user_cost ?? 0).toFixed(2) - // 账户类型显示标签 const antigravityTierLabel = computed(() => { switch (antigravityTier.value) { @@ -1273,8 +1285,24 @@ const isAnthropicOAuthOrSetupToken = computed(() => { return props.account.platform === 'anthropic' && (props.account.type === 'oauth' || props.account.type === 'setup-token') }) +const requestParentBatchUsage = (options?: { force?: boolean }) => { + if (!isBatchManaged.value || !shouldFetchUsage.value) return + props.requestBatchedUsage?.(props.account, options) +} + +const syncManagedUsageState = () => { + if (!isBatchManaged.value) return + usageInfo.value = props.batchedUsage ?? null + error.value = props.batchedUsageError ?? null + loading.value = props.batchedUsageLoading === true +} + const loadUsage = async (options?: { source?: 'passive' | 'active'; bypassCache?: boolean }) => { if (!shouldFetchUsage.value) return + if (isBatchManaged.value) { + requestParentBatchUsage({ force: options?.bypassCache === true }) + return + } // Check cache if (!options?.bypassCache) { @@ -1290,9 +1318,9 @@ const loadUsage = async (options?: { source?: 'passive' | 'active'; bypassCache? error.value = null try { - const fetchFn = () => options?.source - ? adminAPI.accounts.getUsage(props.account.id, options.source) - : adminAPI.accounts.getUsage(props.account.id) + const fetchFn = () => options?.source + ? adminAPI.accounts.getUsage(props.account.id, options.source, options.bypassCache === true) + : adminAPI.accounts.getUsage(props.account.id) const result = await enqueueUsageRequest(props.account, fetchFn) if (!unmounted.value) { usageInfo.value = result @@ -1369,54 +1397,10 @@ const loadActiveUsage = async () => { } } -const handleGrokProbed = (result: GrokQuotaProbeResult) => { - const current = usageInfo.value - if (!current) return - const snapshot = result.snapshot - const statusCode = snapshot?.status_code ?? result.status_code - const hasActiveProbeSnapshot = snapshot != null && ( - result.source === 'active_probe' || - result.source === 'hybrid_probe' || - snapshot.observation_source === 'active_probe' - ) - const probeSucceeded = hasActiveProbeSnapshot && - statusCode != null && statusCode >= 200 && statusCode < 300 - const snapshotEntitlement = snapshot?.entitlement_status?.trim() - const currentEntitlement = current.grok_entitlement_status?.trim() - const entitlementStatus = snapshotEntitlement || ( - probeSucceeded && currentEntitlement?.toLowerCase() === 'forbidden' - ? undefined - : current.grok_entitlement_status - ) - const merged: AccountUsageInfo = { - ...current, - grok_billing: result.billing ?? current.grok_billing, - grok_local_usage_24h: result.local_usage_24h ?? current.grok_local_usage_24h, - grok_local_usage_7d: result.local_usage_7d ?? current.grok_local_usage_7d, - grok_local_usage_monthly: result.local_usage_monthly ?? current.grok_local_usage_monthly, - grok_request_quota: snapshot?.requests ?? current.grok_request_quota, - grok_token_quota: snapshot?.tokens ?? current.grok_token_quota, - grok_retry_after_seconds: snapshot?.retry_after_seconds ?? current.grok_retry_after_seconds, - grok_entitlement_status: entitlementStatus, - grok_quota_snapshot_state: result.billing - ? 'billing_observed' - : snapshot?.headers_observed - ? 'observed' - : current.grok_quota_snapshot_state, - grok_last_quota_probe_at: result.billing?.fetched_at ?? snapshot?.last_probe_at ?? current.grok_last_quota_probe_at, - grok_last_headers_seen_at: snapshot?.last_headers_seen_at ?? current.grok_last_headers_seen_at, - grok_last_status_code: result.status_code ?? snapshot?.status_code ?? current.grok_last_status_code, - is_forbidden: probeSucceeded ? false : current.is_forbidden, - forbidden_reason: probeSucceeded ? undefined : current.forbidden_reason, - forbidden_type: probeSucceeded ? undefined : current.forbidden_type, - validation_url: probeSucceeded ? undefined : current.validation_url, - needs_verify: probeSucceeded ? false : current.needs_verify, - is_banned: probeSucceeded ? false : current.is_banned, - error: result.billing || snapshot ? undefined : current.error, - error_code: result.billing || snapshot ? undefined : current.error_code - } - usageInfo.value = merged - _usageCache.set(props.account.id, { data: merged, ts: Date.now() }) +// The probe persists upstream quota state; refresh this cell so its compact +// bars and entitlement status reflect the newly observed snapshot. +const handleGrokProbed = async () => { + await loadUsage({ source: 'active', bypassCache: true }) } // ===== API Key quota progress bars ===== @@ -1529,11 +1513,49 @@ onMounted(() => { } } + if (isBatchManaged.value) { + syncManagedUsageState() + requestParentBatchUsage() + return + } + if (!shouldAutoLoadUsageOnMount.value) return const source = isAnthropicOAuthOrSetupToken.value ? 'passive' : undefined requestAutoLoad(source) }) +watch( + () => [props.batchedUsage, props.batchedUsageError, props.batchedUsageLoading, isBatchManaged.value] as const, + () => { + syncManagedUsageState() + }, + { immediate: true, deep: true } +) + +watch(isBatchManaged, (managed, wasManaged) => { + if (managed && !wasManaged) { + syncManagedUsageState() + requestParentBatchUsage() + } +}) + +watch( + () => [props.account.id, props.account.platform, props.account.type, isBatchManaged.value] as const, + ([accountID, platform, accountType, managed], [previousAccountID, previousPlatform, previousAccountType]) => { + if ( + accountID === previousAccountID && + platform === previousPlatform && + accountType === previousAccountType + ) { + return + } + if (!managed || !shouldFetchUsage.value) return + syncManagedUsageState() + requestParentBatchUsage() + }, + { flush: 'post' } +) + watch(openAIUsageRefreshKey, (nextKey, prevKey) => { if (!prevKey || nextKey === prevKey) return if (props.account.platform !== 'openai' || props.account.type !== 'oauth') return @@ -1542,6 +1564,11 @@ watch(openAIUsageRefreshKey, (nextKey, prevKey) => { return } + if (isBatchManaged.value) { + requestParentBatchUsage({ force: true }) + return + } + _usageCache.delete(props.account.id) requestAutoLoad() }) @@ -1552,6 +1579,11 @@ watch( if (nextToken === prevToken) return if (!shouldFetchUsage.value) return + if (isBatchManaged.value) { + requestParentBatchUsage({ force: true }) + return + } + const source = isAnthropicOAuthOrSetupToken.value ? 'passive' : undefined _usageCache.delete(props.account.id) loadUsage({ source, bypassCache: true }).catch((e) => { diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 22a499ee65..37fbcf6923 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -3212,6 +3212,7 @@ :show-agent-identity-option="form.platform === 'openai'" :show-codex-pat-option="form.platform === 'openai'" :show-sso-option="form.platform === 'grok'" + :show-email-password-option="false" :show-manual-option="true" :initial-input-method="'manual'" :platform="form.platform" @@ -3224,6 +3225,7 @@ @import-codex-session="handleOpenAIImportCodexSession" @import-codex-pat="handleOpenAIImportCodexPAT" @import-sso="handleGrokImportSSO" + @authorize-password="handleGrokAuthorizePassword" />
@@ -3314,8 +3316,8 @@
@@ -5476,6 +5478,116 @@ const handleGrokImportSSO = async (ssoInput: string) => { } } +/** + * Grok password login: each line is email----password. + * Password is only used for the authorize API call; buildCredentials never stores it. + */ +const handleGrokAuthorizePassword = async (emailPasswordInput: string) => { + if (!emailPasswordInput.trim()) return + if (!validateGrokOAuthUpstreamConfig()) return + + const lines = emailPasswordInput + .split('\n') + // Keep the password portion byte-for-byte; trim is only for determining + // whether this textarea line is blank. + .filter((line) => line.trim() && line.includes('----')) + + if (lines.length === 0) { + grokOAuth.error.value = t( + 'admin.accounts.oauth.grok.pleaseEnterPassword', + 'Please enter email----password (one per line)' + ) + return + } + + grokOAuth.loading.value = true + grokOAuth.error.value = '' + + let successCount = 0 + let failedCount = 0 + const errors: string[] = [] + + try { + for (let i = 0; i < lines.length; i++) { + try { + const tokenInfo = await grokOAuth.authorizePassword(lines[i], form.proxy_id) + if (!tokenInfo) { + failedCount++ + errors.push(`#${i + 1}: ${grokOAuth.error.value || 'Authorization failed'}`) + grokOAuth.error.value = '' + continue + } + + const credentials = grokOAuth.buildCredentials(tokenInfo) + applyGrokOAuthUpstreamConfig(credentials) + const extra = grokOAuth.buildExtraInfo(tokenInfo) + const accountName = + lines.length > 1 + ? `${form.name || tokenInfo.email || 'Grok OAuth Account'} #${i + 1}` + : form.name || tokenInfo.email || 'Grok OAuth Account' + + const modelMapping = buildModelMappingObject( + modelRestrictionMode.value, + allowedModels.value, + modelMappings.value + ) + if (modelMapping) { + credentials.model_mapping = modelMapping + } + if (!applyTempUnschedConfig(credentials)) { + return + } + + await adminAPI.accounts.create({ + name: accountName, + notes: form.notes, + platform: 'grok', + type: 'oauth', + credentials, + extra, + proxy_id: form.proxy_id, + concurrency: form.concurrency, + load_factor: form.load_factor ?? undefined, + priority: form.priority, + rate_multiplier: form.rate_multiplier, + group_ids: form.group_ids, + expires_at: form.expires_at, + auto_pause_on_expired: autoPauseOnExpired.value + }) + successCount++ + } catch (error: any) { + failedCount++ + const errMsg = error.response?.data?.detail || error.message || 'Unknown error' + errors.push(`#${i + 1}: ${errMsg}`) + } + } + + if (successCount > 0 && failedCount === 0) { + appStore.showSuccess( + lines.length > 1 + ? t('admin.accounts.oauth.batchSuccess', { count: successCount }) + : t('admin.accounts.accountCreated') + ) + emit('created') + handleClose() + } else if (successCount > 0) { + appStore.showWarning( + t('admin.accounts.oauth.batchPartialSuccess', { + success: successCount, + failed: failedCount + }) + ) + grokOAuth.error.value = errors.join('\n') + emit('created') + } else { + grokOAuth.error.value = errors.join('\n') + appStore.showError(t('admin.accounts.oauth.batchFailed')) + } + } finally { + grokOAuth.loading.value = false + } +} + // OpenAI OAuth 授权码兑换 const handleOpenAIExchange = async (authCode: string) => { const oauthClient = openaiOAuth diff --git a/frontend/src/components/account/EditAccountModal.vue b/frontend/src/components/account/EditAccountModal.vue index 51e0af9289..3bff385cd5 100644 --- a/frontend/src/components/account/EditAccountModal.vue +++ b/frontend/src/components/account/EditAccountModal.vue @@ -1375,6 +1375,40 @@
+ +
+
+
+ +

+ {{ t('admin.accounts.accountSchedulingThresholdOverrideHint') }} +

+
+ +
+
+ + +

{{ t('admin.accounts.accountSchedulingThresholdOverrideDisabledHint') }}

+
+
+
([]) const antigravityModelMappings = ref([]) const isSyncingAntigravityUpstream = ref(false) const tempUnschedEnabled = ref(false) +const accountSchedulingThresholdOverrideEnabled = ref(false) +const accountSchedulingThresholdOverrideValue = ref(100) +const ACCOUNT_SCHEDULING_THRESHOLD_CREDENTIAL_KEY = 'account_scheduling_threshold' +const supportsAccountSchedulingThresholdOverride = computed(() => + supportsAccountSchedulingThresholdOverridePlatform(props.account?.platform) +) const tempUnschedRules = ref([]) const getModelMappingKey = createStableObjectKeyResolver('edit-model-mapping') const getOpenAICompactModelMappingKey = createStableObjectKeyResolver('edit-openai-compact-model-mapping') @@ -3512,6 +3552,7 @@ const syncFormFromAccount = (newAccount: Account | null) => { loadQuotaControlSettings(newAccount) loadTempUnschedRules(credentials) + loadAccountSchedulingThresholdOverride(newAccount.platform, credentials) // Load header override state (anthropic/openai apikey + grok apikey/oauth) headerOverrideEnabled.value = false @@ -3874,6 +3915,69 @@ const applyTempUnschedConfig = (credentials: Record) => { return true } + +function supportsAccountSchedulingThresholdOverridePlatform(platform: Account['platform'] | undefined) { + return platform === 'openai' || platform === 'anthropic' || platform === 'grok' +} + +function normalizeAccountSchedulingThresholdOverride(value: unknown): number | null { + if (value === null || value === undefined || value === '') { + return null + } + const numeric = Number(value) + if (!Number.isFinite(numeric)) { + return null + } + const integer = Math.trunc(numeric) + if (integer < 1 || integer > 100) { + return null + } + return integer +} + +function clampAccountSchedulingThresholdOverride(value: unknown): number { + return Math.min(100, Math.max(1, Math.trunc(Number(value) || 100))) +} + +function loadAccountSchedulingThresholdOverride( + platform: Account['platform'] | undefined, + credentials: Record | undefined +) { + if (!supportsAccountSchedulingThresholdOverridePlatform(platform)) { + accountSchedulingThresholdOverrideEnabled.value = false + accountSchedulingThresholdOverrideValue.value = 100 + return + } + const value = normalizeAccountSchedulingThresholdOverride( + credentials?.[ACCOUNT_SCHEDULING_THRESHOLD_CREDENTIAL_KEY] + ) + accountSchedulingThresholdOverrideEnabled.value = value !== null + accountSchedulingThresholdOverrideValue.value = value ?? 100 +} + +const applyAccountSchedulingThresholdOverridePatch = ( + credentials: Record, + currentCredentials: Record, + platform: Account['platform'] | undefined = props.account?.platform +) => { + if (!supportsAccountSchedulingThresholdOverridePlatform(platform)) { + return + } + const current = normalizeAccountSchedulingThresholdOverride( + currentCredentials[ACCOUNT_SCHEDULING_THRESHOLD_CREDENTIAL_KEY] + ) + if (!accountSchedulingThresholdOverrideEnabled.value) { + if (current !== null) { + credentials[ACCOUNT_SCHEDULING_THRESHOLD_CREDENTIAL_KEY] = null + } + return + } + const next = clampAccountSchedulingThresholdOverride(accountSchedulingThresholdOverrideValue.value) + if (current !== next) { + credentials[ACCOUNT_SCHEDULING_THRESHOLD_CREDENTIAL_KEY] = next + } +} + function loadTempUnschedRules(credentials?: Record) { tempUnschedEnabled.value = credentials?.temp_unschedulable_enabled === true const rawRules = credentials?.temp_unschedulable_rules @@ -4233,6 +4337,7 @@ const handleSubmit = async () => { // Add intercept warmup requests setting applyInterceptWarmup(newCredentials, interceptWarmupRequests.value, 'edit') + applyAccountSchedulingThresholdOverridePatch(newCredentials, currentCredentials) if (!applyTempUnschedConfig(newCredentials)) { return } @@ -4251,6 +4356,7 @@ const handleSubmit = async () => { // Add intercept warmup requests setting applyInterceptWarmup(newCredentials, interceptWarmupRequests.value, 'edit') + applyAccountSchedulingThresholdOverridePatch(newCredentials, currentCredentials) if (!applyTempUnschedConfig(newCredentials)) { return } @@ -4299,6 +4405,7 @@ const handleSubmit = async () => { } applyInterceptWarmup(newCredentials, interceptWarmupRequests.value, 'edit') + applyAccountSchedulingThresholdOverridePatch(newCredentials, currentCredentials) if (!applyTempUnschedConfig(newCredentials)) { return } @@ -4356,6 +4463,7 @@ const handleSubmit = async () => { } applyInterceptWarmup(newCredentials, interceptWarmupRequests.value, 'edit') + applyAccountSchedulingThresholdOverridePatch(newCredentials, currentCredentials) if (!applyTempUnschedConfig(newCredentials)) { return } @@ -4367,6 +4475,7 @@ const handleSubmit = async () => { const newCredentials: Record = { ...currentCredentials } applyInterceptWarmup(newCredentials, interceptWarmupRequests.value, 'edit') + applyAccountSchedulingThresholdOverridePatch(newCredentials, currentCredentials) if (!applyTempUnschedConfig(newCredentials)) { return } diff --git a/frontend/src/components/account/GrokQuotaProbeCell.vue b/frontend/src/components/account/GrokQuotaProbeCell.vue index fa6fe29bec..e5f84077a2 100644 --- a/frontend/src/components/account/GrokQuotaProbeCell.vue +++ b/frontend/src/components/account/GrokQuotaProbeCell.vue @@ -24,18 +24,13 @@ {{ t('admin.accounts.usageWindow.grokProbe') }} - -
-
+ +
{{ summary }}
@@ -48,12 +43,17 @@ import { computed, ref, watch } from 'vue' import { useI18n } from 'vue-i18n' import { adminAPI } from '@/api/admin' -import type { GrokQuotaProbeResult, GrokQuotaWindow } from '@/api/admin/grok' +import type { GrokQuotaProbeResult } from '@/api/admin/grok' import type { Account } from '@/types' -const props = defineProps<{ - account: Account -}>() +const props = withDefaults( + defineProps<{ + account: Account + /** When true, only show the probe button (+ errors). No duplicate weekly summary. */ + compact?: boolean + }>(), + { compact: false } +) const emit = defineEmits<{ probed: [result: GrokQuotaProbeResult] }>() @@ -79,42 +79,16 @@ const extractErrorMessage = (e: unknown): string => { ) } -const formatWindow = (label: string, window?: GrokQuotaWindow | null): string | null => { - if (!window || window.limit == null || window.remaining == null) return null - return `${label} ${window.remaining}/${window.limit}` -} - -const retryAfterLabel = computed(() => { - const seconds = data.value?.snapshot?.retry_after_seconds - if (seconds == null || seconds <= 0) return null - if (seconds < 60) return `${seconds}s` - return `${Math.ceil(seconds / 60)}m` -}) - const summary = computed(() => { - const snapshot = data.value?.snapshot - if (!data.value) return '' + if (props.compact || !data.value) return '' + // Non-compact fallback (rarely used): brief weekly percent if present. const billing = data.value.billing - const parts: Array = [] if (billing?.period_type?.toLowerCase() === 'weekly' && billing.usage_percent != null) { - parts.push(t('admin.accounts.usageWindow.grokWeeklyUsage', { + return t('admin.accounts.usageWindow.grokWeeklyUsage', { percent: Math.round(Math.min(100, Math.max(0, billing.usage_percent))) - })) + }) } - if (snapshot) { - parts.push( - formatWindow(t('admin.accounts.usageWindow.grokRequests'), snapshot.requests), - formatWindow(t('admin.accounts.usageWindow.grokTokens'), snapshot.tokens) - ) - } - if (retryAfterLabel.value) { - parts.push(t('admin.accounts.usageWindow.grokRetryAfter', { time: retryAfterLabel.value })) - } - if (snapshot?.entitlement_status) { - parts.push(snapshot.entitlement_status) - } - const visibleParts = parts.filter((part): part is string => Boolean(part)) - return visibleParts.length > 0 ? visibleParts.join(' | ') : t('admin.accounts.usageWindow.grokNoHeaders') + return '' }) const truncatedError = computed(() => { diff --git a/frontend/src/components/account/OAuthAuthorizationFlow.vue b/frontend/src/components/account/OAuthAuthorizationFlow.vue index f45aeeb5d7..0c16efb857 100644 --- a/frontend/src/components/account/OAuthAuthorizationFlow.vue +++ b/frontend/src/components/account/OAuthAuthorizationFlow.vue @@ -59,6 +59,17 @@ t(getOAuthKey('ssoCookieAuth')) }} +
+ +
+
+

+ {{ t(getOAuthKey('emailPasswordDesc')) }} +

+
+ + +

+ {{ t(getOAuthKey('emailPasswordHint')) }} +

+
+
+

+ {{ error }} +

+
+ +
+
+
(), { showAgentIdentityOption: false, showCodexPatOption: false, showSsoOption: false, + showEmailPasswordOption: false, showManualOption: true, initialInputMethod: 'manual', + initialEmailPassword: '', platform: 'anthropic', showProjectId: true }) @@ -877,10 +971,15 @@ const emit = defineEmits<{ 'import-codex-session': [content: string] 'import-codex-pat': [accessToken: string] 'import-sso': [content: string] + 'authorize-password': [emailPasswordInput: string] 'update:inputMethod': [method: AuthInputMethod] }>() const { t } = useI18n() +const passwordAuthEnabled = ref(false) +const emailPasswordOptionEnabled = computed( + () => props.showEmailPasswordOption && props.platform === 'grok' && passwordAuthEnabled.value +) const showLocalCallbackNotice = computed(() => props.platform === 'openai' || props.platform === 'grok') @@ -922,10 +1021,30 @@ const sessionTokenInput = ref('') const codexSessionInput = ref('') const codexPATInput = ref('') const ssoCookieInput = ref('') +const emailPasswordInput = ref(props.initialEmailPassword || '') const showHelpDialog = ref(false) const oauthState = ref('') const projectId = ref('') +watch( + () => [props.platform, props.showEmailPasswordOption] as const, + async ([platform, requested]) => { + passwordAuthEnabled.value = false + if (platform !== 'grok' || !requested) return + try { + const capabilities = await adminAPI.grok.getCapabilities() + passwordAuthEnabled.value = capabilities.password_auth_enabled + } catch { + // Fail closed; the backend enforces the same capability. + } + }, + { immediate: true } +) + +watch(emailPasswordOptionEnabled, (enabled) => { + if (!enabled && inputMethod.value === 'email_password') inputMethod.value = 'manual' +}) + // Computed: show method selection only when there is something to choose. const methodOptionCount = computed(() => [ props.showManualOption, @@ -937,7 +1056,8 @@ const methodOptionCount = computed(() => [ props.showCodexSessionImportOption, props.showAgentIdentityOption, props.showCodexPatOption, - props.showSsoOption + props.showSsoOption, + emailPasswordOptionEnabled.value ].filter(Boolean).length) const showMethodSelection = computed(() => methodOptionCount.value > 1) @@ -977,11 +1097,34 @@ const parsedSSOCount = computed(() => { .filter((item) => item).length }) +const parsedEmailPasswordCount = computed(() => { + return emailPasswordInput.value + .split('\n') + .map((item) => item.trim()) + .filter((item) => item && item.includes('----')).length +}) + +const handleAuthorizePassword = () => { + if (emailPasswordInput.value.trim()) { + emit('authorize-password', emailPasswordInput.value) + } +} + // Watchers watch(() => props.initialInputMethod, (newVal) => { inputMethod.value = newVal }) +watch( + () => props.initialEmailPassword, + (newVal) => { + // Only prefill when the field is empty so we never overwrite operator input. + if (newVal && !emailPasswordInput.value.trim()) { + emailPasswordInput.value = newVal + } + } +) + watch(inputMethod, (newVal) => { emit('update:inputMethod', newVal) }) @@ -1081,6 +1224,7 @@ defineExpose({ codexSession: codexSessionInput, codexPAT: codexPATInput, ssoCookie: ssoCookieInput, + emailPassword: emailPasswordInput, inputMethod, reset: () => { authCodeInput.value = '' @@ -1092,6 +1236,7 @@ defineExpose({ codexSessionInput.value = '' codexPATInput.value = '' ssoCookieInput.value = '' + emailPasswordInput.value = '' inputMethod.value = props.initialInputMethod showHelpDialog.value = false } diff --git a/frontend/src/components/account/TempUnschedStatusModal.vue b/frontend/src/components/account/TempUnschedStatusModal.vue index a3e64c487e..bc1ed1685c 100644 --- a/frontend/src/components/account/TempUnschedStatusModal.vue +++ b/frontend/src/components/account/TempUnschedStatusModal.vue @@ -101,6 +101,14 @@ {{ state?.error_message || '-' }}
+ +
+ {{ triggerEvidenceText }} +
@@ -176,10 +184,28 @@ const isActive = computed(() => { }) const ruleIndexDisplay = computed(() => { - if (!state.value) return '-' + if (!state.value || !state.value.matched_keyword || state.value.rule_index < 0) return '-' return state.value.rule_index + 1 }) +const hasThresholdEvidence = computed(() => (state.value?.trigger_count || 0) > 1) + +const triggerEvidenceText = computed(() => { + const count = state.value?.trigger_count || 0 + const threshold = state.value?.trigger_threshold || 0 + const minutes = state.value?.trigger_window_minutes || 0 + if (threshold > 0 && minutes > 0) { + return t('admin.accounts.tempUnschedulable.multipleErrorTrigger', { count, threshold, minutes }) + } + if (threshold > 0) { + return t('admin.accounts.tempUnschedulable.multipleErrorTriggerNoWindow', { count, threshold }) + } + if (minutes > 0) { + return t('admin.accounts.tempUnschedulable.multipleErrorCountInWindow', { count, minutes }) + } + return t('admin.accounts.tempUnschedulable.multipleErrorCount', { count }) +}) + const triggeredAtText = computed(() => { if (!state.value?.triggered_at_unix) return '-' return formatDateTime(new Date(state.value.triggered_at_unix * 1000)) diff --git a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts index 6fcc963145..46f3908254 100644 --- a/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts +++ b/frontend/src/components/account/__tests__/AccountUsageCell.spec.ts @@ -651,7 +651,7 @@ describe('AccountUsageCell', () => { expect(badges.some(node => node.attributes('title') === 'usage.userBilled')).toBe(true) }) - it('Grok OAuth 会展示本地 user billed 用量并把耗尽配额显示为 0% 剩余', async () => { + it('Grok OAuth compact UI drops local chips and header quota bars', async () => { getUsage.mockResolvedValue({ grok_local_usage: { requests: 4, @@ -670,12 +670,7 @@ describe('AccountUsageCell', () => { const wrapper = mount(AccountUsageCell, { props: { - account: makeAccount({ - id: 3861, - platform: 'grok', - type: 'oauth', - extra: {} - }) + account: makeAccount({ id: 3861, platform: 'grok', type: 'oauth', extra: {} }) }, global: { stubs: { @@ -683,66 +678,49 @@ describe('AccountUsageCell', () => { props: ['label', 'utilization', 'resetsAt', 'color'], template: '
{{ label }}|{{ utilization }}|{{ resetsAt }}
' }, - AccountQuotaInfo: true, - GrokQuotaProbeCell: true + AccountQuotaInfo: true } } }) await flushPromises() - expect(getUsage).toHaveBeenCalledWith(3861) - expect(wrapper.text()).toContain('4 req') - expect(wrapper.text()).toContain('1.2K') - expect(wrapper.text()).toContain('A $0.12') - expect(wrapper.text()).toContain('U $0.34') - expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokRequests|0|2026-07-09T16:00:00Z') - - const badges = wrapper.findAll('span[title]') - expect(badges.some(node => node.attributes('title') === 'usage.accountBilled')).toBe(true) - expect(badges.some(node => node.attributes('title') === 'usage.userBilled')).toBe(true) + expect(wrapper.text()).not.toContain('4 req') + expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokRequests|') }) - it('Grok OAuth 配额条按剩余容量显示 100% 满格和 25% 低量', async () => { + it('Grok paid monthly limits show 30d bar without free 24h', async () => { getUsage.mockResolvedValue({ - grok_request_quota: { - limit: 100, - remaining: 100, - reset_at: '2026-07-09T16:00:00Z' + grok_billing: { + period_type: 'weekly', + usage_percent: null, + used_percent: 12, + monthly_limit_cents: 25_000, + used_cents: 3_000, + plan: '' }, - grok_token_quota: { - limit: 1000, - remaining: 250, - reset_at: '2026-07-09T16:00:00Z' - }, - grok_quota_snapshot_state: 'observed' + grok_entitlement_status: 'free', + grok_token_quota: { limit: 1_000, remaining: 250 } }) const wrapper = mount(AccountUsageCell, { props: { - account: makeAccount({ - id: 4073, - platform: 'grok', - type: 'oauth', - extra: {} - }) + account: makeAccount({ id: 4402, platform: 'grok', type: 'oauth', extra: {} }) }, global: { stubs: { UsageProgressBar: { - props: ['label', 'utilization', 'resetsAt', 'color', 'remainingCapacity'], - template: '
{{ label }}|{{ utilization }}|{{ remainingCapacity }}
' + props: ['label', 'utilization'], + template: '
{{ label }}|{{ utilization }}
' }, - AccountQuotaInfo: true, - GrokQuotaProbeCell: true + AccountQuotaInfo: true } } }) await flushPromises() - - expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokRequests|100|true') - expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokTokens|25|true') + expect(wrapper.text()).toContain('30d|') + expect(wrapper.text()).not.toContain('24h|') }) it('Grok OAuth uses the official weekly billing percentage when available', async () => { @@ -775,7 +753,6 @@ describe('AccountUsageCell', () => { template: '
{{ label }}|{{ utilization }}|{{ resetsAt }}|{{ remainingCapacity }}
' }, AccountQuotaInfo: true, - GrokQuotaProbeCell: true } } }) @@ -789,11 +766,11 @@ describe('AccountUsageCell', () => { }) it.each([ - { tokens: 0, expected: 0, compact: '0' }, - { tokens: 500_000, expected: 50, compact: '500.0K' }, - { tokens: 1_000_000, expected: 100, compact: '1.0M' }, - { tokens: 1_100_000, expected: 100, compact: '1.1M' } - ])('Grok Free derives its 1M quota from local tokens: $tokens -> $expected%', async ({ tokens, expected, compact }) => { + { tokens: 0, expected: 0 }, + { tokens: 500_000, expected: 50 }, + { tokens: 1_000_000, expected: 100 }, + { tokens: 1_100_000, expected: 100 } + ])('Grok Free derives its 1M quota from local tokens: $tokens -> $expected%', async ({ tokens, expected }) => { getUsage.mockResolvedValue({ grok_free_token_limit: 1_000_000, grok_billing: { @@ -823,7 +800,6 @@ describe('AccountUsageCell', () => { template: '
{{ label }}|{{ utilization }}
' }, AccountQuotaInfo: true, - GrokQuotaProbeCell: true } } }) @@ -831,10 +807,10 @@ describe('AccountUsageCell', () => { await flushPromises() expect(wrapper.text()).toContain(`24h|${expected}`) - expect(wrapper.findAll('span').filter((node) => node.text() === compact)).toHaveLength(1) expect(wrapper.findAll('.usage-bar')).toHaveLength(1) expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokRequests|') expect(wrapper.text()).not.toContain('admin.accounts.usageWindow.grokTokens|') + expect(wrapper.text()).not.toContain('7d|') }) it('Grok Free uses rolling 24h usage instead of today-only usage', async () => { @@ -872,7 +848,6 @@ describe('AccountUsageCell', () => { template: '
{{ label }}|{{ utilization }}|{{ title }}
' }, AccountQuotaInfo: true, - GrokQuotaProbeCell: true } } }) @@ -880,7 +855,6 @@ describe('AccountUsageCell', () => { await flushPromises() expect(wrapper.text()).toContain('24h|75|admin.accounts.usageWindow.grokFreeQuota24hHint') - expect(wrapper.text()).toContain('750.0K') expect(wrapper.text()).not.toContain('7d|') expect(wrapper.text()).not.toContain('200.0K') expect(wrapper.text()).not.toContain('250.0K') @@ -917,7 +891,6 @@ describe('AccountUsageCell', () => { template: '
{{ label }}|{{ utilization }}
' }, AccountQuotaInfo: true, - GrokQuotaProbeCell: true } } }) @@ -930,85 +903,6 @@ describe('AccountUsageCell', () => { expect(wrapper.text()).not.toContain('250.0K') }) - it('Grok paid plans are not mistaken for Free when weekly usage is temporarily missing', async () => { - getUsage.mockResolvedValue({ - grok_billing: { - period_type: 'weekly', - usage_percent: null, - plan: 'SuperGrok Heavy' - }, - grok_entitlement_status: 'free', - grok_local_usage: { - requests: 2, - tokens: 2_000_000, - cost: 1, - standard_cost: 1 - }, - grok_token_quota: { limit: 1_000, remaining: 250 } - }) - - const wrapper = mount(AccountUsageCell, { - props: { - account: makeAccount({ id: 4401, platform: 'grok', type: 'oauth', extra: {} }) - }, - global: { - stubs: { - UsageProgressBar: { - props: ['label', 'utilization'], - template: '
{{ label }}|{{ utilization }}
' - }, - AccountQuotaInfo: true, - GrokQuotaProbeCell: true - } - } - }) - - await flushPromises() - - expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokTokens|25') - expect(wrapper.text()).not.toContain('2M|') - }) - - it('Grok custom paid monthly limits override stale Free entitlement', async () => { - getUsage.mockResolvedValue({ - grok_billing: { - period_type: 'weekly', - usage_percent: null, - monthly_limit_cents: 25_000, - plan: '' - }, - grok_entitlement_status: 'free', - grok_local_usage: { - requests: 2, - tokens: 2_000_000, - cost: 1, - standard_cost: 1 - }, - grok_token_quota: { limit: 1_000, remaining: 250 } - }) - - const wrapper = mount(AccountUsageCell, { - props: { - account: makeAccount({ id: 4402, platform: 'grok', type: 'oauth', extra: {} }) - }, - global: { - stubs: { - UsageProgressBar: { - props: ['label', 'utilization'], - template: '
{{ label }}|{{ utilization }}
' - }, - AccountQuotaInfo: true, - GrokQuotaProbeCell: true - } - } - }) - - await flushPromises() - - expect(wrapper.text()).toContain('admin.accounts.usageWindow.grokTokens|25') - expect(wrapper.text()).not.toContain('2M|') - }) - it('Grok credential Free tier keeps the 1M fallback when billing is unavailable', async () => { getUsage.mockResolvedValue({ grok_free_token_limit: 1_000_000, @@ -1032,7 +926,6 @@ describe('AccountUsageCell', () => { template: '
{{ label }}|{{ utilization }}
' }, AccountQuotaInfo: true, - GrokQuotaProbeCell: true } } }) @@ -1042,257 +935,6 @@ describe('AccountUsageCell', () => { expect(wrapper.text()).toContain('24h|100') }) - it('Grok paid manual probes keep the weekly/local summary when 24h usage is returned', async () => { - getUsage.mockResolvedValue({ - grok_quota_snapshot_state: 'no_headers', - error: 'stale error', - error_code: 'quota_unknown' - }) - - const wrapper = mount(AccountUsageCell, { - props: { - account: makeAccount({ id: 4501, platform: 'grok', type: 'oauth', extra: {} }) - }, - global: { - stubs: { - UsageProgressBar: { - props: ['label', 'utilization', 'resetsAt'], - template: '
{{ label }}|{{ utilization }}|{{ resetsAt }}
' - }, - AccountQuotaInfo: true, - GrokQuotaProbeCell: { - emits: ['probed'], - template: `` - } - } - } - }) - - await flushPromises() - await wrapper.get('.probe').trigger('click') - - expect(wrapper.text()).toContain('7d|42|2026-07-17T00:00:00Z') - expect(wrapper.text()).toContain('1.0M') - expect(wrapper.text()).not.toContain('750.0K') - expect(wrapper.text()).toContain('ACTIVE') - expect(wrapper.text()).not.toContain('stale error') - }) - - it('Grok successful probes immediately clear stale forbidden state', async () => { - getUsage.mockResolvedValue({ - is_forbidden: true, - forbidden_reason: 'stale forbidden response', - forbidden_type: 'validation', - validation_url: 'https://example.com/verify', - needs_verify: true, - is_banned: true, - grok_entitlement_status: 'forbidden', - grok_quota_snapshot_state: 'no_headers', - error: 'stale forbidden response', - error_code: 'forbidden' - }) - - const wrapper = mount(AccountUsageCell, { - props: { - account: makeAccount({ id: 4503, platform: 'grok', type: 'oauth', extra: {} }) - }, - global: { - stubs: { - UsageProgressBar: true, - AccountQuotaInfo: true, - GrokQuotaProbeCell: true - } - } - }) - - await flushPromises() - expect(wrapper.text()).toContain('forbidden') - - const setupState = wrapper.vm.$.setupState as { - handleGrokProbed: (result: Record) => void - usageInfo: Record | null - } - setupState.handleGrokProbed({ - source: 'active_probe', - snapshot: { - headers_observed: false, - updated_at: '2026-07-18T00:00:00Z', - status_code: 200 - }, - status_code: 200, - headers_observed: false, - reset_supported: false, - fetched_at: 1 - }) - await wrapper.vm.$nextTick() - - expect(setupState.usageInfo).toMatchObject({ - is_forbidden: false, - needs_verify: false, - is_banned: false, - grok_last_status_code: 200 - }) - expect(setupState.usageInfo?.forbidden_reason).toBeUndefined() - expect(setupState.usageInfo?.forbidden_type).toBeUndefined() - expect(setupState.usageInfo?.validation_url).toBeUndefined() - expect(setupState.usageInfo?.grok_entitlement_status).toBeUndefined() - expect(wrapper.text()).not.toContain('admin.accounts.forbidden') - }) - - it('Grok successful probes preserve the entitlement reported by the latest snapshot', async () => { - getUsage.mockResolvedValue({ - is_forbidden: true, - grok_entitlement_status: 'forbidden' - }) - - const wrapper = mount(AccountUsageCell, { - props: { - account: makeAccount({ id: 4504, platform: 'grok', type: 'oauth', extra: {} }) - }, - global: { - stubs: { - UsageProgressBar: true, - AccountQuotaInfo: true, - GrokQuotaProbeCell: true - } - } - }) - - await flushPromises() - - const setupState = wrapper.vm.$.setupState as { - handleGrokProbed: (result: Record) => void - usageInfo: Record | null - } - setupState.handleGrokProbed({ - source: 'active_probe', - snapshot: { - headers_observed: true, - updated_at: '2026-07-18T00:00:00Z', - entitlement_status: 'ACTIVE', - status_code: 200 - }, - status_code: 200, - headers_observed: true, - reset_supported: false, - fetched_at: 1 - }) - await wrapper.vm.$nextTick() - - expect(setupState.usageInfo?.grok_entitlement_status).toBe('ACTIVE') - expect(wrapper.text()).toContain('ACTIVE') - expect(wrapper.text()).not.toContain('admin.accounts.forbidden') - }) - - it('Grok billing-only success does not clear an active-probe forbidden state', async () => { - getUsage.mockResolvedValue({ - is_forbidden: true, - forbidden_type: 'forbidden', - needs_verify: true, - is_banned: true, - grok_entitlement_status: 'forbidden' - }) - - const wrapper = mount(AccountUsageCell, { - props: { - account: makeAccount({ id: 4505, platform: 'grok', type: 'oauth', extra: {} }) - }, - global: { - stubs: { - UsageProgressBar: true, - AccountQuotaInfo: true, - GrokQuotaProbeCell: true - } - } - }) - - await flushPromises() - - const setupState = wrapper.vm.$.setupState as { - handleGrokProbed: (result: Record) => void - usageInfo: Record | null - } - setupState.handleGrokProbed({ - source: 'billing_probe', - billing: { - period_type: 'weekly', - usage_percent: 10, - plan: 'SuperGrok' - }, - status_code: 200, - headers_observed: false, - reset_supported: false, - fetched_at: 1 - }) - await wrapper.vm.$nextTick() - - expect(setupState.usageInfo).toMatchObject({ - is_forbidden: true, - forbidden_type: 'forbidden', - needs_verify: true, - is_banned: true, - grok_entitlement_status: 'forbidden' - }) - expect(wrapper.text()).toContain('forbidden') - }) - - it('Grok Free manual probes merge rolling 24h usage', async () => { - getUsage.mockResolvedValue({ - grok_free_token_limit: 1_000_000, - subscription_tier: 'FREE', - grok_quota_snapshot_state: 'no_headers' - }) - - const wrapper = mount(AccountUsageCell, { - props: { - account: makeAccount({ id: 4502, platform: 'grok', type: 'oauth', extra: {} }) - }, - global: { - stubs: { - UsageProgressBar: { - props: ['label', 'utilization'], - template: '
{{ label }}|{{ utilization }}
' - }, - AccountQuotaInfo: true, - GrokQuotaProbeCell: { - emits: ['probed'], - template: `` - } - } - } - }) - - await flushPromises() - await wrapper.get('.probe').trigger('click') - - expect(wrapper.text()).toContain('24h|75') - expect(wrapper.text()).toContain('750.0K') - expect(wrapper.text()).not.toContain('7d|') - }) - it('Key 账号在 today stats loading 时显示骨架屏', async () => { const wrapper = mount(AccountUsageCell, { props: { @@ -1424,7 +1066,6 @@ describe('AccountUsageCell', () => { template: '
{{ label }}|{{ utilization }}
' }, AccountQuotaInfo: true, - GrokQuotaProbeCell: true } } }) @@ -1468,7 +1109,6 @@ describe('AccountUsageCell', () => { template: '
{{ label }}|{{ utilization }}
' }, AccountQuotaInfo: true, - GrokQuotaProbeCell: true } } }) diff --git a/frontend/src/components/account/__tests__/CreateAccountModal.grok.spec.ts b/frontend/src/components/account/__tests__/CreateAccountModal.grok.spec.ts index 1d0111e42c..c0273e194f 100644 --- a/frontend/src/components/account/__tests__/CreateAccountModal.grok.spec.ts +++ b/frontend/src/components/account/__tests__/CreateAccountModal.grok.spec.ts @@ -23,9 +23,13 @@ describe('CreateAccountModal Grok account types', () => { expect(source).toContain('form.platform === \'grok\' && isOAuthFlow') }) - it('validates and applies upstream config on all three Grok OAuth create paths', () => { - // 授权码兑换 / RT 批量 / SSO 批量 3 处调用(定义为箭头函数,不计入) - expect(source.match(/validateGrokOAuthUpstreamConfig\(\)/g)?.length).toBe(3) - expect(source.match(/applyGrokOAuthUpstreamConfig\(credentials\)/g)?.length).toBe(3) + it('validates and applies upstream config on Grok OAuth create paths', () => { + // 授权码兑换 / RT 批量 / SSO 批量(密码授权已隐藏) + expect(source.match(/validateGrokOAuthUpstreamConfig\(\)/g)?.length).toBeGreaterThanOrEqual(3) + expect(source.match(/applyGrokOAuthUpstreamConfig\(credentials\)/g)?.length).toBeGreaterThanOrEqual(3) + }) + + it('hides Grok password authorize option in the create flow', () => { + expect(source).toContain(':show-email-password-option="false"') }) }) diff --git a/frontend/src/components/admin/account/AccountTestModal.vue b/frontend/src/components/admin/account/AccountTestModal.vue index 0a8f853ebb..766686e3de 100644 --- a/frontend/src/components/admin/account/AccountTestModal.vue +++ b/frontend/src/components/admin/account/AccountTestModal.vue @@ -41,13 +41,28 @@ -
+ +
+ +
-
+