Files
sub2api/backend/internal/service/openai_gateway_response_handling.go
T

1167 lines
39 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"bufio"
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strconv"
"strings"
"sync/atomic"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
"github.com/Wei-Shaw/sub2api/internal/util/responseheaders"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
)
// openaiStreamingResult streaming response result
type openaiStreamingResult struct {
usage *OpenAIUsage
firstTokenMs *int
responseID string
imageCount int
imageOutputSizes []string
}
type openaiNonStreamingResult struct {
*OpenAIUsage
usage *OpenAIUsage
responseID string
imageCount int
imageOutputSizes []string
}
func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, startTime time.Time, originalModel, mappedModel string) (*openaiStreamingResult, error) {
if s.responseHeaderFilter != nil {
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
}
// Set SSE response headers
c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache")
c.Header("Connection", "keep-alive")
c.Header("X-Accel-Buffering", "no")
// Pass through other headers
if v := resp.Header.Get("x-request-id"); v != "" {
c.Header("x-request-id", v)
}
w := c.Writer
flusher, ok := w.(http.Flusher)
if !ok {
return nil, errors.New("streaming not supported")
}
bufferedWriter := bufio.NewWriterSize(w, 4*1024)
flushBuffered := func() error {
if err := bufferedWriter.Flush(); err != nil {
return err
}
flusher.Flush()
return nil
}
usage := &OpenAIUsage{}
imageCounter := newOpenAIImageOutputCounter()
var firstTokenMs *int
responseID := ""
scanner := bufio.NewScanner(resp.Body)
maxLineSize := defaultMaxLineSize
if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
maxLineSize = s.cfg.Gateway.MaxLineSize
}
scanBuf := getSSEScannerBuf64K()
scanner.Buffer(scanBuf[:0], maxLineSize)
streamInterval := time.Duration(0)
if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 {
streamInterval = time.Duration(s.cfg.Gateway.StreamDataIntervalTimeout) * time.Second
}
// 仅监控上游数据间隔超时,不被下游写入阻塞影响
var intervalTicker *time.Ticker
if streamInterval > 0 {
intervalTicker = time.NewTicker(streamInterval)
defer intervalTicker.Stop()
}
var intervalCh <-chan time.Time
if intervalTicker != nil {
intervalCh = intervalTicker.C
}
keepaliveInterval := time.Duration(0)
if s.cfg != nil && s.cfg.Gateway.StreamKeepaliveInterval > 0 {
keepaliveInterval = time.Duration(s.cfg.Gateway.StreamKeepaliveInterval) * time.Second
}
// 下游 keepalive 仅用于防止代理空闲断开
var keepaliveTicker *time.Ticker
if keepaliveInterval > 0 {
keepaliveTicker = time.NewTicker(keepaliveInterval)
defer keepaliveTicker.Stop()
}
var keepaliveCh <-chan time.Time
if keepaliveTicker != nil {
keepaliveCh = keepaliveTicker.C
}
// Track downstream writes separately from upstream reads: pre-output failover
// can buffer response.created / response.in_progress, so keepalive must be
// based on downstream idle time.
lastDownstreamWriteAt := time.Now()
// 仅发送一次错误事件,避免多次写入导致协议混乱。
// 注意:OpenAI `/v1/responses` streaming 事件必须符合 OpenAI Responses schema;
// 否则下游 SDK(例如 OpenCode)会因为类型校验失败而报错。
errorEventSent := false
clientDisconnected := false // 客户端断开后继续 drain 上游以收集 usage
sawTerminalEvent := false
sawFailedEvent := false
failedMessage := ""
clientOutputStarted := false
upstreamRequestID := strings.TrimSpace(resp.Header.Get("x-request-id"))
var streamEarlyErr error
sendErrorEvent := func(reason string) {
if errorEventSent || clientDisconnected {
return
}
errorEventSent = true
payload := `{"type":"error","sequence_number":0,"error":{"type":"upstream_error","message":` + strconv.Quote(reason) + `,"code":` + strconv.Quote(reason) + `}}`
if err := flushBuffered(); err != nil {
clientDisconnected = true
return
}
if _, err := bufferedWriter.WriteString("data: " + payload + "\n\n"); err != nil {
clientDisconnected = true
return
}
if err := flushBuffered(); err != nil {
clientDisconnected = true
return
}
clientOutputStarted = true
lastDownstreamWriteAt = time.Now()
}
needModelReplace := originalModel != mappedModel
streamOutputAccumulator := apicompat.NewBufferedResponseAccumulator()
streamImageOutputs := make([]json.RawMessage, 0, 1)
streamSeenImages := make(map[string]struct{})
resultWithUsage := func() *openaiStreamingResult {
return &openaiStreamingResult{
usage: usage,
firstTokenMs: firstTokenMs,
responseID: responseID,
imageCount: imageCounter.Count(),
imageOutputSizes: imageCounter.Sizes(),
}
}
finalizeStream := func() (*openaiStreamingResult, error) {
if !sawTerminalEvent {
if !openAIStreamClientOutputStarted(c, clientOutputStarted) {
return resultWithUsage(), s.newOpenAIStreamFailoverError(
c,
account,
false,
upstreamRequestID,
nil,
"OpenAI stream ended before a terminal event",
)
}
return resultWithUsage(), fmt.Errorf("stream usage incomplete: missing terminal event")
}
if sawFailedEvent {
return resultWithUsage(), fmt.Errorf("upstream response failed: %s", failedMessage)
}
if !clientDisconnected {
hadBufferedData := bufferedWriter.Buffered() > 0
if err := flushBuffered(); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during final flush, returning collected usage")
} else if hadBufferedData {
clientOutputStarted = true
lastDownstreamWriteAt = time.Now()
}
}
return resultWithUsage(), nil
}
handleScanErr := func(scanErr error) (*openaiStreamingResult, error, bool) {
if scanErr == nil {
return nil, nil, false
}
if sawTerminalEvent && !sawFailedEvent {
logger.LegacyPrintf("service.openai_gateway", "Upstream scan ended after terminal event: %v", scanErr)
return resultWithUsage(), nil, true
}
if sawFailedEvent {
return resultWithUsage(), fmt.Errorf("upstream response failed: %s", failedMessage), true
}
// 客户端断开/取消请求时,上游读取往往会返回 context canceled。
// /v1/responses 的 SSE 事件必须符合 OpenAI 协议;这里不注入自定义 error event,避免下游 SDK 解析失败。
if errors.Is(scanErr, context.Canceled) || errors.Is(scanErr, context.DeadlineExceeded) {
return resultWithUsage(), fmt.Errorf("stream usage incomplete: %w", scanErr), true
}
if errors.Is(scanErr, bufio.ErrTooLong) {
logger.LegacyPrintf("service.openai_gateway", "SSE line too long: account=%d max_size=%d error=%v", account.ID, maxLineSize, scanErr)
sendErrorEvent("response_too_large")
return resultWithUsage(), scanErr, true
}
if !openAIStreamClientOutputStarted(c, clientOutputStarted) {
msg := "OpenAI stream disconnected before completion"
if errText := strings.TrimSpace(scanErr.Error()); errText != "" {
msg += ": " + errText
}
return resultWithUsage(), s.newOpenAIStreamFailoverError(c, account, false, upstreamRequestID, nil, msg), true
}
// 客户端已断开时,上游出错仅影响体验,不影响计费;返回已收集 usage
if clientDisconnected {
return resultWithUsage(), fmt.Errorf("stream usage incomplete after disconnect: %w", scanErr), true
}
sendErrorEvent("stream_read_error")
return resultWithUsage(), fmt.Errorf("stream read error: %w", scanErr), true
}
processSSELine := func(line string, queueDrained bool) {
if streamEarlyErr != nil {
return
}
// Extract data from SSE line (supports both "data: " and "data:" formats)
if data, ok := extractOpenAISSEDataLine(line); ok {
dataBytes := []byte(data)
if openAIStreamEventIsTerminal(data) {
sawTerminalEvent = true
}
eventType := strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String())
if responseID == "" {
responseID = extractOpenAIResponseIDFromJSONBytes(dataBytes)
}
forceFlushFailedEvent := false
if eventType == "response.failed" {
failedMessage = extractOpenAISSEErrorMessage(dataBytes)
// response.failed 自带上游已消耗的 usage(input token 通常已扣);必须先解析
// 再打 cyber 标记,否则 mark 记到的是解析前的 0,导致流式 cyber 按 0 token 计费
// 而漏记真实用量。对齐 WS V2 / Chat 流式路径(均先解析 usage 再 Mark)。
s.parseSSEUsageBytes(dataBytes, usage)
if hit, code, msg := detectOpenAICyberPolicy(dataBytes); hit {
MarkOpsCyberPolicy(c, CyberPolicyMark{
Code: code,
Message: msg,
Body: truncateString(string(dataBytes), 4096),
UpstreamStatus: http.StatusOK,
UpstreamInTok: usage.InputTokens,
UpstreamOutTok: usage.OutputTokens,
})
}
if !openAIStreamClientOutputStarted(c, clientOutputStarted) {
if status, errType, errMsg, matched := applyOpenAIStreamFailedErrorPassthroughRule(c, dataBytes, failedMessage); matched {
sawFailedEvent = true
MarkResponseCommitted(c)
c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8")
c.JSON(status, gin.H{
"error": gin.H{
"type": errType,
"message": errMsg,
},
})
streamEarlyErr = fmt.Errorf("upstream response failed: passthrough rule matched message=%s", errMsg)
return
}
if openAIStreamFailedEventShouldFailover(dataBytes, failedMessage) {
sawFailedEvent = true
streamEarlyErr = s.newOpenAIStreamFailoverError(c, account, false, upstreamRequestID, dataBytes, failedMessage)
return
}
}
forceFlushFailedEvent = true
sawFailedEvent = true
}
imageCounter.AddSSEData(dataBytes)
// Correct Codex tool calls if needed (apply_patch -> edit, etc.)
if correctedData, corrected := s.toolCorrector.CorrectToolCallsInSSEBytes(dataBytes); corrected {
dataBytes = correctedData
data = string(correctedData)
line = "data: " + data
eventType = strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String())
}
if imageOutput, ok := extractImageGenerationOutputFromSSEData(dataBytes, streamSeenImages); ok {
streamImageOutputs = append(streamImageOutputs, imageOutput)
}
if responsesStreamEventMayContributeToOutput(eventType) {
var streamEvent apicompat.ResponsesStreamEvent
if err := json.Unmarshal(dataBytes, &streamEvent); err == nil {
streamOutputAccumulator.ProcessEvent(&streamEvent)
}
}
if normalizedData, normalized := normalizeResponsesStreamingTerminalOutput(dataBytes, streamOutputAccumulator, streamImageOutputs); normalized {
dataBytes = normalizedData
data = string(normalizedData)
line = "data: " + data
eventType = strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String())
}
if sanitizedData, sanitized := sanitizeOpenAIResponseFailedEventForClient(
dataBytes,
eventType,
openAIStreamClientOutputStarted(c, clientOutputStarted),
); sanitized {
dataBytes = sanitizedData
data = string(sanitizedData)
line = "data: " + data
}
// Replace model in response if needed.
// Fast path: most events do not contain model field values.
if needModelReplace && mappedModel != "" && strings.Contains(line, mappedModel) {
line = s.replaceModelInSSELine(line, mappedModel, originalModel)
}
startsClientOutput := forceFlushFailedEvent || openAIStreamDataStartsClientOutput(data, eventType)
// 写入客户端(客户端断开后继续 drain 上游)
if !clientDisconnected {
shouldFlush := queueDrained && (clientOutputStarted || startsClientOutput)
if firstTokenMs == nil && startsClientOutput {
// 保证首个 token 事件尽快出站,避免影响 TTFT。
shouldFlush = true
}
if _, err := bufferedWriter.WriteString(line); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming, continuing to drain upstream for billing")
} else if _, err := bufferedWriter.WriteString("\n"); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming, continuing to drain upstream for billing")
} else if shouldFlush {
if err := flushBuffered(); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming flush, continuing to drain upstream for billing")
} else {
clientOutputStarted = true
lastDownstreamWriteAt = time.Now()
}
}
}
// Record first token time
if firstTokenMs == nil && startsClientOutput {
ms := int(time.Since(startTime).Milliseconds())
firstTokenMs = &ms
}
s.parseSSEUsageBytes(dataBytes, usage)
return
}
// Forward non-data lines as-is
if !clientDisconnected {
if _, err := bufferedWriter.WriteString(line); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming, continuing to drain upstream for billing")
} else if _, err := bufferedWriter.WriteString("\n"); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming, continuing to drain upstream for billing")
} else if queueDrained && clientOutputStarted {
if err := flushBuffered(); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming flush, continuing to drain upstream for billing")
} else {
clientOutputStarted = true
lastDownstreamWriteAt = time.Now()
}
}
}
}
// 无超时/无 keepalive 的常见路径走同步扫描,减少 goroutine 与 channel 开销。
if streamInterval <= 0 && keepaliveInterval <= 0 {
defer putSSEScannerBuf64K(scanBuf)
for scanner.Scan() {
processSSELine(scanner.Text(), true)
if streamEarlyErr != nil {
return resultWithUsage(), streamEarlyErr
}
}
if result, err, done := handleScanErr(scanner.Err()); done {
return result, err
}
return finalizeStream()
}
type scanEvent struct {
line string
err error
}
// 独立 goroutine 读取上游,避免读取阻塞影响 keepalive/超时处理
events := make(chan scanEvent, 16)
done := make(chan struct{})
sendEvent := func(ev scanEvent) bool {
select {
case events <- ev:
return true
case <-done:
return false
}
}
var lastReadAt int64
atomic.StoreInt64(&lastReadAt, time.Now().UnixNano())
go func(scanBuf *sseScannerBuf64K) {
defer putSSEScannerBuf64K(scanBuf)
defer close(events)
for scanner.Scan() {
atomic.StoreInt64(&lastReadAt, time.Now().UnixNano())
if !sendEvent(scanEvent{line: scanner.Text()}) {
return
}
}
if err := scanner.Err(); err != nil {
_ = sendEvent(scanEvent{err: err})
}
}(scanBuf)
defer close(done)
for {
select {
case ev, ok := <-events:
if !ok {
return finalizeStream()
}
if result, err, done := handleScanErr(ev.err); done {
return result, err
}
processSSELine(ev.line, len(events) == 0)
if streamEarlyErr != nil {
return resultWithUsage(), streamEarlyErr
}
case <-intervalCh:
lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt))
if time.Since(lastRead) < streamInterval {
continue
}
if clientDisconnected {
return resultWithUsage(), fmt.Errorf("stream usage incomplete after timeout")
}
logger.LegacyPrintf("service.openai_gateway", "Stream data interval timeout: account=%d model=%s interval=%s", account.ID, originalModel, streamInterval)
// 处理流超时,可能标记账户为临时不可调度或错误状态
if s.rateLimitService != nil {
s.rateLimitService.HandleStreamTimeout(ctx, account, originalModel)
}
sendErrorEvent("stream_timeout")
return resultWithUsage(), fmt.Errorf("stream data interval timeout")
case <-keepaliveCh:
if clientDisconnected {
continue
}
if time.Since(lastDownstreamWriteAt) < keepaliveInterval {
continue
}
if _, err := bufferedWriter.WriteString(":\n\n"); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming, continuing to drain upstream for billing")
continue
}
if err := flushBuffered(); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during keepalive flush, continuing to drain upstream for billing")
} else {
lastDownstreamWriteAt = time.Now()
}
}
}
}
// extractOpenAISSEDataLine 低开销提取 SSE `data:` 行内容。
// 兼容 `data: xxx` 与 `data:xxx` 两种格式。
func extractOpenAISSEDataLine(line string) (string, bool) {
if !strings.HasPrefix(line, "data:") {
return "", false
}
start := len("data:")
for start < len(line) {
if line[start] != ' ' && line[start] != ' ' {
break
}
start++
}
return line[start:], true
}
func extractOpenAISSEEventLine(line string) (string, bool) {
if !strings.HasPrefix(line, "event:") {
return "", false
}
start := len("event:")
for start < len(line) {
if line[start] != ' ' && line[start] != ' ' {
break
}
start++
}
return strings.TrimSpace(line[start:]), true
}
type openAICompatSSEFrame struct {
EventType string
Data string
}
type openAICompatSSEFrameParser struct {
eventType string
dataLines []string
}
func (p *openAICompatSSEFrameParser) AddLine(line string) (openAICompatSSEFrame, bool) {
if line == "" {
return p.dispatch()
}
if strings.HasPrefix(line, ":") {
return openAICompatSSEFrame{}, false
}
if eventType, ok := extractOpenAISSEEventLine(line); ok {
p.eventType = eventType
return openAICompatSSEFrame{}, false
}
if data, ok := extractOpenAISSEDataLine(line); ok {
p.dataLines = append(p.dataLines, data)
}
return openAICompatSSEFrame{}, false
}
func (p *openAICompatSSEFrameParser) Finish() (openAICompatSSEFrame, bool) {
return p.dispatch()
}
func (p *openAICompatSSEFrameParser) dispatch() (openAICompatSSEFrame, bool) {
frame := openAICompatSSEFrame{
EventType: p.eventType,
Data: strings.Join(p.dataLines, "\n"),
}
p.eventType = ""
p.dataLines = nil
return frame, frame.Data != ""
}
func openAICompatPayloadWithEventType(payload, eventType string) string {
eventType = strings.TrimSpace(eventType)
if eventType == "" || strings.TrimSpace(payload) == "" || strings.TrimSpace(payload) == "[DONE]" {
return payload
}
if gjson.Get(payload, "type").Exists() {
return payload
}
patched, err := sjson.Set(payload, "type", eventType)
if err != nil {
return payload
}
return patched
}
func (s *OpenAIGatewayService) replaceModelInSSELine(line, fromModel, toModel string) string {
data, ok := extractOpenAISSEDataLine(line)
if !ok {
return line
}
if data == "" || data == "[DONE]" {
return line
}
// 使用 gjson 精确检查 model 字段,避免全量 JSON 反序列化
if m := gjson.Get(data, "model"); m.Exists() && m.Str == fromModel {
newData, err := sjson.Set(data, "model", toModel)
if err != nil {
return line
}
return "data: " + newData
}
// 检查嵌套的 response.model 字段
if m := gjson.Get(data, "response.model"); m.Exists() && m.Str == fromModel {
newData, err := sjson.Set(data, "response.model", toModel)
if err != nil {
return line
}
return "data: " + newData
}
return line
}
// correctToolCallsInResponseBody 修正响应体中的工具调用
func (s *OpenAIGatewayService) correctToolCallsInResponseBody(body []byte) []byte {
if len(body) == 0 {
return body
}
updated := body
if s != nil && s.toolCorrector != nil {
if corrected, changed := s.toolCorrector.CorrectToolCallsInSSEBytes(updated); changed {
updated = corrected
}
}
if normalized, changed := normalizeOpenAIResponsesFunctionCallArguments(updated); changed {
updated = normalized
}
return updated
}
func normalizeOpenAIResponsesFunctionCallArguments(data []byte) ([]byte, bool) {
if len(bytes.TrimSpace(data)) == 0 || !bytes.Contains(data, []byte(`"arguments"`)) {
return data, false
}
if !gjson.ValidBytes(data) {
return data, false
}
updated := data
changed := false
setDedupedArgument := func(path string) {
arg := gjson.GetBytes(updated, path)
if !arg.Exists() || arg.Type != gjson.String {
return
}
deduped, ok := dedupeRepeatedJSONArgumentString(arg.Str)
if !ok {
return
}
next, err := sjson.SetBytes(updated, path, deduped)
if err != nil {
return
}
updated = next
changed = true
}
eventType := strings.TrimSpace(gjson.GetBytes(updated, "type").String())
if eventType == "response.function_call_arguments.done" {
setDedupedArgument("arguments")
}
if itemType := strings.TrimSpace(gjson.GetBytes(updated, "item.type").String()); isResponsesFunctionCallItemType(itemType) {
setDedupedArgument("item.arguments")
}
dedupeResponsesFunctionCallOutputArguments(updated, "response.output", setDedupedArgument)
dedupeResponsesFunctionCallOutputArguments(updated, "output", setDedupedArgument)
return updated, changed
}
func dedupeResponsesFunctionCallOutputArguments(data []byte, outputPath string, setDedupedArgument func(string)) {
output := gjson.GetBytes(data, outputPath)
if !output.Exists() || !output.IsArray() {
return
}
for i, item := range output.Array() {
if !isResponsesFunctionCallItemType(strings.TrimSpace(item.Get("type").String())) {
continue
}
setDedupedArgument(outputPath + "." + strconv.Itoa(i) + ".arguments")
}
}
func isResponsesFunctionCallItemType(itemType string) bool {
return itemType == "function_call" || itemType == "custom_tool_call"
}
func dedupeRepeatedJSONArgumentString(arguments string) (string, bool) {
if len(arguments) == 0 || len(arguments)%2 != 0 {
return "", false
}
halfLen := len(arguments) / 2
first := arguments[:halfLen]
if first != arguments[halfLen:] {
return "", false
}
trimmed := strings.TrimSpace(first)
if trimmed == "" || (!strings.HasPrefix(trimmed, "{") && !strings.HasPrefix(trimmed, "[")) {
return "", false
}
if !json.Valid([]byte(first)) {
return "", false
}
return first, true
}
func (s *OpenAIGatewayService) parseSSEUsage(data string, usage *OpenAIUsage) {
s.parseSSEUsageBytes([]byte(data), usage)
}
func (s *OpenAIGatewayService) parseSSEUsageBytes(data []byte, usage *OpenAIUsage) {
if usage == nil || len(data) == 0 || bytes.Equal(data, []byte("[DONE]")) {
return
}
// 选择性解析:仅在数据中包含终止事件标识时才进入字段提取。
if len(data) < 72 {
return
}
eventType := gjson.GetBytes(data, "type").String()
if eventType != "response.completed" && eventType != "response.done" && eventType != "response.failed" &&
eventType != "response.incomplete" && eventType != "response.cancelled" && eventType != "response.canceled" {
return
}
if parsedUsage, ok := extractOpenAIUsageFromJSONBytes(data); ok {
*usage = parsedUsage
}
}
func extractOpenAIUsageFromJSONBytes(body []byte) (OpenAIUsage, bool) {
if len(body) == 0 || !gjson.ValidBytes(body) {
return OpenAIUsage{}, false
}
if usage, ok := openAIUsageFromGJSON(gjson.GetBytes(body, "usage")); ok {
return usage, true
}
return openAIUsageFromGJSON(gjson.GetBytes(body, "response.usage"))
}
func extractOpenAIResponseIDFromJSONBytes(body []byte) string {
if len(body) == 0 || !gjson.ValidBytes(body) {
return ""
}
if id := strings.TrimSpace(gjson.GetBytes(body, "id").String()); id != "" {
return id
}
return strings.TrimSpace(gjson.GetBytes(body, "response.id").String())
}
func (s *OpenAIGatewayService) bindHTTPResponseAccount(ctx context.Context, c *gin.Context, account *Account, responseID string) {
if s == nil || account == nil || account.ID <= 0 {
return
}
responseID = strings.TrimSpace(responseID)
if responseID == "" {
return
}
store := s.getOpenAIWSStateStore()
if store == nil {
return
}
groupID := getOpenAIGroupIDFromContext(c)
ttl := s.openAIWSResponseStickyTTL()
logOpenAIWSBindResponseAccountWarn(groupID, account.ID, responseID, store.BindResponseAccount(ctx, groupID, responseID, account.ID, ttl))
}
func openAIUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) {
if !value.Exists() || !value.IsObject() {
return OpenAIUsage{}, false
}
inputTokens := value.Get("input_tokens").Int()
if inputTokens == 0 {
inputTokens = value.Get("prompt_tokens").Int()
}
outputTokens := value.Get("output_tokens").Int()
if outputTokens == 0 {
outputTokens = value.Get("completion_tokens").Int()
}
cacheReadTokens := value.Get("input_tokens_details.cached_tokens").Int()
if cacheReadTokens == 0 {
cacheReadTokens = value.Get("prompt_tokens_details.cached_tokens").Int()
}
imageOutputTokens := value.Get("output_tokens_details.image_tokens").Int()
if imageOutputTokens == 0 {
imageOutputTokens = value.Get("completion_tokens_details.image_tokens").Int()
}
return OpenAIUsage{
InputTokens: int(inputTokens),
OutputTokens: int(outputTokens),
CacheCreationInputTokens: int(value.Get("cache_creation_input_tokens").Int()),
CacheReadInputTokens: int(cacheReadTokens),
ImageOutputTokens: int(imageOutputTokens),
}, true
}
func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, originalModel, mappedModel string) (*openaiNonStreamingResult, error) {
body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError)
if err != nil {
return nil, err
}
// Detect SSE responses for ALL account types via Content-Type header.
// Some OpenAI-compatible upstreams (including other sub2api instances)
// may return SSE even when stream=false was requested.
if isEventStreamResponse(resp.Header) {
return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel)
}
// bodyLooksLikeSSE is a line-level heuristic: real SSE framing requires
// "data:"/"event:" field names at the very start of a physical line. A
// plain bytes.Contains scan would also match ordinary JSON responses
// whose string content merely echoes the literal text "data:" or
// "event:" (e.g. compact tool output), causing those JSON bodies to be
// misrouted into handleSSEToJSON and lose their usage accounting.
bodyLooksLikeSSE := bodyHasSSEFraming(body)
// For OAuth accounts, also fall back to a body-content heuristic because
// the upstream may omit the Content-Type header while still sending SSE.
// This heuristic is NOT applied to API-key accounts to avoid false
// positives on JSON responses that coincidentally contain "data:" or
// "event:" in their text content.
if account.Type == AccountTypeOAuth && bodyLooksLikeSSE {
return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel)
}
usageValue, usageOK := extractOpenAIUsageFromJSONBytes(body)
if !usageOK {
if bodyLooksLikeSSE {
return s.handleSSEToJSON(resp, c, body, originalModel, mappedModel)
}
return nil, fmt.Errorf("parse response: invalid json response")
}
usage := &usageValue
// Replace model in response if needed
if originalModel != mappedModel {
body = s.replaceModelInResponseBody(body, mappedModel, originalModel)
}
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
contentType := "application/json"
if s.cfg != nil && !s.cfg.Security.ResponseHeaders.Enabled {
if upstreamType := resp.Header.Get("Content-Type"); upstreamType != "" {
contentType = upstreamType
}
}
c.Data(resp.StatusCode, contentType, body)
return &openaiNonStreamingResult{
OpenAIUsage: usage,
usage: usage,
responseID: extractOpenAIResponseIDFromJSONBytes(body),
imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body),
imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body),
}, nil
}
func isEventStreamResponse(header http.Header) bool {
contentType := strings.ToLower(header.Get("Content-Type"))
return strings.Contains(contentType, "text/event-stream")
}
// bodyHasSSEFraming reports whether body contains genuine SSE framing by
// scanning for physical lines that begin with the "data:" or "event:"
// field names, per the SSE spec. Unlike a raw substring scan, this does not
// match when those strings only appear embedded inside JSON string values
// (e.g. "data: foo" quoted as part of an assistant text field), since such
// occurrences never start a physical line in a valid JSON encoding.
func bodyHasSSEFraming(body []byte) bool {
for _, line := range bytes.Split(body, []byte("\n")) {
line = bytes.TrimRight(line, "\r")
if bytes.HasPrefix(line, []byte("data:")) || bytes.HasPrefix(line, []byte("event:")) {
return true
}
}
return false
}
func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Context, body []byte, originalModel, mappedModel string) (*openaiNonStreamingResult, error) {
bodyText := string(body)
finalResponse, ok := extractCodexFinalResponse(bodyText)
usage := &OpenAIUsage{}
if ok {
if parsedUsage, parsed := extractOpenAIUsageFromJSONBytes(finalResponse); parsed {
*usage = parsedUsage
}
// When the terminal event has an empty output array, reconstruct
// output from accumulated delta events so the client gets full content.
// gjson Array() returns empty slice for null, missing, or empty arrays.
if len(gjson.GetBytes(finalResponse, "output").Array()) == 0 {
if outputJSON, reconstructed := reconstructResponseOutputFromSSE(bodyText); reconstructed {
if patched, err := sjson.SetRawBytes(finalResponse, "output", outputJSON); err == nil {
finalResponse = patched
}
}
}
body = finalResponse
if originalModel != mappedModel {
body = s.replaceModelInResponseBody(body, mappedModel, originalModel)
}
// Correct tool calls in final response
body = s.correctToolCallsInResponseBody(body)
} else {
terminalType, terminalPayload, terminalOK := extractOpenAISSETerminalEvent(bodyText)
if terminalOK && terminalType == "response.failed" {
msg := extractOpenAISSEErrorMessage(terminalPayload)
if msg == "" {
msg = "Upstream compact response failed"
}
return nil, s.writeOpenAINonStreamingProtocolError(resp, c, msg)
}
usage = s.parseSSEUsageFromBody(bodyText)
if originalModel != mappedModel {
bodyText = s.replaceModelInSSEBody(bodyText, mappedModel, originalModel)
}
body = []byte(bodyText)
}
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
contentType := "application/json; charset=utf-8"
if !ok {
contentType = resp.Header.Get("Content-Type")
if contentType == "" {
contentType = "text/event-stream"
}
}
c.Data(resp.StatusCode, contentType, body)
return &openaiNonStreamingResult{
OpenAIUsage: usage,
usage: usage,
responseID: extractOpenAIResponseIDFromJSONBytes(body),
imageCount: countOpenAIImageOutputsFromSSEBody(bodyText),
imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText),
}, nil
}
func extractOpenAISSETerminalEvent(body string) (string, []byte, bool) {
var terminalType string
var terminalPayload []byte
forEachOpenAISSEDataPayload(body, func(data []byte) {
if terminalPayload != nil {
return
}
eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String())
switch eventType {
case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled":
terminalType = eventType
terminalPayload = append([]byte(nil), data...)
}
})
if terminalPayload != nil {
return terminalType, terminalPayload, true
}
return "", nil, false
}
func extractOpenAISSEErrorMessage(payload []byte) string {
if len(payload) == 0 {
return ""
}
for _, path := range []string{"response.error.message", "error.message", "message"} {
if msg := strings.TrimSpace(gjson.GetBytes(payload, path).String()); msg != "" {
return sanitizeUpstreamErrorMessage(msg)
}
}
return sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(payload)))
}
func sanitizeOpenAIResponseFailedEventForClient(payload []byte, eventType string, clientOutputStarted bool) ([]byte, bool) {
if eventType != "response.failed" || len(payload) == 0 || !gjson.ValidBytes(payload) {
return payload, false
}
updated := payload
if clientOutputStarted && isOpenAIContextWindowError(extractOpenAISSEErrorMessage(payload), payload) {
errorPath := ""
switch {
case gjson.GetBytes(updated, "response.error").Exists():
errorPath = "response.error"
case gjson.GetBytes(updated, "error").Exists():
errorPath = "error"
}
if errorPath != "" {
next, err := sjson.SetBytes(updated, errorPath+".type", "invalid_request_error")
if err != nil {
return payload, false
}
updated = next
next, err = sjson.SetBytes(updated, errorPath+".code", "context_length_exceeded")
if err != nil {
return payload, false
}
updated = next
}
}
if !gjson.GetBytes(updated, "response").Exists() {
return updated, !bytes.Equal(updated, payload)
}
for _, path := range []string{
"response.instructions",
"response.output",
"response.usage",
"response.metadata",
"response.reasoning",
"response.tools",
"response.tool_choice",
"response.parallel_tool_calls",
"response.text",
"response.truncation",
"response.max_output_tokens",
"response.incomplete_details",
} {
next, err := sjson.DeleteBytes(updated, path)
if err != nil {
return payload, false
}
updated = next
}
return updated, !bytes.Equal(updated, payload)
}
func (s *OpenAIGatewayService) writeOpenAINonStreamingProtocolError(resp *http.Response, c *gin.Context, message string) error {
message = sanitizeUpstreamErrorMessage(strings.TrimSpace(message))
if message == "" {
message = "Upstream returned an invalid non-streaming response"
}
setOpsUpstreamError(c, http.StatusBadGateway, message, "")
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8")
c.JSON(http.StatusBadGateway, gin.H{
"error": gin.H{
"type": "upstream_error",
"message": message,
},
})
return fmt.Errorf("non-streaming openai protocol error: %s", message)
}
func extractCodexFinalResponse(body string) ([]byte, bool) {
var finalResponse []byte
forEachOpenAISSEDataPayload(body, func(data []byte) {
if finalResponse != nil {
return
}
eventType := gjson.GetBytes(data, "type").String()
if eventType == "response.done" || eventType == "response.completed" {
if response := gjson.GetBytes(data, "response"); response.Exists() && response.Type == gjson.JSON && response.Raw != "" {
finalResponse = []byte(response.Raw)
}
}
})
if finalResponse != nil {
return finalResponse, true
}
return nil, false
}
func normalizeResponsesStreamingTerminalOutput(data []byte, acc *apicompat.BufferedResponseAccumulator, imageOutputs []json.RawMessage) ([]byte, bool) {
eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String())
switch eventType {
case "response.completed", "response.done", "response.incomplete", "response.cancelled", "response.canceled":
default:
return data, false
}
output := gjson.GetBytes(data, "response.output")
hasAccumulatedOutput := (acc != nil && acc.HasContent()) || len(imageOutputs) > 0
if output.Exists() && output.IsArray() {
if len(output.Array()) > 0 || !hasAccumulatedOutput {
return data, false
}
}
outputJSON := []byte("[]")
if reconstructed, ok := buildResponsesOutputJSON(acc, imageOutputs); ok {
outputJSON = reconstructed
}
updated, err := sjson.SetRawBytes(data, "response.output", outputJSON)
if err != nil {
return data, false
}
return updated, true
}
func responsesStreamEventMayContributeToOutput(eventType string) bool {
switch eventType {
case "response.output_text.delta",
"response.output_item.added",
"response.function_call_arguments.delta",
"response.reasoning_summary_text.delta":
return true
default:
return false
}
}
// reconstructResponseOutputFromSSE scans raw SSE body text for delta events and
// returns a JSON-encoded output array reconstructed from accumulated deltas.
// Returns (nil, false) if no content was found in deltas.
func reconstructResponseOutputFromSSE(bodyText string) ([]byte, bool) {
acc := apicompat.NewBufferedResponseAccumulator()
imageOutputs := make([]json.RawMessage, 0, 1)
seenImages := make(map[string]struct{})
forEachOpenAISSEDataPayload(bodyText, func(data []byte) {
if imageOutput, ok := extractImageGenerationOutputFromSSEData(data, seenImages); ok {
imageOutputs = append(imageOutputs, imageOutput)
}
eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String())
if responsesStreamEventMayContributeToOutput(eventType) {
var event apicompat.ResponsesStreamEvent
if err := json.Unmarshal(data, &event); err == nil {
acc.ProcessEvent(&event)
}
}
})
return buildResponsesOutputJSON(acc, imageOutputs)
}
func buildResponsesOutputJSON(acc *apicompat.BufferedResponseAccumulator, imageOutputs []json.RawMessage) ([]byte, bool) {
if (acc == nil || !acc.HasContent()) && len(imageOutputs) == 0 {
return nil, false
}
var output []json.RawMessage
if acc != nil && acc.HasContent() {
outputJSON, err := json.Marshal(acc.BuildOutput())
if err == nil {
_ = json.Unmarshal(outputJSON, &output)
}
}
output = append(output, imageOutputs...)
if len(output) == 0 {
return nil, false
}
outputJSON, err := json.Marshal(output)
if err != nil {
return nil, false
}
return outputJSON, true
}
func extractImageGenerationOutputFromSSEData(data []byte, seen map[string]struct{}) (json.RawMessage, bool) {
if len(data) == 0 || !gjson.ValidBytes(data) {
return nil, false
}
if gjson.GetBytes(data, "type").String() != "response.output_item.done" {
return nil, false
}
item := gjson.GetBytes(data, "item")
if !item.Exists() || !item.IsObject() || item.Get("type").String() != "image_generation_call" {
return nil, false
}
if strings.TrimSpace(item.Get("result").String()) == "" {
return nil, false
}
key := strings.TrimSpace(item.Get("id").String())
if key == "" {
key = strings.TrimSpace(item.Get("output_format").String()) + "|" + strings.TrimSpace(item.Get("result").String())
}
if key != "" && seen != nil {
if _, exists := seen[key]; exists {
return nil, false
}
seen[key] = struct{}{}
}
return json.RawMessage(item.Raw), true
}
func (s *OpenAIGatewayService) parseSSEUsageFromBody(body string) *OpenAIUsage {
usage := &OpenAIUsage{}
forEachOpenAISSEDataPayload(body, func(data []byte) {
s.parseSSEUsageBytes(data, usage)
})
return usage
}
func (s *OpenAIGatewayService) replaceModelInSSEBody(body, fromModel, toModel string) string {
lines := strings.Split(body, "\n")
for i, line := range lines {
if _, ok := extractOpenAISSEDataLine(line); !ok {
continue
}
lines[i] = s.replaceModelInSSELine(line, fromModel, toModel)
}
return strings.Join(lines, "\n")
}