Merge pull request #5408 from IanShaw027/feat/grok-complete-integration

feat(grok): 完善 Grok 平台集成 — 授权、模型映射、媒体/Voice/搜索计费与调度门禁
This commit is contained in:
Wesley Liddick
2026-08-09 11:31:09 +08:00
committed by GitHub
240 changed files with 19246 additions and 1452 deletions
+3
View File
@@ -143,3 +143,6 @@ docs/*
frontend/coverage/
aicodex
output/
# Vitest / Vite cache at repo root
.vite/
+3 -3
View File
@@ -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)
@@ -193,11 +193,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)
+71 -2
View File
@@ -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(", ")
+43
View File
@@ -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()
+230
View File
@@ -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))
+482
View File
@@ -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) {
+364
View File
@@ -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)
}
+6 -1
View File
@@ -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",
+502 -1
View File
@@ -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
+34 -18
View File
@@ -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()
+32
View File
@@ -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).
+10
View File
@@ -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=
+56
View File
@@ -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)
}
+13
View File
@@ -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)
@@ -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:<mime>;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"`
@@ -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)
@@ -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
}
@@ -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)
@@ -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) {
@@ -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) {
+70 -50
View File
@@ -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,
@@ -372,6 +372,10 @@ func (h *SettingHandler) GetSettings(c *gin.Context) {
ChannelMonitorEnabled: settings.ChannelMonitorEnabled,
ChannelMonitorDefaultIntervalSeconds: settings.ChannelMonitorDefaultIntervalSeconds,
GrokDefaultTextModel: settings.GrokDefaultTextModel,
GrokCrossClientModelMapEnabled: settings.GrokCrossClientModelMapEnabled,
GrokDefaultBaseURLMode: settings.GrokDefaultBaseURLMode,
AvailableChannelsEnabled: settings.AvailableChannelsEnabled,
ModelPlazaEnabled: settings.ModelPlazaEnabled,
@@ -380,7 +384,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)
@@ -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]
@@ -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",
@@ -330,6 +330,11 @@ type UpdateSettingsRequest struct {
ChannelMonitorEnabled *bool `json:"channel_monitor_enabled"`
ChannelMonitorDefaultIntervalSeconds *int `json:"channel_monitor_default_interval_seconds"`
// 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"`
@@ -354,6 +359,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"`
@@ -1476,7 +1484,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,
@@ -1860,6 +1869,24 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
}
return previousSettings.ChannelMonitorDefaultIntervalSeconds
}(),
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
@@ -2293,6 +2320,10 @@ func (h *SettingHandler) UpdateSettings(c *gin.Context) {
ChannelMonitorEnabled: updatedSettings.ChannelMonitorEnabled,
ChannelMonitorDefaultIntervalSeconds: updatedSettings.ChannelMonitorDefaultIntervalSeconds,
GrokDefaultTextModel: updatedSettings.GrokDefaultTextModel,
GrokCrossClientModelMapEnabled: updatedSettings.GrokCrossClientModelMapEnabled,
GrokDefaultBaseURLMode: updatedSettings.GrokDefaultBaseURLMode,
AvailableChannelsEnabled: updatedSettings.AvailableChannelsEnabled,
ModelPlazaEnabled: updatedSettings.ModelPlazaEnabled,
@@ -2304,6 +2335,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 {
+5
View File
@@ -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,
+8
View File
@@ -303,6 +303,11 @@ type SystemSettings struct {
ChannelMonitorEnabled bool `json:"channel_monitor_enabled"`
ChannelMonitorDefaultIntervalSeconds int `json:"channel_monitor_default_interval_seconds"`
// 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"`
@@ -327,6 +332,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"`
}
+7 -1
View File
@@ -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"`
@@ -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 {
@@ -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.")
}
+338
View File
@@ -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 ""
}
@@ -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")
}
}
+203 -4
View File
@@ -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),
+16 -4
View File
@@ -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))
})
}
}
@@ -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
}
@@ -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,
@@ -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")
}
+125
View File
@@ -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
}
@@ -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)
}
+93 -27
View File
@@ -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
}
+36 -1
View File
@@ -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{
+79
View File
@@ -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))
}
@@ -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"))
}
+269 -23
View File
@@ -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
}
}
+65
View File
@@ -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"))
}
+133 -17
View File
@@ -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()
@@ -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"))
}
+20 -8
View File
@@ -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")
}
+71 -7
View File
@@ -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
+101
View File
@@ -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))
}
+49 -19
View File
@@ -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 {
+29 -1
View File
@@ -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 {
@@ -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,
@@ -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) {
@@ -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)
@@ -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
}
@@ -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(`<html>403 Forbidden</html>`))
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)
}
+26
View File
@@ -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 {
+15 -9
View File
@@ -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"
)
// 默认配置常量
@@ -85,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
)
@@ -450,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
@@ -461,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
@@ -83,6 +83,11 @@ var migrationChecksumCompatibilityRules = map[string]migrationChecksumCompatibil
"123_fix_legacy_auth_source_grant_on_signup_defaults.sql": newMigrationChecksumCompatibilityRule("2ce43c2cd89e9f9e1febd34a407ed9e84d177386c5544b6f02c1f58a21129f57", "6cd33422f215dcd1f486ab6f35c0ea5805d9ca69bb25906d94bc649156657145"),
"159_batch_image_foundation.sql": newMigrationChecksumCompatibilityRule("d902b70982025ec519749faf058aab7631e82c3f48167b9a4ae4db718eb72cce", "82da85b5d98e67a0507647b873a40373e84538e4adafdeed6767c0ac8b6570b2"),
"161_batch_image_pricing_snapshot.sql": newMigrationChecksumCompatibilityRule("4012af3e43636cb6af22e0176d59d1fcc70615c0f310194329461ae462c4fbd6", "96d915c9b7a6941ae99039e0ff3f1a61481eb9bddd933d11c6fadb2274554e87"),
// 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 迁移文件应用到指定的数据库。
@@ -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": "",
@@ -1155,6 +1163,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,
@@ -1274,6 +1285,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": "",
+4
View File
@@ -379,6 +379,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)
@@ -463,9 +464,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)
+120
View File
@@ -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 ""
@@ -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")
@@ -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)
+40 -13
View File
@@ -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
}
@@ -1311,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.
@@ -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",
@@ -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
}
@@ -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"},
@@ -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
}
@@ -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))
}
@@ -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)
}
@@ -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
}
@@ -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%")
}
+2 -2
View File
@@ -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 {
File diff suppressed because it is too large Load Diff
@@ -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")
}
+133 -25
View File
@@ -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)
@@ -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])
}
}
@@ -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
+10 -57
View File
@@ -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{
+25
View File
@@ -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 {
@@ -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),
@@ -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])
+20 -4
View File
@@ -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 平台使用)
+31 -26
View File
@@ -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.
@@ -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,
@@ -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})
@@ -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 }
+91 -1
View File
@@ -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
@@ -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
}
@@ -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))
}
@@ -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 {
@@ -397,6 +408,19 @@ const (
// pre-filled when creating a new channel monitor from the admin UI. Range: [15, 3600].
SettingKeyChannelMonitorDefaultIntervalSeconds = "channel_monitor_default_interval_seconds"
// 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).
@@ -576,6 +600,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 {
@@ -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 {
@@ -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
+51 -3
View File
@@ -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) {
+102
View File
@@ -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,6 +584,11 @@ 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
@@ -595,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
@@ -1225,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 == "" {
@@ -1236,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 {
@@ -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 分组请求的计费模型:来源覆盖把计费模型
@@ -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)
}
@@ -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()

Some files were not shown because too many files have changed in this diff Show More