fix(grok): harden voice request ids, video pending, and search pricing alerts

Mint durable grok_audio/grok_realtime usage ids, avoid CLI headers on api.x.ai
voice, retry video pending store and fail-closed when snapshot is missing without
status duration, align pure-video ImageCount tests, and escalate unset search
price_per_1k to error-level logs.
This commit is contained in:
IanShaw027
2026-08-08 09:45:12 +08:00
parent 12db0f906a
commit e01ce90d47
7 changed files with 126 additions and 26 deletions
+8
View File
@@ -101,6 +101,8 @@ func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) {
// 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()},
@@ -234,6 +236,12 @@ func (h *OpenAIGatewayHandler) recordGrokVoiceUsage(
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)
+36 -7
View File
@@ -422,7 +422,8 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
}
// Defer billing until status polling observes video.url. Persist create-time
// model/duration/resolution so status can still price if upstream omits them.
if err := h.gatewayService.StoreGrokVideoPendingBilling(requestCtx, result.ResponseID, subject.UserID, apiKey.ID, service.GrokVideoPendingBilling{
// Retry once: missing pending causes silent underpricing (status omits resolution).
pending := service.GrokVideoPendingBilling{
Model: requestModel,
BillingModel: firstNonEmptyString(result.BillingModel, requestModel),
UpstreamModel: result.UpstreamModel,
@@ -431,12 +432,22 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
OriginalModel: clientRequestedModel(c, requestModel),
// Wall-clock start for usage duration_ms: create accepted → first done discovery.
CreatedAt: videoCreateStartedAt,
}); err != nil {
reqLog.Warn("grok_media.store_video_pending_billing_failed",
}
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),
)
}
}
}
// Status poll OR content download can observe official done+video.url.
@@ -540,6 +551,28 @@ func prepareGrokVideoCompletionBilling(
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))
@@ -549,10 +582,6 @@ func prepareGrokVideoCompletionBilling(
reqLog.Debug("grok_media.video_billing_already_claimed", zap.String("request_id", taskRequestID))
return nil
}
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))
}
// Re-merge with pending: resolution is request-only; model/duration fill gaps.
merged := *statusResult
if pending != nil {
@@ -233,6 +233,32 @@ func isForcedUsageBillingRequestID(requestID string) bool {
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
@@ -4,6 +4,7 @@ package service
import (
"context"
"strings"
"testing"
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
@@ -28,5 +29,31 @@ 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)
}
+12 -7
View File
@@ -76,10 +76,11 @@ func (s *OpenAIGatewayService) ForwardGrokVoice(ctx context.Context, c *gin.Cont
contentType = "application/json"
}
req.Header.Set("Content-Type", contentType)
// Voice hits api.x.ai (not CLI proxy). Still stamp CLI identity headers for
// consistency with other Grok outbound calls; transport will not rewrite
// non-CLI-proxy hosts.
applyGrokCLIHeaders(req.Header)
// Match media path: CLI identity headers only on the CLI chat proxy.
// Official api.x.ai voice rejects or mistreats OAuth when CLI headers are stamped.
if account.IsGrokOAuth() && isGrokCLIProxyTarget(targetURL) {
applyGrokCLIHeaders(req.Header)
}
account.ApplyHeaderOverrides(req.Header)
proxyURL := ""
@@ -102,8 +103,10 @@ func (s *OpenAIGatewayService) ForwardGrokVoice(ctx context.Context, c *gin.Cont
}
writeGrokMediaResponse(c, resp, data, s.responseHeaderFilter)
audioUsage := estimateGrokVoiceAudioUsage(baseEndpoint, body, contentType, data, time.Since(started))
upstreamID := firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id"))
return &OpenAIForwardResult{
RequestID: firstNonEmpty(resp.Header.Get("x-request-id"), resp.Header.Get("xai-request-id")),
// Forced durable money-event id so usage_billing_dedup cannot collapse under a reused client id.
RequestID: StableGrokAudioBillingRequestID(upstreamID),
Model: baseEndpoint,
UpstreamModel: baseEndpoint,
Duration: time.Since(started),
@@ -132,8 +135,10 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con
u.Scheme = "wss"
u.RawQuery = "model=" + url.QueryEscape(firstNonEmpty(model, "grok-voice-latest"))
headers := http.Header{"Authorization": []string{"Bearer " + token}}
// Stamp CLI identity for consistency (host is api.x.ai; no CLI-proxy rewrite).
applyGrokCLIHeaders(headers)
// Match media/voice HTTP: CLI headers only on CLI proxy hosts.
if account.IsGrokOAuth() && isGrokCLIProxyTarget(u.String()) {
applyGrokCLIHeaders(headers)
}
if account != nil {
account.ApplyHeaderOverrides(headers)
}
@@ -2045,7 +2045,8 @@ func TestGrokVideoBillingUsesSeparateVideoRateMultiplier(t *testing.T) {
ResponseID: "video-request-123",
Model: "grok-imagine-video-1.5",
BillingModel: "grok-imagine-video-1.5",
ImageCount: 1,
// Pure video completion clears ImageCount (handler contract).
ImageCount: 0,
VideoCount: 1,
VideoResolution: VideoBillingResolution480P,
VideoDurationSeconds: 1,
@@ -2073,7 +2074,7 @@ func TestGrokVideoBillingUsesSeparateVideoRateMultiplier(t *testing.T) {
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, "grok-imagine-video-1.5", usageRepo.lastLog.Model)
require.Equal(t, 1, usageRepo.lastLog.ImageCount)
require.Equal(t, 0, usageRepo.lastLog.ImageCount)
require.Nil(t, usageRepo.lastLog.ImageSize)
require.InDelta(t, 0.08, usageRepo.lastLog.TotalCost, 1e-12)
require.InDelta(t, 0.02, usageRepo.lastLog.ActualCost, 1e-12)
@@ -2098,7 +2099,7 @@ func TestOpenAIGatewayServiceRecordUsage_GrokVideoUsesDefaultRateCard(t *testing
ResponseID: "video-default-rate-card",
Model: "grok-imagine-video-1.5",
BillingModel: "grok-imagine-video-1.5",
ImageCount: 1,
ImageCount: 0,
VideoCount: 1,
VideoResolution: VideoBillingResolution720P,
Duration: time.Second,
@@ -2122,7 +2123,7 @@ func TestOpenAIGatewayServiceRecordUsage_GrokVideoUsesDefaultRateCard(t *testing
// 结果未携带 duration 时按上游默认 8 秒计费:0.14 USD/s × 8s。
require.InDelta(t, 0.14*8, usageRepo.lastLog.TotalCost, 1e-12)
require.InDelta(t, 0.14*8, usageRepo.lastLog.ActualCost, 1e-12)
require.Equal(t, 1, usageRepo.lastLog.ImageCount)
require.Equal(t, 0, usageRepo.lastLog.ImageCount)
require.NotNil(t, usageRepo.lastLog.BillingMode)
require.Equal(t, string(BillingModeVideo), *usageRepo.lastLog.BillingMode)
require.Equal(t, 1, usageRepo.lastLog.VideoCount)
@@ -2165,7 +2166,7 @@ func TestOpenAIGatewayServiceRecordUsage_GroupImagePriceOverridesChannelImagePri
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, 1, usageRepo.lastLog.ImageCount)
require.Equal(t, 0, usageRepo.lastLog.ImageCount)
require.Equal(t, ImageBillingSize2K, *usageRepo.lastLog.ImageSize)
require.InDelta(t, 0.021, usageRepo.lastLog.TotalCost, 1e-12)
require.InDelta(t, 0.021, usageRepo.lastLog.ActualCost, 1e-12)
@@ -2186,7 +2187,7 @@ func TestOpenAIGatewayServiceRecordUsage_GroupVideoPriceOverridesChannelImagePri
RequestID: "resp_grok_video_group_price",
Model: "grok-imagine-video",
BillingModel: "grok-imagine-video",
ImageCount: 1,
ImageCount: 0,
VideoCount: 1,
VideoResolution: VideoBillingResolution720P,
VideoDurationSeconds: 1,
@@ -2210,7 +2211,7 @@ func TestOpenAIGatewayServiceRecordUsage_GroupVideoPriceOverridesChannelImagePri
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, 1, usageRepo.lastLog.ImageCount)
require.Equal(t, 0, usageRepo.lastLog.ImageCount)
require.Nil(t, usageRepo.lastLog.ImageSize)
require.InDelta(t, 0.037, usageRepo.lastLog.TotalCost, 1e-12)
require.InDelta(t, 0.037, usageRepo.lastLog.ActualCost, 1e-12)
@@ -2232,7 +2233,7 @@ func TestOpenAIGatewayServiceRecordUsage_GroupVideoModelPriceOverridesFlatAndCha
RequestID: "resp_grok_video_model_price",
Model: "grok-imagine-video-1.5-preview",
BillingModel: "grok-imagine-video-1.5-preview",
ImageCount: 1,
ImageCount: 0,
VideoCount: 1,
VideoResolution: VideoBillingResolution720P,
VideoDurationSeconds: 2,
@@ -2335,7 +2336,7 @@ func TestOpenAIGatewayServiceRecordUsage_HydratesGroupVideoPriceWhenAuthSnapshot
RequestID: "resp_grok_video_hydrated_price",
Model: "grok-imagine-video",
BillingModel: "grok-imagine-video",
ImageCount: 1,
ImageCount: 0,
VideoCount: 1,
VideoResolution: VideoBillingResolution720P,
VideoDurationSeconds: 1,
@@ -2376,7 +2377,7 @@ func TestOpenAIGatewayServiceRecordUsage_GrokVideoWithTokenChannelPricingKeepsVi
RequestID: "resp_grok_video_token_channel",
Model: "grok-imagine-video",
BillingModel: "grok-imagine-video",
ImageCount: 1,
ImageCount: 0,
VideoCount: 1,
VideoResolution: VideoBillingResolution720P,
VideoDurationSeconds: 5,
@@ -2401,7 +2402,7 @@ func TestOpenAIGatewayServiceRecordUsage_GrokVideoWithTokenChannelPricingKeepsVi
require.NotNil(t, usageRepo.lastLog.BillingMode)
require.Equal(t, string(BillingModeToken), *usageRepo.lastLog.BillingMode)
require.Nil(t, usageRepo.lastLog.ImageSize)
require.Equal(t, 1, usageRepo.lastLog.ImageCount)
require.Equal(t, 0, usageRepo.lastLog.ImageCount)
require.Equal(t, 1, usageRepo.lastLog.VideoCount)
require.NotNil(t, usageRepo.lastLog.VideoResolution)
require.Equal(t, VideoBillingResolution720P, *usageRepo.lastLog.VideoResolution)
@@ -489,9 +489,13 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
if result != nil && result.SearchCount > 0 {
price := groupSearchPricePer1kFromAPIKey(apiKey)
if price == nil || *price <= 0 {
logger.L().Warn("openai_usage.grok_search_price_per_1k_unset_free",
// 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",
zap.Int("search_count", result.SearchCount),
zap.String("model", billingModel),
zap.Int64("api_key_id", apiKey.ID),
zap.Any("group_id", apiKey.GroupID),
)
}
searchCost = s.billingService.CalculateSearchCost(result.SearchCount, price, webSearchMultiplier)