fix: Realtime 仅在观察到音频后计费,并修正标志位求值顺序

原先 elapsed>0 就出账,握手失败也会扣费;随后用 audioObserved,
但又在同一个 return 里先 Load 再收 errCh,标志位恒为 false,会话全部漏计。

- 先等中继结束再读 audioObserved
- 无音频或零时长不出账;每次连接独立 request id
This commit is contained in:
IanShaw027
2026-08-13 08:38:10 +08:00
parent c4d883b8da
commit 678eb22a40
4 changed files with 103 additions and 18 deletions
+14 -11
View File
@@ -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
@@ -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)
}
}
+40 -7
View File
@@ -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.
@@ -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")