mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:28:39 +08:00
fix(grok): filter billing ping response events
This commit is contained in:
@@ -179,11 +179,12 @@ func (s *OpenAIGatewayService) forwardGrokResponses(
|
||||
var firstTokenMs *int
|
||||
responseID := ""
|
||||
if reqStream {
|
||||
maxLineSize := defaultMaxLineSize
|
||||
if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
|
||||
maxLineSize = s.cfg.Gateway.MaxLineSize
|
||||
}
|
||||
resp.Body = newGrokResponsesBillingPingFilterBody(resp.Body, account, maxLineSize)
|
||||
if hasGrokResponsesClientToolMapping(clientToolMapping) {
|
||||
maxLineSize := defaultMaxLineSize
|
||||
if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
|
||||
maxLineSize = s.cfg.Gateway.MaxLineSize
|
||||
}
|
||||
resp.Body = newGrokResponsesClientToolStreamBody(resp.Body, clientToolMapping, maxLineSize)
|
||||
}
|
||||
streamResult, err := s.handleStreamingResponse(ctx, resp, c, account, startTime, originalModel, upstreamModel)
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type grokResponsesBillingPingFilterBody struct {
|
||||
*io.PipeReader
|
||||
source io.Closer
|
||||
closeOnce sync.Once
|
||||
closeErr error
|
||||
}
|
||||
|
||||
func (b *grokResponsesBillingPingFilterBody) Close() error {
|
||||
readerErr := b.PipeReader.Close()
|
||||
sourceErr := b.closeSource()
|
||||
if readerErr != nil {
|
||||
return readerErr
|
||||
}
|
||||
return sourceErr
|
||||
}
|
||||
|
||||
func (b *grokResponsesBillingPingFilterBody) closeSource() error {
|
||||
b.closeOnce.Do(func() { b.closeErr = b.source.Close() })
|
||||
return b.closeErr
|
||||
}
|
||||
|
||||
func newGrokResponsesBillingPingFilterBody(source io.ReadCloser, account *Account, maxLineSize int) io.ReadCloser {
|
||||
if account == nil || account.Platform != PlatformGrok {
|
||||
return source
|
||||
}
|
||||
reader, writer := io.Pipe()
|
||||
body := &grokResponsesBillingPingFilterBody{PipeReader: reader, source: source}
|
||||
go filterGrokResponsesBillingPings(source, writer, body.closeSource, maxLineSize)
|
||||
return body
|
||||
}
|
||||
|
||||
func filterGrokResponsesBillingPings(
|
||||
source io.Reader,
|
||||
destination *io.PipeWriter,
|
||||
closeSource func() error,
|
||||
maxLineSize int,
|
||||
) {
|
||||
defer func() { _ = closeSource() }()
|
||||
if maxLineSize <= 0 {
|
||||
maxLineSize = defaultMaxLineSize
|
||||
}
|
||||
|
||||
scanner := bufio.NewScanner(source)
|
||||
scanBuf := getSSEScannerBuf64K()
|
||||
defer putSSEScannerBuf64K(scanBuf)
|
||||
initialBufferSize := len(scanBuf)
|
||||
if maxLineSize < initialBufferSize {
|
||||
initialBufferSize = maxLineSize
|
||||
}
|
||||
scanner.Buffer(scanBuf[:0:initialBufferSize], maxLineSize)
|
||||
scanner.Split(scanSSELinesPreservingEndings)
|
||||
frame := make([][]byte, 0, 3)
|
||||
|
||||
writeFrame := func() error {
|
||||
if !isGrokResponsesBillingPingFrame(frame) {
|
||||
for _, line := range frame {
|
||||
if _, err := destination.Write(line); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
frame = frame[:0]
|
||||
return nil
|
||||
}
|
||||
|
||||
for scanner.Scan() {
|
||||
line := append([]byte(nil), scanner.Bytes()...)
|
||||
frame = append(frame, line)
|
||||
if len(bytes.TrimSuffix(bytes.TrimSuffix(line, []byte("\n")), []byte("\r"))) == 0 {
|
||||
if err := writeFrame(); err != nil {
|
||||
_ = destination.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(frame) > 0 {
|
||||
if err := writeFrame(); err != nil {
|
||||
_ = destination.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
_ = destination.CloseWithError(fmt.Errorf("filter Grok Responses billing ping: %w", err))
|
||||
return
|
||||
}
|
||||
_ = destination.Close()
|
||||
}
|
||||
|
||||
func scanSSELinesPreservingEndings(data []byte, atEOF bool) (advance int, token []byte, err error) {
|
||||
for index, value := range data {
|
||||
switch value {
|
||||
case '\n':
|
||||
return index + 1, data[:index+1], nil
|
||||
case '\r':
|
||||
if index+1 == len(data) && !atEOF {
|
||||
return 0, nil, nil
|
||||
}
|
||||
if index+1 < len(data) && data[index+1] == '\n' {
|
||||
return index + 2, data[:index+2], nil
|
||||
}
|
||||
return index + 1, data[:index+1], nil
|
||||
}
|
||||
}
|
||||
if atEOF && len(data) > 0 {
|
||||
return len(data), data, nil
|
||||
}
|
||||
return 0, nil, nil
|
||||
}
|
||||
|
||||
func isGrokResponsesBillingPingFrame(rawLines [][]byte) bool {
|
||||
eventType := ""
|
||||
data := ""
|
||||
eventLines := 0
|
||||
dataLines := 0
|
||||
for _, rawLine := range rawLines {
|
||||
line := strings.TrimSuffix(strings.TrimSuffix(string(rawLine), "\n"), "\r")
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
if value, ok := extractOpenAISSEEventLine(line); ok {
|
||||
eventType = value
|
||||
eventLines++
|
||||
continue
|
||||
}
|
||||
if value, ok := extractOpenAISSEDataLine(line); ok {
|
||||
data = value
|
||||
dataLines++
|
||||
continue
|
||||
}
|
||||
return false
|
||||
}
|
||||
if eventType != "ping" || eventLines != 1 || dataLines != 1 {
|
||||
return false
|
||||
}
|
||||
|
||||
var payload struct {
|
||||
Type string `json:"type"`
|
||||
OpenCodeType json.RawMessage `json:"x-opencode-type"`
|
||||
Cost json.RawMessage `json:"cost"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(data), &payload); err != nil || payload.Type != "ping" {
|
||||
return false
|
||||
}
|
||||
if len(payload.OpenCodeType) > 0 {
|
||||
var openCodeType string
|
||||
return json.Unmarshal(payload.OpenCodeType, &openCodeType) == nil && openCodeType == "inference-cost"
|
||||
}
|
||||
var cost string
|
||||
return json.Unmarshal(payload.Cost, &cost) == nil && cost == "0"
|
||||
}
|
||||
@@ -0,0 +1,240 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestGrokResponsesBillingPingFilter(t *testing.T) {
|
||||
input := strings.Join([]string{
|
||||
": upstream keepalive",
|
||||
"",
|
||||
"event: response.output_text.delta",
|
||||
`data: {"type":"response.output_text.delta","delta":"hello"}`,
|
||||
"",
|
||||
"event: future.vendor_event",
|
||||
`data: {"type":"future.vendor_event","value":1}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":"ping","x-opencode-type":"inference-cost","cost":2.75,"input-tokens":42}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":"ping","cost":"0"}`,
|
||||
"",
|
||||
"event: response.completed",
|
||||
`data: {"type":"response.completed","response":{"id":"resp_1","usage":{"input_tokens":3,"output_tokens":5}}}`,
|
||||
"",
|
||||
}, "\n")
|
||||
|
||||
body := newGrokResponsesBillingPingFilterBody(
|
||||
io.NopCloser(strings.NewReader(input)),
|
||||
&Account{Platform: PlatformGrok},
|
||||
defaultMaxLineSize,
|
||||
)
|
||||
output, err := io.ReadAll(body)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, body.Close())
|
||||
|
||||
result := string(output)
|
||||
require.NotContains(t, result, `"x-opencode-type":"inference-cost"`)
|
||||
require.NotContains(t, result, `{"type":"ping","cost":"0"}`)
|
||||
require.Contains(t, result, ": upstream keepalive\n\n")
|
||||
require.Contains(t, result, "event: response.output_text.delta")
|
||||
require.Contains(t, result, `{"type":"response.output_text.delta","delta":"hello"}`)
|
||||
require.Contains(t, result, "event: future.vendor_event")
|
||||
require.Contains(t, result, `{"type":"future.vendor_event","value":1}`)
|
||||
require.Contains(t, result, "event: response.completed")
|
||||
require.Contains(t, result, `"usage":{"input_tokens":3,"output_tokens":5}`)
|
||||
}
|
||||
|
||||
func TestGrokResponsesBillingPingFilterPreservesUnrelatedPingFrames(t *testing.T) {
|
||||
input := strings.Join([]string{
|
||||
"event: ping",
|
||||
`data: {"type":"ping","kind":"keepalive"}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":"ping","cost":2}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":"ping"}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":"ping","cost":0.0001}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":"ping","cost":0}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":"ping","cost":" 0 "}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":"ping","x-opencode-type":"keepalive","cost":0}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":" ping ","cost":0}`,
|
||||
"",
|
||||
"event: ping",
|
||||
`data: {"type":"ping","x-opencode-type":null,"cost":0}`,
|
||||
"",
|
||||
"event: custom",
|
||||
`data: {"type":"ping","x-opencode-type":"inference-cost"}`,
|
||||
"",
|
||||
}, "\n")
|
||||
|
||||
body := newGrokResponsesBillingPingFilterBody(
|
||||
io.NopCloser(strings.NewReader(input)),
|
||||
&Account{Platform: PlatformGrok},
|
||||
defaultMaxLineSize,
|
||||
)
|
||||
output, err := io.ReadAll(body)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, body.Close())
|
||||
require.Equal(t, input, string(output))
|
||||
}
|
||||
|
||||
func TestGrokResponsesBillingPingFilterPreservesFramingAndMalformedInput(t *testing.T) {
|
||||
input := "event: ping\r\ndata: {not-json}\r\n\r\n" +
|
||||
"event: ping\r\ndata: {\"type\":\"ping\",\"cost\":\"0\"} trailing\r\n\r\n" +
|
||||
"event: future.response.event\r\ndata: {\"type\":\"future.response.event\"}"
|
||||
body := newGrokResponsesBillingPingFilterBody(
|
||||
io.NopCloser(strings.NewReader(input)),
|
||||
&Account{Platform: PlatformGrok},
|
||||
defaultMaxLineSize,
|
||||
)
|
||||
output, err := io.ReadAll(body)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, body.Close())
|
||||
require.Equal(t, input, string(output))
|
||||
}
|
||||
|
||||
func TestGrokResponsesBillingPingFilterHandlesBareCRFrames(t *testing.T) {
|
||||
input := "event: ping\rdata: {\"type\":\"ping\",\"cost\":\"0\"}\r\r" +
|
||||
"event: future.event\rdata: {\"type\":\"future.event\"}\r\r"
|
||||
want := "event: future.event\rdata: {\"type\":\"future.event\"}\r\r"
|
||||
body := newGrokResponsesBillingPingFilterBody(
|
||||
io.NopCloser(strings.NewReader(input)),
|
||||
&Account{Platform: PlatformGrok},
|
||||
defaultMaxLineSize,
|
||||
)
|
||||
output, err := io.ReadAll(body)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, body.Close())
|
||||
require.Equal(t, want, string(output))
|
||||
}
|
||||
|
||||
func TestGrokResponsesBillingPingFilterDoesNotFilterNonGrokAccounts(t *testing.T) {
|
||||
input := "event: ping\ndata: {\"type\":\"ping\",\"cost\":\"0\"}\n\n"
|
||||
source := io.NopCloser(strings.NewReader(input))
|
||||
body := newGrokResponsesBillingPingFilterBody(source, &Account{Platform: PlatformOpenAI}, defaultMaxLineSize)
|
||||
|
||||
output, err := io.ReadAll(body)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, body.Close())
|
||||
require.Equal(t, input, string(output))
|
||||
}
|
||||
|
||||
func TestGrokResponsesBillingPingFilterPreservesUsageAndTerminalEvent(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
input := strings.Join([]string{
|
||||
"event: ping",
|
||||
`data: {"type":"ping","x-opencode-type":"inference-cost","cost":"0"}`,
|
||||
"",
|
||||
"event: response.completed",
|
||||
`data: {"type":"response.completed","response":{"id":"resp_1","usage":{"input_tokens":3,"output_tokens":5}}}`,
|
||||
"",
|
||||
}, "\n")
|
||||
account := &Account{ID: 1, Platform: PlatformGrok}
|
||||
resp := &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{},
|
||||
Body: newGrokResponsesBillingPingFilterBody(
|
||||
io.NopCloser(strings.NewReader(input)), account, defaultMaxLineSize,
|
||||
),
|
||||
}
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}},
|
||||
toolCorrector: NewCodexToolCorrector(),
|
||||
}
|
||||
|
||||
result, err := svc.handleStreamingResponse(context.Background(), resp, c, account, time.Now(), "grok-4.5", "grok-4.5")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 3, result.usage.InputTokens)
|
||||
require.Equal(t, 5, result.usage.OutputTokens)
|
||||
require.Equal(t, "resp_1", result.responseID)
|
||||
require.Contains(t, recorder.Body.String(), "response.completed")
|
||||
require.NotContains(t, recorder.Body.String(), "inference-cost")
|
||||
}
|
||||
|
||||
type grokPingFilterTestReadCloser struct {
|
||||
reader io.ReadCloser
|
||||
closeCount atomic.Int32
|
||||
}
|
||||
|
||||
func (r *grokPingFilterTestReadCloser) Read(p []byte) (int, error) { return r.reader.Read(p) }
|
||||
func (r *grokPingFilterTestReadCloser) Close() error {
|
||||
r.closeCount.Add(1)
|
||||
return r.reader.Close()
|
||||
}
|
||||
|
||||
func TestGrokResponsesBillingPingFilterCloseCancelsSourceOnce(t *testing.T) {
|
||||
upstreamReader, upstreamWriter := io.Pipe()
|
||||
source := &grokPingFilterTestReadCloser{reader: upstreamReader}
|
||||
body := newGrokResponsesBillingPingFilterBody(source, &Account{Platform: PlatformGrok}, defaultMaxLineSize)
|
||||
|
||||
require.NoError(t, body.Close())
|
||||
require.Eventually(t, func() bool { return source.closeCount.Load() == 1 }, time.Second, time.Millisecond)
|
||||
_, err := upstreamWriter.Write([]byte("blocked"))
|
||||
require.Error(t, err)
|
||||
require.NoError(t, upstreamWriter.Close())
|
||||
}
|
||||
|
||||
func TestGrokResponsesBillingPingFilterFlushesCompletedFrames(t *testing.T) {
|
||||
upstreamReader, upstreamWriter := io.Pipe()
|
||||
body := newGrokResponsesBillingPingFilterBody(upstreamReader, &Account{Platform: PlatformGrok}, defaultMaxLineSize)
|
||||
defer body.Close()
|
||||
|
||||
go func() {
|
||||
_, _ = io.WriteString(upstreamWriter, "event: future.event\ndata: {\"type\":\"future.event\"}\n\n")
|
||||
}()
|
||||
|
||||
result := make(chan error, 1)
|
||||
go func() {
|
||||
buffer := make([]byte, 64)
|
||||
n, err := body.Read(buffer)
|
||||
if err == nil && !strings.Contains(string(buffer[:n]), "future.event") {
|
||||
err = errors.New("completed frame was not forwarded")
|
||||
}
|
||||
result <- err
|
||||
}()
|
||||
select {
|
||||
case err := <-result:
|
||||
require.NoError(t, err)
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("completed frame was buffered until upstream EOF")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGrokResponsesBillingPingFilterReportsOversizedLine(t *testing.T) {
|
||||
body := newGrokResponsesBillingPingFilterBody(
|
||||
io.NopCloser(strings.NewReader("data: 123456789\n\n")),
|
||||
&Account{Platform: PlatformGrok},
|
||||
8,
|
||||
)
|
||||
_, err := io.ReadAll(body)
|
||||
require.ErrorContains(t, err, "filter Grok Responses billing ping")
|
||||
require.NoError(t, body.Close())
|
||||
}
|
||||
Reference in New Issue
Block a user