fix(grok): close free-by-default billing and related review blockers

H1/H2: bill search and voice with code defaults when group prices are
nil (explicit 0 remains free); bump API key auth snapshot to v19 and
refresh incomplete media/search/audio projections.

M1–M6: free-quota soft gate fails open on cache miss with background
refresh and 60s default TTL; correct password_auth config docs; default
cross-client model map to true (→ grok-4.5); audit /tts and /web_search;
exclude composite from migration 220 video-price clears; never let a
search surcharge mask token pricing failures.
This commit is contained in:
IanShaw027
2026-08-08 14:39:22 +08:00
parent cec922d335
commit 7eb1310701
22 changed files with 342 additions and 99 deletions
+6 -3
View File
@@ -1036,10 +1036,12 @@ type GatewayConfig struct {
// - 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: bound hot-path aggregate query frequency (0 disables cache).
// - 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.
@@ -1050,7 +1052,8 @@ type GatewayGrokConfig struct {
FreeQuotaSoftGatePercent int `mapstructure:"free_quota_soft_gate_percent"`
// FreeQuotaWindowHours controls the local rolling usage window.
FreeQuotaWindowHours int `mapstructure:"free_quota_window_hours"`
// FreeQuotaStatsCacheSeconds bounds hot-path aggregate query frequency.
// 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"`
}
@@ -2347,7 +2350,7 @@ func setDefaults() {
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", 5)
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)
+1 -1
View File
@@ -548,7 +548,7 @@ func TestLoadDefaultGrokFreeQuotaSoftGate(t *testing.T) {
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, 5, cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds)
require.Equal(t, 60, cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds)
}
func TestLoadDefaultOpenAIHTTP2Enabled(t *testing.T) {
+32 -5
View File
@@ -71,6 +71,31 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) {
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 {
@@ -174,11 +199,13 @@ func (h *GatewayHandler) WebSearch(c *gin.Context) {
// 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 && (apiKey.Group.GetSearchPricePer1k() == nil || *apiKey.Group.GetSearchPricePer1k() <= 0) {
logger.L().With(
zap.String("component", "handler.gateway.web_search"),
zap.Int64("group_id", apiKey.Group.ID),
).Warn("gateway.web_search.search_price_per_1k_unset_free")
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{
+38
View File
@@ -2,6 +2,7 @@ package handler
import (
"context"
"encoding/json"
"errors"
"io"
"net/http"
@@ -149,6 +150,24 @@ func (h *OpenAIGatewayHandler) GrokVoice(c *gin.Context, endpoint string) {
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"
@@ -298,3 +317,22 @@ func readGrokVoiceGatewayBody(c *gin.Context) ([]byte, error) {
}
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 ""
}
+4 -4
View File
@@ -63,12 +63,12 @@ const (
)
// ModelMappingOptions controls optional expansions of the default mapping.
// Cross-client wildcards (gpt-*/claude-*) are OFF unless explicitly enabled —
// silent rewrite of foreign model names is opt-in for operators who want
// Codex/Claude clients to talk to Grok groups without renaming models.
// 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.
// Empty → DefaultTextModel (grok-4.5).
DefaultText string
// EnableCrossClientMap merges gpt-*/codex-*/o*/claude-* → DefaultText.
EnableCrossClientMap bool
@@ -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 迁移文件应用到指定的数据库。
+2 -2
View File
@@ -886,7 +886,7 @@ func TestAPIContracts(t *testing.T) {
"hide_ccs_import_button": false,
"grok_default_text_model": "grok-4.5",
"grok_default_base_url_mode": "cli",
"grok_cross_client_model_map_enabled": false,
"grok_cross_client_model_map_enabled": true,
"purchase_subscription_enabled": false,
"purchase_subscription_url": "",
"table_default_page_size": 20,
@@ -1165,7 +1165,7 @@ func TestAPIContracts(t *testing.T) {
"hide_ccs_import_button": false,
"grok_default_text_model": "grok-4.5",
"grok_default_base_url_mode": "cli",
"grok_cross_client_model_map_enabled": false,
"grok_cross_client_model_map_enabled": true,
"purchase_subscription_enabled": false,
"purchase_subscription_url": "",
"table_default_page_size": 20,
@@ -46,14 +46,14 @@ func TestEveryGatewayPOSTRouteIsClassifiedForPromptAuditCoverage(t *testing.T) {
"/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",
"/tts": "voice synthesis input is not a text-generation prompt",
"/stt": "speech transcription is not a text-generation prompt",
"/custom-voices": "voice profile management has no model prompt",
"/web_search": "search query is handled by the dedicated search workflow",
}
unclassified := make([]string, 0)
@@ -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
@@ -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})
@@ -10,7 +10,10 @@ func TestCalculateSearchCost(t *testing.T) {
t.Parallel()
s := &BillingService{}
require.Equal(t, 0.0, s.CalculateSearchCost(0, floatPtr(10), 1).ActualCost)
require.Equal(t, 0.0, s.CalculateSearchCost(5, nil, 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
@@ -27,6 +30,13 @@ func TestCalculateAudioCost(t *testing.T) {
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 }
+23 -3
View File
@@ -1437,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 网页搜索按次费用。
@@ -1465,18 +1474,25 @@ func (s *BillingService) CalculateWebSearchCost(callCount int, groupPrice *float
}
// CalculateSearchCost bills search/tool invocations (e.g. web_search) per 1k calls.
// Uses explicit group search_price_per_1k when set; otherwise returns zero cost.
// 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{}
}
if groupPricePer1k == nil || *groupPricePer1k <= 0 {
pricePer1k := defaultSearchPricePer1k
if groupPricePer1k != nil {
if *groupPricePer1k < 0 {
return &CostBreakdown{}
}
pricePer1k = *groupPricePer1k
}
if pricePer1k == 0 {
return &CostBreakdown{}
}
if rateMultiplier < 0 {
rateMultiplier = 0
}
unit := *groupPricePer1k / 1000.0
unit := pricePer1k / 1000.0
total := unit * float64(numCalls)
return &CostBreakdown{
TotalCost: total,
@@ -1492,6 +1508,7 @@ type audioPriceConfig struct {
}
// 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{}
@@ -1499,14 +1516,17 @@ func (s *BillingService) CalculateAudioCost(mode string, durationOrUnits float64
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
}
@@ -864,8 +864,8 @@ func (s *GatewayService) calculateRecordUsageCost(
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 unset; search calls free group_model=%s count=%d", billingModel, result.SearchCount)
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) {
@@ -19,7 +19,7 @@ import (
// - free_quota_token_limit (int64, default 500_000)
// - free_quota_soft_gate_percent (int, default 95) — stop scheduling before the nominal limit
// - free_quota_window_hours (int, default 24) — local usage rolling window
// - free_quota_stats_cache_seconds (int, default 5) — bound hot-path aggregate query frequency
// - free_quota_stats_cache_seconds (int, default 60) — stats cache TTL; hot path never waits on DB
//
// Soft-gate applies only to *explicit* free OAuth (subscription_tier/plan_type ==
// "free"). Media/cache free detection uses isKnownGrokFreeAccount instead.
@@ -127,6 +127,9 @@ func (s *GatewayService) filterGrokFreeQuotaAccountsForGateway(ctx context.Conte
var gatewayGrokFreeQuotaGateCache sync.Map
var openaiGrokFreeQuotaGateCache sync.Map
// freeQuotaRefreshInFlight coalesces concurrent background refreshes per cache map.
var freeQuotaRefreshInFlight sync.Map // *sync.Map -> *sync.Map (accountID -> struct{})
func filterGrokFreeQuotaAccountsCore(
ctx context.Context,
cfg *config.Config,
@@ -160,6 +163,7 @@ func filterGrokFreeQuotaAccountsCore(
continue
}
}
// Miss / stale: fail open on this request; refresh asynchronously.
if _, exists := seenMissing[account.ID]; !exists {
seenMissing[account.ID] = struct{}{}
missingIDs = append(missingIDs, account.ID)
@@ -167,40 +171,7 @@ func filterGrokFreeQuotaAccountsCore(
}
if len(missingIDs) > 0 {
statsByID, err := queryGrokFreeQuotaWindowStats(ctx, usageLogRepo, missingIDs, now.Add(-settings.window))
if err != nil {
grokFreeQuotaGateQueryFailureTotal.Add(1)
if settings.cacheTTL > 0 {
for _, accountID := range missingIDs {
cache.Store(accountID, grokFreeQuotaGateCacheEntry{checkedAt: now})
}
}
slog.Warn("grok_free_quota_soft_gate_stats_failed",
"account_count", len(missingIDs),
"window_hours", settings.window.Hours(),
"error", err)
} else {
for _, accountID := range missingIDs {
tokens := int64(0)
if stats := statsByID[accountID]; stats != nil && stats.Tokens > 0 {
tokens = stats.Tokens
}
tokensByID[accountID] = tokens
cache.Store(accountID, grokFreeQuotaGateCacheEntry{tokens: tokens, checkedAt: now, known: true})
if tokens >= settings.gateTokens {
grokFreeQuotaGateBlockedTotal.Add(1)
slog.Info("grok_free_quota_soft_gate_blocked",
"account_id", accountID,
"tokens", tokens,
"gate_tokens", settings.gateTokens,
"limit_tokens", settings.limitTokens,
"window_hours", settings.window.Hours())
}
}
}
// Only sweep on the query path: it is already TTL-bounded, so the Range
// cost stays off the cache-hit hot path.
sweepGrokFreeQuotaGateCache(cache, now, settings.cacheTTL)
scheduleGrokFreeQuotaStatsRefresh(usageLogRepo, cache, settings, missingIDs)
}
filtered := make([]Account, 0, len(accounts))
@@ -216,6 +187,76 @@ func filterGrokFreeQuotaAccountsCore(
return filtered
}
// scheduleGrokFreeQuotaStatsRefresh loads usage stats off the request path.
// Concurrent callers for the same accountID are coalesced via in-flight markers.
func scheduleGrokFreeQuotaStatsRefresh(
usageLogRepo UsageLogRepository,
cache *sync.Map,
settings grokFreeQuotaGateSettings,
accountIDs []int64,
) {
if usageLogRepo == nil || cache == nil || len(accountIDs) == 0 {
return
}
inFlightRoot, _ := freeQuotaRefreshInFlight.LoadOrStore(cache, &sync.Map{})
inFlight := inFlightRoot.(*sync.Map)
toFetch := make([]int64, 0, len(accountIDs))
for _, id := range accountIDs {
if _, loaded := inFlight.LoadOrStore(id, struct{}{}); !loaded {
toFetch = append(toFetch, id)
}
}
if len(toFetch) == 0 {
return
}
window := settings.window
gateTokens := settings.gateTokens
limitTokens := settings.limitTokens
cacheTTL := settings.cacheTTL
go func() {
defer func() {
for _, id := range toFetch {
inFlight.Delete(id)
}
}()
now := time.Now().UTC()
statsByID, err := queryGrokFreeQuotaWindowStats(context.Background(), usageLogRepo, toFetch, now.Add(-window))
if err != nil {
grokFreeQuotaGateQueryFailureTotal.Add(1)
if cacheTTL > 0 {
for _, accountID := range toFetch {
cache.Store(accountID, grokFreeQuotaGateCacheEntry{checkedAt: now})
}
}
slog.Warn("grok_free_quota_soft_gate_stats_failed",
"account_count", len(toFetch),
"window_hours", window.Hours(),
"error", err)
sweepGrokFreeQuotaGateCache(cache, now, cacheTTL)
return
}
for _, accountID := range toFetch {
tokens := int64(0)
if stats := statsByID[accountID]; stats != nil && stats.Tokens > 0 {
tokens = stats.Tokens
}
cache.Store(accountID, grokFreeQuotaGateCacheEntry{tokens: tokens, checkedAt: now, known: true})
if tokens >= gateTokens {
grokFreeQuotaGateBlockedTotal.Add(1)
slog.Info("grok_free_quota_soft_gate_blocked",
"account_id", accountID,
"tokens", tokens,
"gate_tokens", gateTokens,
"limit_tokens", limitTokens,
"window_hours", window.Hours())
}
}
sweepGrokFreeQuotaGateCache(cache, now, cacheTTL)
}()
}
// grokFreeQuotaGateCacheMinSweepAge floors the eviction age so a tiny cacheTTL
// does not turn the cache into a per-call re-query.
const grokFreeQuotaGateCacheMinSweepAge = 5 * time.Minute
@@ -59,7 +59,7 @@ func grokFreeQuotaTestConfig() *config.Config {
cfg.Gateway.Grok.FreeQuotaTokenLimit = 500_000
cfg.Gateway.Grok.FreeQuotaSoftGatePercent = 95
cfg.Gateway.Grok.FreeQuotaWindowHours = 24
cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds = 5
cfg.Gateway.Grok.FreeQuotaStatsCacheSeconds = 60
return cfg
}
@@ -67,6 +67,8 @@ func TestFilterGrokFreeQuotaAccountsOnlyBlocksExplicitFreeOAuth(t *testing.T) {
repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{
1: {Tokens: 475_000}, // 95% of 500k
}}
// Clear shared cache for deterministic unit tests.
openaiGrokFreeQuotaGateCache = sync.Map{}
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}}
accounts := []Account{
{ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "FREE"}},
@@ -75,15 +77,26 @@ func TestFilterGrokFreeQuotaAccountsOnlyBlocksExplicitFreeOAuth(t *testing.T) {
{ID: 4, Platform: PlatformGrok, Type: AccountTypeAPIKey, Credentials: map[string]any{"subscription_tier": "FREE"}},
}
// First pass: cache miss fails open (does not block) and schedules background refresh.
filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)
require.Equal(t, []int64{1, 2, 3, 4}, accountIDs(filtered), "miss fails open on hot path")
require.Eventually(t, func() bool {
repo.mu.Lock()
defer repo.mu.Unlock()
return repo.calls >= 1
}, 2*time.Second, 10*time.Millisecond)
// Second pass: uses refreshed cache and blocks over-gate free OAuth.
filtered = scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)
require.Equal(t, []int64{2, 3, 4}, accountIDs(filtered), "paid and unknown fail-open; API-key free marker is not gated")
require.Equal(t, 1, repo.calls)
require.Equal(t, []int64{1}, repo.lastIDs, "paid, unknown, and API-key accounts must not enter the local free-tier query")
require.WithinDuration(t, time.Now().UTC().Add(-24*time.Hour), repo.start, time.Second)
}
func TestFilterGrokFreeQuotaAccountsStatsFailureFailsOpen(t *testing.T) {
repo := &grokFreeQuotaUsageRepoStub{err: errors.New("usage database unavailable")}
openaiGrokFreeQuotaGateCache = sync.Map{}
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}}
accounts := []Account{{
ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth,
@@ -92,7 +105,12 @@ func TestFilterGrokFreeQuotaAccountsStatsFailureFailsOpen(t *testing.T) {
filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)
require.Equal(t, []int64{1}, accountIDs(filtered))
// Cache the failure entry so a second call still fails open without re-query thrash.
require.Eventually(t, func() bool {
repo.mu.Lock()
defer repo.mu.Unlock()
return repo.calls >= 1
}, 2*time.Second, 10*time.Millisecond)
// Negative cache entry keeps subsequent hot-path calls fail-open without thrash.
filtered = scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)
require.Equal(t, []int64{1}, accountIDs(filtered))
require.Equal(t, 1, repo.calls)
@@ -118,23 +136,36 @@ func TestFilterGrokFreeQuotaAccountsRecoversAfterRollingUsageFalls(t *testing.T)
repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{
1: {Tokens: 490_000},
}}
openaiGrokFreeQuotaGateCache = sync.Map{}
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}}
accounts := []Account{{
ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth,
Credentials: map[string]any{"plan_type": "free"},
}}
require.Empty(t, scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts))
// Miss fails open, then background fill blocks over-gate account.
require.Equal(t, []int64{1}, accountIDs(scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)))
require.Eventually(t, func() bool {
filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)
return len(filtered) == 0
}, 2*time.Second, 10*time.Millisecond)
repo.mu.Lock()
repo.stats[1] = &usagestats.AccountStats{Tokens: 100_000}
repo.mu.Unlock()
require.Empty(t, scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts), "fresh cache keeps the short soft-gate hold")
// Fresh positive cache still holds the soft-gate until TTL expires.
require.Empty(t, scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts), "fresh cache keeps the soft-gate hold")
// Expire entry → miss fails open and schedules refresh with recovered usage.
scheduler.grokFreeQuotaGateCache.Store(int64(1), grokFreeQuotaGateCacheEntry{
tokens: 490_000, checkedAt: time.Now().Add(-time.Minute), known: true,
})
require.Equal(t, []int64{1}, accountIDs(scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)))
require.Equal(t, 2, repo.calls)
require.Eventually(t, func() bool {
filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)
return len(filtered) == 1 && filtered[0].ID == 1
}, 2*time.Second, 10*time.Millisecond)
require.GreaterOrEqual(t, repo.calls, 2)
}
func TestResolveGrokFreeQuotaGateSettingsDefaultsToNinetyFivePercent(t *testing.T) {
@@ -159,6 +190,7 @@ func TestIsExplicitGrokFreeOAuthAccount_OnlyExactFree(t *testing.T) {
func TestOpenAIAccountSchedulerLoadBalanceAppliesGrokFreeQuotaGate(t *testing.T) {
cfg := grokFreeQuotaTestConfig()
cfg.RunMode = config.RunModeSimple
openaiGrokFreeQuotaGateCache = sync.Map{}
accounts := []Account{
{ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Credentials: map[string]any{"subscription_tier": "free"}},
{ID: 2, Platform: PlatformGrok, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, Concurrency: 1, Credentials: map[string]any{"subscription_tier": "pro"}},
@@ -172,6 +204,13 @@ func TestOpenAIAccountSchedulerLoadBalanceAppliesGrokFreeQuotaGate(t *testing.T)
}
scheduler := &defaultOpenAIAccountScheduler{service: svc, stats: newOpenAIAccountRuntimeStats()}
// Warm cache via background refresh so load-balance sees the soft-gate.
_ = scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)
require.Eventually(t, func() bool {
filtered := scheduler.filterGrokFreeQuotaAccounts(context.Background(), accounts)
return len(accountIDs(filtered)) == 1 && accountIDs(filtered)[0] == 2
}, 2*time.Second, 10*time.Millisecond)
selection, _, _, _, err := scheduler.selectByLoadBalance(context.Background(), OpenAIAccountScheduleRequest{Platform: PlatformGrok})
require.NoError(t, err)
require.NotNil(t, selection)
@@ -191,9 +230,13 @@ func TestGrokFreeQuotaGateIsSchedulerOnlyAdminPathUnfiltered(t *testing.T) {
repo := &grokFreeQuotaUsageRepoStub{stats: map[int64]*usagestats.AccountStats{
9: {Tokens: 500_000},
}}
openaiGrokFreeQuotaGateCache = sync.Map{}
scheduler := &defaultOpenAIAccountScheduler{service: &OpenAIGatewayService{cfg: grokFreeQuotaTestConfig(), usageLogRepo: repo}}
overGate := Account{ID: 9, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "FREE"}}
require.Empty(t, scheduler.filterGrokFreeQuotaAccounts(context.Background(), []Account{overGate}))
require.Eventually(t, func() bool {
_ = scheduler.filterGrokFreeQuotaAccounts(context.Background(), []Account{overGate})
return len(scheduler.filterGrokFreeQuotaAccounts(context.Background(), []Account{overGate})) == 0
}, 2*time.Second, 10*time.Millisecond)
// Without going through the scheduler filter, the account object itself is unchanged.
require.True(t, isExplicitGrokFreeOAuthAccount(&overGate))
require.Equal(t, int64(9), overGate.ID)
@@ -240,13 +283,15 @@ func TestFilterGrokFreeQuotaAccountsEvictsDepartedAccounts(t *testing.T) {
accounts := []Account{
{ID: 1, Platform: PlatformGrok, Type: AccountTypeOAuth, Credentials: map[string]any{"subscription_tier": "FREE"}},
}
// First call schedules async refresh + may not have finished sweep yet.
_ = filterGrokFreeQuotaAccountsCore(context.Background(), grokFreeQuotaTestConfig(), repo, &cache, accounts)
require.Eventually(t, func() bool {
_, departedStillCached := cache.Load(int64(99))
_, freshCached := cache.Load(int64(1))
return !departedStillCached && freshCached
}, 2*time.Second, 10*time.Millisecond)
filtered := filterGrokFreeQuotaAccountsCore(context.Background(), grokFreeQuotaTestConfig(), repo, &cache, accounts)
require.Equal(t, []int64{1}, accountIDs(filtered))
_, departedStillCached := cache.Load(int64(99))
require.False(t, departedStillCached)
_, freshCached := cache.Load(int64(1))
require.True(t, freshCached)
}
func accountIDs(accounts []Account) []int64 {
@@ -82,4 +82,37 @@ func TestGroupMediaPricingLooksIncomplete_VideoModelPricesComplete(t *testing.T)
"grok-imagine-video": {"720p": 0.1},
},
}))
price := 10.0
require.False(t, groupMediaPricingLooksIncomplete(&Group{SearchPricePer1k: &price}))
require.False(t, groupMediaPricingLooksIncomplete(&Group{AudioRealtimePricePerMin: &price}))
// Legacy video price alone still marks complete (existing path).
require.False(t, groupMediaPricingLooksIncomplete(&Group{VideoPrice720P: &price}))
}
func TestCalculateOpenAIRecordUsageCost_TokenPricingErrorNotSwallowedBySearch(t *testing.T) {
t.Parallel()
price := 10.0
svc := &OpenAIGatewayService{
billingService: newTestBillingService(),
}
apiKey := &APIKey{
Group: &Group{SearchPricePer1k: &price},
}
// Unknown model → token pricing fails; search must not replace that with $0/$search bill.
cost, err := svc.calculateOpenAIRecordUsageCost(
context.Background(),
&OpenAIForwardResult{SearchCount: 100},
apiKey,
[]string{"totally-unknown-model-xyz-no-pricing"},
1.0,
1.0,
1.0,
1.0,
UsageTokens{InputTokens: 1000, OutputTokens: 500},
"",
false,
)
require.Error(t, err)
require.Nil(t, cost)
}
@@ -496,14 +496,13 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
lastErr = err
}
}
// Search-only (e.g. no model / zero tokens): still bill search when priced.
// Search surcharge is additive. Never let a zero/default search cost mask a
// real token-pricing failure for requests that attempted token billing.
searchCost := (*CostBreakdown)(nil)
if result != nil && result.SearchCount > 0 {
price := groupSearchPricePer1kFromAPIKey(apiKey)
if price == nil || *price <= 0 {
// Silent free search is a revenue leak; error-level so ops/alerts notice.
// Billing still proceeds at $0 so requests are not failed mid-flight.
logger.L().Error("openai_usage.search_price_per_1k_unset_free",
if price != nil && *price == 0 {
logger.L().Info("openai_usage.search_price_per_1k_explicit_free",
zap.Int("search_count", result.SearchCount),
zap.String("model", billingModel),
zap.Int64("api_key_id", apiKey.ID),
@@ -513,18 +512,23 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
searchCost = s.billingService.CalculateSearchCost(result.SearchCount, price, webSearchMultiplier)
}
if tokenCost == nil && searchCost == nil {
if lastErr == nil {
if len(billingModels) == 0 || billingModel == "" {
return nil, errors.New("openai usage billing model is empty")
tokenBillingAttempted := len(billingModels) > 0 && billingModel != ""
if tokenCost == nil {
if tokenBillingAttempted {
if lastErr == nil {
lastErr = errors.New("no non-empty billing model candidates")
}
lastErr = errors.New("no non-empty billing model candidates")
return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr)
}
// Search-only (no model / pure tool path): allow search billing alone.
if searchCost != nil {
return searchCost, nil
}
if lastErr == nil {
lastErr = errors.New("openai usage billing model is empty")
}
return nil, fmt.Errorf("calculate OpenAI usage cost failed for billing models %s: %w", strings.Join(billingModels, ","), lastErr)
}
if tokenCost == nil {
return searchCost, nil
}
if searchCost == nil || (searchCost.TotalCost == 0 && searchCost.ActualCost == 0) {
return tokenCost, nil
}
@@ -705,10 +709,14 @@ func (s *OpenAIGatewayService) apiKeyWithFreshGroupMediaPricing(ctx context.Cont
return &clone
}
// groupMediaPricingLooksIncomplete 判断分组对象是否可能缺失媒体计费字段(例如由不含
// 这些字段的旧快照或手工构造的上下文对象生成)。image/video 独立倍率在数据库中的
// 默认值均为 1.0,正常加载的分组不可能两个倍率同时为 0 且未开启独立倍率、全部媒体
// 价为 nil——只有这种情况才回源查库,避免对未配置覆盖价的分组每条媒体用量都多打一次 DB 查询。
// groupMediaPricingLooksIncomplete 判断分组对象是否可能缺失媒体/搜索/语音计费字段
// (例如由不含这些字段的旧快照或手工构造的上下文对象生成)。image/video 独立倍率在
// 数据库中的默认值均为 1.0;正常加载的分组不可能两个倍率同时为 0 且未开启独立倍率、
// 全部媒体/搜索/语音价为 nil——只有这种情况才回源查库,避免对未配置覆盖价的分组每条
// 用量都多打一次 DB 查询。
//
// 注意:apiKeyAuthSnapshotVersion 升级会强制刷新存量快照;本函数是热路径上的二次兜底,
// 不能仅凭 legacy video_price_* 判定完整而跳过 VideoModelPrices/search/audio 的回源。
func groupMediaPricingLooksIncomplete(group *Group) bool {
if group == nil {
return true
@@ -719,11 +727,17 @@ func groupMediaPricingLooksIncomplete(group *Group) bool {
if group.ImageRateMultiplier != 0 || group.VideoRateMultiplier != 0 {
return false
}
// Per-model video prices are first-class billing config; a projection that
// already carries them is complete enough to skip a DB refresh.
// Any first-class pricing field present means the projection is not a blank shell.
if len(group.VideoModelPrices) > 0 {
return false
}
if group.SearchPricePer1k != nil ||
group.AudioRealtimePricePerMin != nil ||
group.AudioTTSPricePerMillionChars != nil ||
group.AudioSTTPricePerHour != nil ||
group.WebSearchPricePerCall != nil {
return false
}
return group.ImagePrice1K == nil && group.ImagePrice2K == nil && group.ImagePrice4K == nil &&
group.VideoPrice480P == nil && group.VideoPrice720P == nil && group.VideoPrice1080P == nil
}
+4 -2
View File
@@ -190,7 +190,7 @@ func (s *SettingService) InitializeDefaultSettings(ctx context.Context) error {
// Grok: safe defaults — no cross-vendor model rewrite unless operators enable it.
SettingKeyGrokDefaultTextModel: "grok-4.5",
SettingKeyGrokCrossClientModelMapEnabled: "false",
SettingKeyGrokCrossClientModelMapEnabled: "true",
SettingKeyGrokDefaultBaseURLMode: GrokDefaultBaseURLModeCLI,
// Available channels feature (default disabled; opt-in)
@@ -796,7 +796,9 @@ func (s *SettingService) parseSettings(settings map[string]string) *SystemSettin
if result.GrokDefaultTextModel == "" {
result.GrokDefaultTextModel = "grok-4.5"
}
result.GrokCrossClientModelMapEnabled = settings[SettingKeyGrokCrossClientModelMapEnabled] == "true"
// Default true (missing/empty → enabled) so Claude/Codex→Grok mapping keeps working.
// Operators can set false to disable silent cross-client rewrite.
result.GrokCrossClientModelMapEnabled = !isFalseSettingValue(settings[SettingKeyGrokCrossClientModelMapEnabled])
result.GrokDefaultBaseURLMode = normalizeGrokDefaultBaseURLMode(settings[SettingKeyGrokDefaultBaseURLMode])
// Available channels feature (default: disabled; strict true)
@@ -1,5 +1,5 @@
-- Grok Voice 显式定价:realtime / TTS / STT。
-- NULL 表示未配置(当前中继不强制计费;字段预留给分组配置与后续计费接线)。
-- NULL = 使用代码默认单价;显式 0 = 免费;>0 = 分组覆盖价。
ALTER TABLE groups ADD COLUMN IF NOT EXISTS audio_realtime_price_per_min DECIMAL(20,8);
ALTER TABLE groups ADD COLUMN IF NOT EXISTS audio_tts_price_per_million_chars DECIMAL(20,8);
ALTER TABLE groups ADD COLUMN IF NOT EXISTS audio_stt_price_per_hour DECIMAL(20,8);
@@ -1,2 +1,3 @@
-- Grok / 通用搜索工具显式定价(per 1000 calls,USD)。
-- NULL = 使用代码默认 $10/1k;显式 0 = 免费;>0 = 分组覆盖价。
ALTER TABLE groups ADD COLUMN IF NOT EXISTS search_price_per_1k DECIMAL(20,8);
@@ -16,6 +16,7 @@ SELECT id AS group_id,
now() AS backed_up_at
FROM groups
WHERE platform IS DISTINCT FROM 'grok'
AND platform IS DISTINCT FROM 'composite'
AND (
video_price_480p IS NOT NULL
OR video_price_720p IS NOT NULL
@@ -24,7 +25,7 @@ WHERE platform IS DISTINCT FROM 'grok'
);
COMMENT ON TABLE groups_video_price_backup_220 IS
'迁移 220 清空非 Grok 分组视频价前的快照。确认无需回滚后可安全 DROP;回滚方式:UPDATE groups g SET video_price_480p = b.video_price_480p, ... FROM groups_video_price_backup_220 b WHERE g.id = b.group_id';
'迁移 220 清空非 Grok/非 composite 分组视频价前的快照。composite 可能路由到 Grok 账号,予以保留。确认无需回滚后可安全 DROP;回滚方式:UPDATE groups g SET video_price_480p = b.video_price_480p, ... FROM groups_video_price_backup_220 b WHERE g.id = b.group_id';
UPDATE groups
SET video_price_480p = NULL,
@@ -32,6 +33,7 @@ SET video_price_480p = NULL,
video_price_1080p = NULL,
video_model_prices = NULL
WHERE platform IS DISTINCT FROM 'grok'
AND platform IS DISTINCT FROM 'composite'
AND (
video_price_480p IS NOT NULL
OR video_price_720p IS NOT NULL
+6 -4
View File
@@ -444,15 +444,17 @@ gateway:
# Enabled by default because free detection requires an explicit subscription_tier/plan_type of "free".
# Stats/query failures fail open so DB issues do not block all Grok traffic.
grok:
# Email/password authorization is hidden and hard-disabled. Use SSO cookie,
# browser OAuth, or refresh_token re-auth instead. The flag is retained for
# config compatibility only and is ignored by the server.
# Email/password OAuth is off by default and hidden in the admin UI.
# Setting true enables POST /admin/grok/oauth/password (password → SSO → Build OAuth).
# Prefer SSO cookie, browser OAuth, or refresh_token re-auth in production.
password_auth_enabled: false
free_quota_soft_gate_enabled: true
free_quota_token_limit: 500000
free_quota_soft_gate_percent: 95
free_quota_window_hours: 24
free_quota_stats_cache_seconds: 5
# Stats cache for free-tier soft gate. Hot path never waits on DB: misses
# fail open and refresh in the background. Prefer >= 60s in production.
free_quota_stats_cache_seconds: 60
# HTTP upstream connection pool settings (HTTP/2 + multi-proxy scenario defaults)
# HTTP 上游连接池配置(HTTP/2 + 多代理场景默认值)
# Max idle connections across all hosts