mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:08:02 +08:00
fix: Realtime 仅在观察到音频后计费,并修正标志位求值顺序
原先 elapsed>0 就出账,握手失败也会扣费;随后用 audioObserved, 但又在同一个 return 里先 Load 再收 errCh,标志位恒为 false,会话全部漏计。 - 先等中继结束再读 audioObserved - 无音频或零时长不出账;每次连接独立 request id
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user