From 678eb22a405526c85ab226c35f3ed1e113369d0f Mon Sep 17 00:00:00 2001 From: IanShaw027 Date: Thu, 13 Aug 2026 08:38:10 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20Realtime=20=E4=BB=85=E5=9C=A8=E8=A7=82?= =?UTF-8?q?=E5=AF=9F=E5=88=B0=E9=9F=B3=E9=A2=91=E5=90=8E=E8=AE=A1=E8=B4=B9?= =?UTF-8?q?=EF=BC=8C=E5=B9=B6=E4=BF=AE=E6=AD=A3=E6=A0=87=E5=BF=97=E4=BD=8D?= =?UTF-8?q?=E6=B1=82=E5=80=BC=E9=A1=BA=E5=BA=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 原先 elapsed>0 就出账,握手失败也会扣费;随后用 audioObserved, 但又在同一个 return 里先 Load 再收 errCh,标志位恒为 false,会话全部漏计。 - 先等中继结束再读 audioObserved - 无音频或零时长不出账;每次连接独立 request id --- backend/internal/handler/grok_audio.go | 25 +++++----- .../handler/grok_audio_billing_test.go | 27 +++++++++++ backend/internal/service/grok_audio.go | 47 ++++++++++++++++--- backend/internal/service/grok_audio_test.go | 22 +++++++++ 4 files changed, 103 insertions(+), 18 deletions(-) diff --git a/backend/internal/handler/grok_audio.go b/backend/internal/handler/grok_audio.go index 542d70d423..4ac1ae53bf 100644 --- a/backend/internal/handler/grok_audio.go +++ b/backend/internal/handler/grok_audio.go @@ -89,7 +89,7 @@ func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) { model = "grok-voice-latest" } started := time.Now() - proxyErr := h.gatewayService.ProxyGrokRealtime(c.Request.Context(), c, conn, selection.Account, token, model) + audioObserved, 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)) @@ -98,20 +98,23 @@ func (h *OpenAIGatewayHandler) GrokRealtime(c *gin.Context) { 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()}, - } + if result := grokRealtimeBillingResult(model, elapsed, audioObserved); result != nil { h.recordGrokVoiceUsage(c, apiKey, selection.Account, subscription, "realtime", nil, result) } } +func grokRealtimeBillingResult(model string, elapsed time.Duration, audioObserved bool) *service.OpenAIForwardResult { + if !audioObserved || elapsed <= 0 { + return nil + } + return &service.OpenAIForwardResult{ + RequestID: service.StableGrokRealtimeBillingRequestID(""), + Model: model, + Duration: elapsed, + AudioUsage: &service.AudioUsage{Mode: "realtime", DurationOrUnits: elapsed.Minutes()}, + } +} + func isExpectedGrokRealtimeClose(err error) bool { if err == nil { return true diff --git a/backend/internal/handler/grok_audio_billing_test.go b/backend/internal/handler/grok_audio_billing_test.go index 6b8a6c81be..8ec5d08af0 100644 --- a/backend/internal/handler/grok_audio_billing_test.go +++ b/backend/internal/handler/grok_audio_billing_test.go @@ -4,6 +4,7 @@ package handler import ( "testing" + "time" coderws "github.com/coder/websocket" ) @@ -23,3 +24,29 @@ func TestIsExpectedGrokRealtimeClose(t *testing.T) { t.Fatal("policy violations must not be treated as billable normal closes") } } + +func TestGrokRealtimeBillingResultRequiresObservedAudio(t *testing.T) { + if grokRealtimeBillingResult("grok-voice-latest", time.Second, false) != nil { + t.Fatal("a session without observed audio must not be billed") + } + if grokRealtimeBillingResult("grok-voice-latest", 0, true) != nil { + t.Fatal("zero-duration sessions must not be billed") + } +} + +func TestGrokRealtimeBillingResultUsesForcedUniqueID(t *testing.T) { + first := grokRealtimeBillingResult("grok-voice-latest", 90*time.Second, true) + second := grokRealtimeBillingResult("grok-voice-latest", 90*time.Second, true) + if first == nil || second == nil { + t.Fatal("observed audio sessions should be billable") + } + if first.RequestID == "" { + t.Fatalf("unexpected billing request ID %q", first.RequestID) + } + if first.RequestID == second.RequestID { + t.Fatal("independent realtime connections must not share a billing request ID") + } + if first.AudioUsage == nil || first.AudioUsage.Mode != "realtime" || first.AudioUsage.DurationOrUnits != 1.5 { + t.Fatalf("unexpected audio usage: %#v", first.AudioUsage) + } +} diff --git a/backend/internal/service/grok_audio.go b/backend/internal/service/grok_audio.go index 0b16047966..06c8410b8d 100644 --- a/backend/internal/service/grok_audio.go +++ b/backend/internal/service/grok_audio.go @@ -8,6 +8,7 @@ import ( "net/http" "net/url" "strings" + "sync/atomic" "time" coderws "github.com/coder/websocket" @@ -117,20 +118,20 @@ func (s *OpenAIGatewayService) ForwardGrokVoice(ctx context.Context, c *gin.Cont // ProxyGrokRealtime relays JSON Realtime events to xAI's native Voice WS. // Audio is carried as base64 inside JSON events, so preserving the JSON bytes // is sufficient and avoids translating protocol event types. -func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Context, client *coderws.Conn, account *Account, token, model string) error { +func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Context, client *coderws.Conn, account *Account, token, model string) (bool, error) { if s == nil || client == nil || account == nil { - return fmt.Errorf("realtime service, client, and account are required") + return false, fmt.Errorf("realtime service, client, and account are required") } if account.Platform != PlatformGrok { - return fmt.Errorf("account platform %s is not supported for grok realtime", account.Platform) + return false, fmt.Errorf("account platform %s is not supported for grok realtime", account.Platform) } base, err := buildGrokVoiceURL(account, s.cfg, "realtime") if err != nil { - return err + return false, err } u, err := url.Parse(base) if err != nil { - return err + return false, err } u.Scheme = "wss" u.RawQuery = "model=" + url.QueryEscape(firstNonEmpty(model, "grok-voice-latest")) @@ -150,13 +151,14 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con } upstream, _, _, err := dialer.Dial(ctx, u.String(), headers, proxyURL) if err != nil { - return err + return false, err } defer func() { _ = upstream.Close() }() ctx, cancel := context.WithCancel(ctx) defer cancel() errCh := make(chan error, 2) + var audioObserved atomic.Bool // Upstream → client go func() { @@ -166,6 +168,9 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con errCh <- readErr return } + if grokRealtimeEventHasAudio(msg) { + audioObserved.Store(true) + } if writeErr := client.Write(ctx, coderws.MessageText, msg); writeErr != nil { errCh <- writeErr return @@ -184,6 +189,9 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con if kind != coderws.MessageText && kind != coderws.MessageBinary { continue } + if grokRealtimeEventHasAudio(msg) { + audioObserved.Store(true) + } var raw json.RawMessage if unmarshalErr := json.Unmarshal(msg, &raw); unmarshalErr != nil { errCh <- fmt.Errorf("invalid realtime event: %w", unmarshalErr) @@ -196,7 +204,32 @@ func (s *OpenAIGatewayService) ProxyGrokRealtime(ctx context.Context, c *gin.Con } }() - return <-errCh + return awaitGrokRealtimeAudioObserved(errCh, &audioObserved) +} + +func awaitGrokRealtimeAudioObserved(errCh <-chan error, audioObserved *atomic.Bool) (bool, error) { + err := <-errCh + if audioObserved == nil { + return false, err + } + return audioObserved.Load(), err +} + +func grokRealtimeEventHasAudio(msg []byte) bool { + if !gjson.ValidBytes(msg) { + return false + } + eventType := strings.ToLower(strings.TrimSpace(gjson.GetBytes(msg, "type").String())) + if !strings.Contains(eventType, "audio") || strings.Contains(eventType, "transcript") { + return false + } + for _, path := range []string{"audio", "delta", "data"} { + value := gjson.GetBytes(msg, path) + if value.Type == gjson.String && strings.TrimSpace(value.String()) != "" { + return true + } + } + return false } // estimateGrokVoiceAudioUsage derives billing units from the request/response. diff --git a/backend/internal/service/grok_audio_test.go b/backend/internal/service/grok_audio_test.go index 57a60431ad..410af06042 100644 --- a/backend/internal/service/grok_audio_test.go +++ b/backend/internal/service/grok_audio_test.go @@ -2,6 +2,8 @@ package service import ( "context" + "io" + "sync/atomic" "testing" "github.com/Wei-Shaw/sub2api/internal/pkg/xai" @@ -59,6 +61,26 @@ func TestForwardGrokVoice_RejectsNonGrok(t *testing.T) { require.Contains(t, err.Error(), "not supported") } +func TestAwaitGrokRealtimeAudioObservedReadsFlagAfterRelayExits(t *testing.T) { + errCh := make(chan error, 1) + var observed atomic.Bool + go func() { + observed.Store(true) + errCh <- io.EOF + }() + got, err := awaitGrokRealtimeAudioObserved(errCh, &observed) + require.ErrorIs(t, err, io.EOF) + require.True(t, got, "audioObserved must be read after the relay returns, not before <-errCh") +} + +func TestGrokRealtimeEventHasAudio(t *testing.T) { + require.False(t, grokRealtimeEventHasAudio([]byte(`{"type":"session.created"}`))) + require.False(t, grokRealtimeEventHasAudio([]byte(`{"type":"response.audio_transcript.delta","delta":"hi"}`))) + require.False(t, grokRealtimeEventHasAudio([]byte(`{"type":"response.audio.delta","delta":""}`))) + require.True(t, grokRealtimeEventHasAudio([]byte(`{"type":"response.audio.delta","delta":"abc"}`))) + require.True(t, grokRealtimeEventHasAudio([]byte(`{"type":"response.output_audio.delta","audio":"abc"}`))) +} + func TestForwardGrokVoice_RejectsUnknownEndpoint(t *testing.T) { svc := &OpenAIGatewayService{} _, err := svc.ForwardGrokVoice(context.Background(), nil, &Account{Platform: PlatformGrok}, "unknown", []byte(`{}`), "application/json")