mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:08:03 +08:00
Merge pull request #5408 from IanShaw027/feat/grok-complete-integration
feat(grok): 完善 Grok 平台集成 — 授权、模型映射、媒体/Voice/搜索计费与调度门禁
This commit is contained in:
@@ -143,3 +143,6 @@ docs/*
|
||||
frontend/coverage/
|
||||
aicodex
|
||||
output/
|
||||
|
||||
# Vitest / Vite cache at repo root
|
||||
.vite/
|
||||
|
||||
@@ -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
@@ -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(", ")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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.")
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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),
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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": "",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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%")
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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 平台使用)
|
||||
|
||||
@@ -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 }
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user