Files
sub2api/backend/internal/service/openai_gateway_response_handling.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

1903 lines
67 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
searchCount int
}
type openaiNonStreamingResult struct {
*OpenAIUsage
usage *OpenAIUsage
responseID string
imageCount int
imageOutputSizes []string
searchCount int
}
func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, startTime time.Time, originalModel, mappedModel string) (*openaiStreamingResult, error) {
return s.handleStreamingResponseWithReasoning(ctx, resp, c, account, startTime, originalModel, mappedModel, "")
}
func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, startTime time.Time, originalModel, mappedModel, reasoningEffort string) (*openaiStreamingResult, error) {
observer := upstreamResponseModelObserverFromContext(c)
if observer == nil {
observer = beginUpstreamResponseModelObservation(c)
}
firstOutputTimeout := time.Duration(0)
if account != nil && account.Platform == PlatformOpenAI {
firstOutputTimeout = s.openAIFirstOutputTimeout(reasoningEffort)
}
guardFirstOutput := firstOutputTimeout > 0
stageFirstOutput := account != nil && account.Platform == PlatformOpenAI
var attemptResponseHeaders http.Header
if stageFirstOutput {
if s.responseHeaderFilter != nil {
attemptResponseHeaders = responseheaders.FilterHeaders(resp.Header, s.responseHeaderFilter)
} else if requestID := strings.TrimSpace(resp.Header.Get("x-request-id")); requestID != "" {
attemptResponseHeaders = http.Header{"X-Request-Id": []string{requestID}}
}
} else if s.responseHeaderFilter != nil {
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
}
// x-codex-turn-state 不在通用响应头白名单内,按 Codex 协议显式回传:
// 客户端会在同回合的后续请求中回带(openai_codex_turn_state.go)。
// OpenAI 首个语义输出前只暂存,溯源在 applyAttemptResponseHeaders 真正提交时记录。
if stageFirstOutput {
stageOpenAICodexTurnState(&attemptResponseHeaders, resp.Header)
} else {
s.relayOpenAICodexTurnState(c, account, resp.Header)
}
// 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 !stageFirstOutput && resp.Header.Get("x-request-id") != "" {
v := resp.Header.Get("x-request-id")
c.Header("x-request-id", v)
}
applyAttemptResponseHeaders := func() {
if !stageFirstOutput || len(attemptResponseHeaders) == 0 || c.Writer.Written() {
return
}
for key, values := range attemptResponseHeaders {
for _, value := range values {
c.Writer.Header().Add(key, value)
}
}
// 暂存头此刻才真正写给客户端:turn-state 溯源在这里记录(见
// noteStagedOpenAICodexTurnStateCommitted 的 failover 说明)。
s.noteStagedOpenAICodexTurnStateCommitted(c, account, attemptResponseHeaders)
// These headers describe this gateway's SSE stream and are stable across
// account attempts. Keep them authoritative over upstream values.
c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache")
c.Header("Connection", "keep-alive")
c.Header("X-Accel-Buffering", "no")
}
w := c.Writer
flusher, ok := w.(http.Flusher)
if !ok {
return nil, errors.New("streaming not supported")
}
maxLineSize := defaultMaxLineSize
if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
maxLineSize = s.cfg.Gateway.MaxLineSize
}
var firstTokenMs *int
firstOutputProgressObserved := false
bufferedWriter := bufio.NewWriterSize(w, 4*1024)
var firstOutputStage *openAIFirstOutputStage
if stageFirstOutput {
firstOutputStage = newDefaultOpenAIFirstOutputStage()
defer func() {
if err := firstOutputStage.Close(); err != nil {
logger.LegacyPrintf("service.openai_gateway", "OpenAI first-output staging cleanup failed: account=%d model=%s error=%v", account.ID, originalModel, err)
}
}()
}
writePendingString := func(value string) (int, error) {
if firstOutputStage != nil && !firstOutputStage.closed {
return firstOutputStage.WriteString(value)
}
return bufferedWriter.WriteString(value)
}
pendingBytes := func() int64 {
if firstOutputStage != nil && !firstOutputStage.closed {
return firstOutputStage.Buffered()
}
return int64(bufferedWriter.Buffered())
}
flushBuffered := func() error {
if firstOutputStage != nil && !firstOutputStage.closed {
if err := firstOutputStage.CommitTo(w); err != nil {
return err
}
} else {
if err := bufferedWriter.Flush(); err != nil {
return err
}
}
flusher.Flush()
return nil
}
usage := &OpenAIUsage{}
imageCounter := newOpenAIImageOutputCounter()
responseID := ""
var firstOutputScanGuard atomic.Bool
firstOutputScanGuard.Store(stageFirstOutput)
scanner := bufio.NewScanner(resp.Body)
scanBuf := getSSEScannerBuf64K()
scanner.Buffer(scanBuf[:0], maxLineSize)
if stageFirstOutput {
scanner.Split(openAIFirstOutputDynamicScanLines(&firstOutputScanGuard))
}
documentScanner := newOpenAISSEJSONDocumentScanner(scanner)
streamInterval := time.Duration(0)
if s.cfg != nil && s.cfg.Gateway.StreamDataIntervalTimeout > 0 {
streamInterval = time.Duration(s.cfg.Gateway.StreamDataIntervalTimeout) * time.Second
}
// Grok: always enforce an upstream-read idle so hung SSE bodies fail over
// instead of holding the OAuth slot until the client cancels. Prefer the
// global gateway setting when set; otherwise apply a Grok-only default.
if account != nil && account.Platform == PlatformGrok {
cfgSec := 0
if s.cfg != nil {
cfgSec = s.cfg.Gateway.StreamDataIntervalTimeout
}
streamInterval = resolveGrokStreamIdleTimeout(cfgSec)
}
// 仅监控上游数据间隔超时,不被下游写入阻塞影响
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
}
var firstOutputTimer *time.Timer
var firstOutputCh <-chan time.Time
if firstOutputTimeout > 0 {
remaining := time.Until(startTime.Add(firstOutputTimeout))
if remaining <= 0 {
remaining = time.Nanosecond
}
firstOutputTimer = time.NewTimer(remaining)
firstOutputCh = firstOutputTimer.C
defer firstOutputTimer.Stop()
}
stopFirstOutputTimer := func() {
if firstOutputTimer == nil {
return
}
if !firstOutputTimer.Stop() {
select {
case <-firstOutputTimer.C:
default:
}
}
firstOutputTimer = nil
firstOutputCh = nil
}
// 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
responsesSemanticOutputSeen := false
capacityFailoverSuppressedLogged := false
failedMessage := ""
clientOutputStarted := false
upstreamRequestID := strings.TrimSpace(resp.Header.Get("x-request-id"))
var streamEarlyErr error
eventInProgress := false
eventStartsClientOutput := false
eventStartsVisibleOutput := false
eventShouldFlush := false
handlePendingWriteError := func(err error) {
if firstOutputStage != nil && !firstOutputStage.closed {
message := "OpenAI first-output staging failed"
if errors.Is(err, errOpenAIFirstOutputStageLimit) {
message = "OpenAI first-output staging limit exceeded"
}
logger.LegacyPrintf("service.openai_gateway", "%s: account=%d model=%s error=%v", message, account.ID, originalModel, err)
failoverErr := s.newOpenAIStreamFailoverError(c, account, false, upstreamRequestID, nil, message)
failoverErr.SafeToFailoverAfterWrite = true
streamEarlyErr = failoverErr
_ = resp.Body.Close()
return
}
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming, continuing to drain upstream for billing")
}
completeGuardedEvent := func(queueDrained bool) {
completedProgressEvent := eventStartsClientOutput
completedVisibleEvent := eventStartsVisibleOutput
shouldFlush := eventShouldFlush || (queueDrained && clientOutputStarted)
eventInProgress = false
if !clientDisconnected {
if completedProgressEvent {
applyAttemptResponseHeaders()
}
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()
}
}
}
if completedProgressEvent && !firstOutputProgressObserved {
firstOutputScanGuard.Store(false)
firstOutputProgressObserved = true
stopFirstOutputTimer()
}
if completedVisibleEvent && firstTokenMs == nil {
ms := int(time.Since(startTime).Milliseconds())
firstTokenMs = &ms
}
eventStartsClientOutput = false
eventStartsVisibleOutput = false
eventShouldFlush = false
}
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 := writePendingString("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{})
searchCounter := 0
// Dedup search tool calls across SSE events (item.done + response.completed
// both list the same call_id — counting both would ~2× the surcharge).
streamSearchSeen := make(map[string]struct{})
resultWithUsage := func() *openaiStreamingResult {
return &openaiStreamingResult{
usage: usage,
firstTokenMs: firstTokenMs,
responseID: responseID,
imageCount: imageCounter.Count(),
imageOutputSizes: imageCounter.Sizes(),
searchCount: searchCounter,
}
}
flushPending := func(disconnectMessage string) {
if clientDisconnected || pendingBytes() == 0 {
return
}
if err := flushBuffered(); err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "%s", disconnectMessage)
return
}
clientOutputStarted = true
lastDownstreamWriteAt = time.Now()
}
finalizeStream := func() (*openaiStreamingResult, error) {
if stageFirstOutput && eventInProgress {
// EOF dispatches the final SSE event even without a trailing blank line.
completeGuardedEvent(true)
}
if sawTerminalEvent && !sawFailedEvent {
s.clearOpenAIProxyStreamDisconnect(account)
}
if !sawTerminalEvent && !openAIStreamClientOutputStarted(c, clientOutputStarted) && !eventShouldFlush {
return resultWithUsage(), s.newOpenAIStreamFailoverError(
c,
account,
false,
upstreamRequestID,
nil,
"OpenAI stream ended before a terminal event",
)
}
flushPending("Client disconnected during final flush, returning collected usage")
if !sawTerminalEvent {
if openAIStreamClientOutputStarted(c, clientOutputStarted) && !clientDisconnected {
s.recordOpenAIProxyStreamDisconnect(account, errors.New("stream ended before terminal event"), upstreamRequestID)
}
return resultWithUsage(), fmt.Errorf("stream usage incomplete: missing terminal event")
}
if sawFailedEvent {
return resultWithUsage(), fmt.Errorf("upstream response failed: %s", failedMessage)
}
return resultWithUsage(), nil
}
handleScanErr := func(scanErr error) (*openaiStreamingResult, error, bool) {
if scanErr == nil {
return nil, nil, false
}
if errors.Is(scanErr, errOpenAIFirstOutputScannerLimit) && !firstOutputProgressObserved {
logger.LegacyPrintf("service.openai_gateway", "SSE token exceeded guarded first-output limit: account=%d limit=%d error=%v", account.ID, openAIFirstOutputStageMaxBytes+openAIFirstOutputScannerFramingAllowance, scanErr)
failoverErr := s.newOpenAIStreamFailoverError(
c, account, false, upstreamRequestID, nil,
"OpenAI SSE line exceeds guarded first-output limit",
)
failoverErr.SafeToFailoverAfterWrite = true
return resultWithUsage(), failoverErr, true
}
if errors.Is(scanErr, bufio.ErrTooLong) && stageFirstOutput && !firstOutputProgressObserved {
logger.LegacyPrintf("service.openai_gateway", "SSE line too long before first output: account=%d max_size=%d error=%v", account.ID, maxLineSize, scanErr)
failoverErr := s.newOpenAIStreamFailoverError(
c, account, false, upstreamRequestID, nil,
"OpenAI SSE line exceeds guarded first-output limit",
)
failoverErr.SafeToFailoverAfterWrite = true
return resultWithUsage(), failoverErr, true
}
if sawTerminalEvent {
if !sawFailedEvent {
s.clearOpenAIProxyStreamDisconnect(account)
logger.LegacyPrintf("service.openai_gateway", "Upstream scan ended after terminal event: %v", scanErr)
}
result, err := finalizeStream()
return result, err, true
}
// 客户端断开/取消请求时,上游读取往往会返回 context canceled。
// /v1/responses 的 SSE 事件必须符合 OpenAI 协议;这里不注入自定义 error event,避免下游 SDK 解析失败。
if errors.Is(scanErr, context.Canceled) || errors.Is(scanErr, context.DeadlineExceeded) {
if eventShouldFlush {
flushPending("Client disconnected during canceled stream flush, returning collected usage")
}
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) && !eventShouldFlush {
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
}
s.recordOpenAIProxyStreamDisconnect(account, scanErr, upstreamRequestID)
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)
eventTypeRaw := gjson.GetBytes(dataBytes, "type").String()
eventType := strings.TrimSpace(eventTypeRaw)
observer.ObserveOpenAI(dataBytes, eventTypeRaw)
// 初始上游 data 的 type 只解析一次:原始值保持终止事件的精确匹配,规范化值供后续分支复用。
if openAIStreamEventIsTerminalWithType(data, eventTypeRaw) {
sawTerminalEvent = true
}
if responseID == "" {
responseID = extractOpenAIResponseIDFromJSONBytes(dataBytes)
}
forceFlushFailedEvent := false
if !capacityFailoverSuppressedLogged && account != nil && account.Platform == PlatformOpenAI &&
(eventType == "error" || eventType == "response.failed") &&
openAIStreamClientOutputStarted(c, clientOutputStarted) &&
isOpenAIUpstreamCapacityShedEvent(dataBytes) {
logOpenAICapacityFailoverSuppressed(ctx, account, "native_sse", upstreamRequestID, eventType)
capacityFailoverSuppressedLogged = true
}
if eventType == "error" && !openAIStreamClientOutputStarted(c, clientOutputStarted) {
errorMessage := extractOpenAISSEErrorMessage(dataBytes)
if status, errType, errMsg, matched := applyOpenAIStreamFailedErrorPassthroughRule(c, account.Platform, dataBytes, errorMessage); matched {
s.recordOpenAIStreamUpstreamError(c, account, false, upstreamRequestID, "http_error", dataBytes, errorMessage)
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 error event: passthrough rule matched message=%s", errMsg)
return
}
if openAIStreamErrorEventShouldFailover(dataBytes, errorMessage) {
streamEarlyErr = s.newOpenAIStreamFailoverError(c, account, false, upstreamRequestID, dataBytes, errorMessage, resp.Header)
return
}
}
if eventType == "response.failed" {
failedMessage = extractOpenAISSEErrorMessage(dataBytes)
// response.failed 自带上游已消耗的 usageinput 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, account.Platform, dataBytes, failedMessage); matched {
sawFailedEvent = true
// 命中透传规则也要记录 ops 上游错误事件(对齐 CC/Messages 与
// antigravity 先例),否则透传命中的 failed 在监控中不可见。
s.recordOpenAIStreamUpstreamError(c, account, false, upstreamRequestID, "http_error", dataBytes, failedMessage)
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, resp.Header)
return
}
}
forceFlushFailedEvent = true
sawFailedEvent = true
}
if normalizedData, normalized := normalizeCompletedImageGenerationStatus(dataBytes); normalized {
dataBytes = normalizedData
data = string(normalizedData)
line = "data: " + data
}
imageCounter.AddSSEData(dataBytes)
searchCounter += countGrokNativeSearchCallsInSSEDataDedup(dataBytes, streamSearchSeen)
// 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())
}
restoredData, restoreErr := restoreGrokResponsesClientToolPayload(c, dataBytes)
if restoreErr != nil {
streamEarlyErr = fmt.Errorf("restore Grok Responses client tool response: %w", restoreErr)
return
}
restoredData, restoreErr = restoreOpenAIResponsesNamespacePayload(c, restoredData)
if restoreErr != nil {
streamEarlyErr = fmt.Errorf("restore OpenAI namespace response: %w", restoreErr)
return
}
if !bytes.Equal(restoredData, dataBytes) {
dataBytes = restoredData
data = string(restoredData)
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)
startsVisibleOutput := openAIStreamDataStartsVisibleOutput(data, eventType)
if stageFirstOutput {
eventStartsClientOutput = eventStartsClientOutput || startsClientOutput
eventStartsVisibleOutput = eventStartsVisibleOutput || startsVisibleOutput
if startsClientOutput {
firstOutputScanGuard.Store(false)
}
}
if startsClientOutput && !openAIStreamEventTypeIsTerminal(eventType) {
responsesSemanticOutputSeen = true
}
// OpenAI Responses streams that terminate with an empty
// response.completed (no output, no usage, no error, nothing sent
// to the client) are silent upstream refusals: fail over instead of
// recording a successful 0/0 usage turn (issue #5009).
if account != nil && account.Platform == PlatformOpenAI &&
(eventType == "response.completed" || eventType == "response.done") &&
!sawFailedEvent && !responsesSemanticOutputSeen && !clientOutputStarted &&
openAIResponsesCompletedEventIsEmpty(dataBytes, usage) {
sawTerminalEvent = true
streamEarlyErr = newOpenAIResponsesEmptyCompletedFailoverError(c, account, upstreamRequestID)
return
}
// 写入客户端(客户端断开后继续 drain 上游)
if !clientDisconnected {
shouldFlush := queueDrained && (clientOutputStarted || startsClientOutput)
if firstTokenMs == nil && startsVisibleOutput {
// 保证首个 token 事件尽快出站,避免影响 TTFT。
shouldFlush = true
}
eventShouldFlush = eventShouldFlush || shouldFlush
if _, err := writePendingString(line); err != nil {
handlePendingWriteError(err)
} else if _, err := writePendingString("\n"); err != nil {
handlePendingWriteError(err)
} else {
eventInProgress = true
}
}
// Record first token time
if !guardFirstOutput && firstTokenMs == nil && startsVisibleOutput {
ms := int(time.Since(startTime).Milliseconds())
firstTokenMs = &ms
stopFirstOutputTimer()
}
s.parseSSEUsageBytes(dataBytes, usage)
return
}
// A blank line dispatches a guarded event from the attempt-local stage.
if stageFirstOutput && line == "" {
if !clientDisconnected {
if _, err := writePendingString("\n"); err != nil {
handlePendingWriteError(err)
}
}
if streamEarlyErr == nil {
completeGuardedEvent(queueDrained)
}
return
}
// Non-guarded streams retain upstream's event-boundary flushing: a keepalive
// or queue-drain flush must never split an open SSE event.
shouldFlush := false
if line == "" {
shouldFlush = eventShouldFlush || (queueDrained && clientOutputStarted)
eventShouldFlush = false
}
if !clientDisconnected {
if _, err := writePendingString(line); err != nil {
handlePendingWriteError(err)
} else if _, err := writePendingString("\n"); err != nil {
handlePendingWriteError(err)
} else {
eventInProgress = line != ""
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()
}
}
}
}
}
// 无超时/无 keepalive 的常见路径走同步扫描,减少 goroutine 与 channel 开销。
if streamInterval <= 0 && keepaliveInterval <= 0 && firstOutputTimeout <= 0 {
defer putSSEScannerBuf64K(scanBuf)
for documentScanner.Scan() {
processSSELine(documentScanner.Text(), true)
if streamEarlyErr != nil {
return resultWithUsage(), streamEarlyErr
}
}
if result, err, done := handleScanErr(documentScanner.Err()); done {
return result, err
}
return finalizeStream()
}
type scanEvent struct {
line string
err error
processed chan struct{}
}
// 独立 goroutine 读取上游,避免读取阻塞影响 keepalive/超时处理
// Guard mode permits one queued token plus the token being processed. With
// the guarded scanner cap this bounds scanner/channel retention near 16 MiB;
// the timeout-disabled path preserves the legacy depth of 16.
events := make(chan scanEvent, openAIFirstOutputEventQueueSize(guardFirstOutput))
done := make(chan struct{})
sendEvent := func(ev scanEvent) bool {
if firstOutputScanGuard.Load() {
ev.processed = make(chan struct{})
}
select {
case events <- ev:
case <-done:
return false
}
if ev.processed == nil {
return true
}
select {
case <-ev.processed:
return true
case <-done:
return false
}
}
markEventProcessed := func(ev scanEvent) {
if ev.processed != nil {
close(ev.processed)
}
}
var lastReadAt int64
atomic.StoreInt64(&lastReadAt, time.Now().UnixNano())
go func(scanBuf *sseScannerBuf64K) {
defer putSSEScannerBuf64K(scanBuf)
defer close(events)
for documentScanner.Scan() {
atomic.StoreInt64(&lastReadAt, time.Now().UnixNano())
if !sendEvent(scanEvent{line: documentScanner.Text()}) {
return
}
}
if err := documentScanner.Err(); err != nil {
_ = sendEvent(scanEvent{err: err})
}
}(scanBuf)
defer close(done)
for {
select {
case ev, ok := <-events:
if !ok {
if stageFirstOutput && eventInProgress {
// EOF dispatches the final SSE event even without a trailing blank
// line. Do not synthesize extra bytes on the downstream wire.
completeGuardedEvent(true)
}
return finalizeStream()
}
if result, err, done := handleScanErr(ev.err); done {
markEventProcessed(ev)
return result, err
}
processSSELine(ev.line, len(events) == 0)
markEventProcessed(ev)
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)
}
// Grok: short cool + account failover when no client-visible bytes
// were committed yet (pre-commit). After output started we keep the
// legacy stream_timeout path so partial SSE is not dual-written.
if account != nil && account.Platform == PlatformGrok {
s.tempUnscheduleGrok(ctx, account, grokStreamIdleCooldown, "grok stream idle timeout")
if !openAIStreamClientOutputStarted(c, clientOutputStarted) && !eventShouldFlush {
_ = resp.Body.Close()
return resultWithUsage(), grokStreamIdleFailoverError(account, streamInterval)
}
}
sendErrorEvent("stream_timeout")
return resultWithUsage(), fmt.Errorf("stream data interval timeout")
case <-firstOutputCh:
if firstOutputProgressObserved {
stopFirstOutputTimer()
continue
}
_ = resp.Body.Close()
for ev := range events {
markEventProcessed(ev)
}
return resultWithUsage(), s.newOpenAIFirstOutputTimeoutError(
ctx, c, account, startTime, originalModel, reasoningEffort,
firstOutputTimeout, "semantic_output", resp.Header,
)
case <-keepaliveCh:
if clientDisconnected {
continue
}
if eventInProgress {
continue
}
if time.Since(lastDownstreamWriteAt) < keepaliveInterval {
continue
}
if stageFirstOutput {
// Bypass attempt-local buffered frames. The stable SSE headers may be
// committed here, but account headers remain private until semantic output.
n, err := w.Write([]byte(":\n\n"))
recordOpenAIStreamKeepaliveBytes(c, n)
if err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming, continuing to drain upstream for billing")
continue
}
flusher.Flush()
lastDownstreamWriteAt = time.Now()
continue
}
if _, err := writePendingString(":\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
}
// 部分 OpenAI 兼容上游(例如 Cline API)会将标准响应包在 data 字段中:
// {"data":{"choices": [...], "usage": {...}}, "success":true}。
// 按优先级先保留原有路径,再尝试兼容层 data 包装,
// 避免同步请求能正常返回但用量被静默记录为 0。
candidates := []struct {
usagePath string
imageUsagePath string
}{
{usagePath: "usage", imageUsagePath: "tool_usage.image_gen"},
{usagePath: "response.usage", imageUsagePath: "response.tool_usage.image_gen"},
{usagePath: "data.usage", imageUsagePath: "data.tool_usage.image_gen"},
{usagePath: "data.response.usage", imageUsagePath: "data.response.tool_usage.image_gen"},
}
for _, candidate := range candidates {
if usage, ok := openAIUsageFromGJSON(gjson.GetBytes(body, candidate.usagePath)); ok {
mergeHostedImageGenToolUsage(gjson.GetBytes(body, candidate.imageUsagePath), &usage)
return usage, true
}
}
return OpenAIUsage{}, false
}
// openAIResponsesCompletedEventIsEmpty reports whether a response.completed /
// response.done SSE payload carries no usage, no error and no output items.
// The accumulated usage is consulted too, because OpenAI may deliver usage on
// an earlier event. An empty terminal event after a stream with no semantic
// output is treated as a silent upstream refusal (issue #5009).
func openAIResponsesCompletedEventIsEmpty(data []byte, usage *OpenAIUsage) bool {
if len(data) == 0 || !gjson.ValidBytes(data) {
return false
}
if usage != nil && (usage.InputTokens > 0 || usage.OutputTokens > 0 ||
usage.ImageInputTokens > 0 || usage.ImageOutputTokens > 0 ||
usage.CacheCreationInputTokens > 0 || usage.CacheReadInputTokens > 0) {
return false
}
if gjson.GetBytes(data, "usage").Exists() || gjson.GetBytes(data, "response.usage").Exists() {
return false
}
if gjson.GetBytes(data, "error").Exists() || gjson.GetBytes(data, "response.error").Exists() {
return false
}
if output := gjson.GetBytes(data, "response.output"); output.Exists() && output.IsArray() && len(output.Array()) > 0 {
return false
}
return true
}
func mergeHostedImageGenToolUsage(imageGen gjson.Result, usage *OpenAIUsage) {
if !imageGen.Exists() || !imageGen.IsObject() {
return
}
if usage.ImageOutputTokens == 0 {
if v := imageGen.Get("output_tokens_details.image_tokens").Int(); v > 0 {
usage.ImageOutputTokens = int(v)
}
}
if usage.ImageInputTokens == 0 {
if v := imageGen.Get("input_tokens_details.image_tokens").Int(); v > 0 {
usage.ImageInputTokens = int(v)
}
}
}
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 := openAICacheReadTokensFromUsage(value)
cacheCreationTokens := openAICacheCreationTokensFromUsage(value)
imageOutputTokens := value.Get("output_tokens_details.image_tokens").Int()
if imageOutputTokens == 0 {
imageOutputTokens = value.Get("completion_tokens_details.image_tokens").Int()
}
// 图片输入 token(如 gpt-image-2 的 /v1/images/edits 带图请求),
// 上游在 input_tokens_details.image_tokens 单独回传,用于图/文输入分价计费。
// 普通文本请求该字段为 0,走原路径行为不变。
imageInputTokens := firstPositiveGJSONInt(
value.Get("input_tokens_details.image_tokens"),
value.Get("prompt_tokens_details.image_tokens"),
)
return OpenAIUsage{
InputTokens: int(inputTokens),
ImageInputTokens: imageInputTokens,
OutputTokens: int(outputTokens),
CacheCreationInputTokens: cacheCreationTokens,
CacheReadInputTokens: cacheReadTokens,
ImageOutputTokens: int(imageOutputTokens),
}, true
}
func openAICacheReadTokensFromUsage(value gjson.Result) int {
for _, nested := range []gjson.Result{
value.Get("input_tokens_details.cached_tokens"),
value.Get("prompt_tokens_details.cached_tokens"),
} {
if nested.Exists() {
return max(int(nested.Int()), 0)
}
}
return firstPositiveGJSONInt(
value.Get("cache_read_input_tokens"),
value.Get("cache_read_tokens"),
value.Get("cached_tokens"),
)
}
func openAICacheCreationTokensFromUsage(value gjson.Result) int {
for _, nested := range []gjson.Result{
value.Get("input_tokens_details.cache_write_tokens"),
value.Get("prompt_tokens_details.cache_write_tokens"),
value.Get("input_tokens_details.cache_creation_tokens"),
value.Get("prompt_tokens_details.cache_creation_tokens"),
} {
if nested.Exists() {
return max(int(nested.Int()), 0)
}
}
return firstPositiveGJSONInt(
value.Get("cache_write_tokens"),
value.Get("cache_creation_input_tokens"),
value.Get("cache_write_input_tokens"),
value.Get("cache_creation_tokens"),
)
}
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
}
observer := upstreamResponseModelObserverFromContext(c)
if observer == nil {
observer = beginUpstreamResponseModelObservation(c)
}
if bodyHasSSEFraming(body) {
observeOpenAISSEBody(observer, string(body))
} else {
observer.ObserveOpenAI(body, strings.TrimSpace(gjson.GetBytes(body, "type").String()))
}
// 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, account, 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, account, body, originalModel, mappedModel)
}
if account != nil && account.IsGrok() && isOpenAIResponsesCompactPath(c) {
body, err = convertGrokResponseToOpenAICompact(body)
if err != nil {
return nil, fmt.Errorf("convert Grok compact response: %w", err)
}
}
usageValue, usageOK := extractOpenAIUsageFromJSONBytes(body)
if !usageOK {
if bodyLooksLikeSSE {
return s.handleSSEToJSON(resp, c, account, 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)
}
body, err = restoreGrokResponsesClientToolPayload(c, body)
if err != nil {
return nil, fmt.Errorf("restore Grok Responses client tool response: %w", err)
}
body, err = restoreOpenAIResponsesNamespacePayload(c, body)
if err != nil {
return nil, fmt.Errorf("restore OpenAI namespace response: %w", err)
}
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
// Codex 协议要求 /responses/compact JSON 响应携带 x-codex-turn-state
// codex-api/src/endpoint/compact.rs 从响应头捕获),显式回传。
s.relayOpenAICodexTurnState(c, account, resp.Header)
contentType := "application/json"
if s.cfg != nil && !s.cfg.Security.ResponseHeaders.Enabled {
if upstreamType := resp.Header.Get("Content-Type"); upstreamType != "" {
contentType = upstreamType
}
}
if !writeOpenAICompactSSEBridge(c, resp.StatusCode, body) {
c.Data(resp.StatusCode, contentType, body)
}
return &openaiNonStreamingResult{
OpenAIUsage: usage,
usage: usage,
responseID: extractOpenAIResponseIDFromJSONBytes(body),
imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body),
imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body),
searchCount: countGrokNativeSearchCallsFromJSONBytes(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, account *Account, 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
}
}
}
finalResponse = supplementCompactionItemFromSSE(c, finalResponse, bodyText)
body = finalResponse
if originalModel != mappedModel {
body = s.replaceModelInResponseBody(body, mappedModel, originalModel)
}
// Correct tool calls in final response
body = s.correctToolCallsInResponseBody(body)
restoredBody, restoreErr := restoreGrokResponsesClientToolPayload(c, body)
if restoreErr != nil {
return nil, fmt.Errorf("restore Grok Responses client tool response: %w", restoreErr)
}
restoredBody, restoreErr = restoreOpenAIResponsesNamespacePayload(c, restoredBody)
if restoreErr != nil {
return nil, fmt.Errorf("restore OpenAI namespace response: %w", restoreErr)
}
body = restoredBody
} 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)
s.relayOpenAICodexTurnState(c, account, resp.Header)
contentType := "application/json; charset=utf-8"
if !ok {
contentType = resp.Header.Get("Content-Type")
if contentType == "" {
contentType = "text/event-stream"
}
}
if !writeOpenAICompactSSEBridge(c, resp.StatusCode, body) {
c.Data(resp.StatusCode, contentType, body)
}
return &openaiNonStreamingResult{
OpenAIUsage: usage,
usage: usage,
responseID: extractOpenAIResponseIDFromJSONBytes(body),
imageCount: countOpenAIImageOutputsFromSSEBody(bodyText),
imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText),
searchCount: countGrokNativeSearchCallsFromSSEBody(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) {
eventType = strings.TrimSpace(eventType)
isFailedEvent := eventType == "response.failed"
if (!isFailedEvent && eventType != "error") || len(payload) == 0 || !gjson.ValidBytes(payload) {
return payload, false
}
updated := payload
// 容量降载码对 Codex CLI 是致命错误;事件既然要写给客户端(failover 已不可用),
// 就改写为客户端可重试的错误码。error 帧与 response.failed 都要改:上游降载
// 总是先推 error 帧再收 failed,两帧携带同一个错误。
if rewritten, changed := sanitizeOpenAICapacityShedErrorCodeForClient(updated); changed {
updated = rewritten
}
if !isFailedEvent {
return updated, !bytes.Equal(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, "")
// body-signal compact 心跳可能已把响应头提交为 200,此时只能以
// response.failed 终止事件回传错误,不能再写 JSON+状态码。
if openAICompactClientWantsStream(c) && StopOpenAICompactSSEKeepaliveCommitted(c) {
writeOpenAICompactSSEFailureMessage(c, http.StatusBadGateway, "upstream_error", message)
return fmt.Errorf("non-streaming openai protocol error: %s", 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
}
if normalized, changed := normalizeCompletedImageGenerationStatus(data); changed {
data = normalized
}
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 normalizeCompletedImageGenerationStatus(data []byte) ([]byte, bool) {
if len(data) == 0 || !gjson.ValidBytes(data) {
return data, false
}
shouldNormalize := func(item gjson.Result) bool {
if !item.Exists() || !item.IsObject() ||
strings.TrimSpace(item.Get("type").String()) != "image_generation_call" {
return false
}
switch strings.TrimSpace(item.Get("status").String()) {
case "generating", "in_progress":
return strings.TrimSpace(item.Get("result").String()) != ""
default:
return false
}
}
eventType := strings.TrimSpace(gjson.GetBytes(data, "type").String())
switch eventType {
case "response.output_item.done":
if !shouldNormalize(gjson.GetBytes(data, "item")) {
return data, false
}
updated, err := sjson.SetBytes(data, "item.status", "completed")
if err != nil {
return data, false
}
return updated, true
case "response.completed", "response.done":
output := gjson.GetBytes(data, "response.output")
if !output.Exists() || !output.IsArray() {
return data, false
}
updated := data
changed := false
for i, item := range output.Array() {
if !shouldNormalize(item) {
continue
}
next, err := sjson.SetBytes(updated, "response.output."+strconv.Itoa(i)+".status", "completed")
if err != nil {
return data, false
}
updated = next
changed = true
}
return updated, changed
default:
return data, 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
}
}
// collectRawResponsesOutputItemsFromSSE 按到达顺序收集 SSE 流中
// response.output_item.done 携带的原始 item。除已产生结果但仍停留在进行中
// 的图片状态外,item 以 raw JSON 逐字节保留,
// 避免经窄结构体重建时丢弃 encrypted_content/summary/opaque 等 compact
// 专属或未来新增字段(#3777 问题 2)。若整条流没有任何 done 事件,退回
// 收集 output_item.added 中的 compaction 类 item——compaction 结果没有
// delta 事件,部分上游只在 added 事件中携带完整 item。
func collectRawResponsesOutputItemsFromSSE(bodyText string) ([]byte, bool) {
var items []json.RawMessage
seen := make(map[string]struct{})
hasCompactionItem := false
appendItem := func(item gjson.Result) {
if !item.Exists() || !item.IsObject() {
return
}
key := strings.TrimSpace(item.Get("id").String())
if key == "" {
key = item.Raw
}
if _, dup := seen[key]; dup {
return
}
seen[key] = struct{}{}
if isResponsesCompactionItemType(item.Get("type").String()) {
hasCompactionItem = true
}
items = append(items, json.RawMessage(item.Raw))
}
forEachOpenAISSEDataPayload(bodyText, func(data []byte) {
if normalized, changed := normalizeCompletedImageGenerationStatus(data); changed {
data = normalized
}
if strings.TrimSpace(gjson.GetBytes(data, "type").String()) != "response.output_item.done" {
return
}
appendItem(gjson.GetBytes(data, "item"))
})
// done 事件未携带 compaction item 时再看 added:覆盖"其他 item 有 done、
// compaction 只在 added 中"的混合形态;done 已含 compaction 时跳过,
// 避免同一 item 在无 id 可去重时被收集两份(Codex 要求恰好一个)。
if !hasCompactionItem {
forEachOpenAISSEDataPayload(bodyText, func(data []byte) {
if strings.TrimSpace(gjson.GetBytes(data, "type").String()) != "response.output_item.added" {
return
}
item := gjson.GetBytes(data, "item")
if !isResponsesCompactionItemType(item.Get("type").String()) {
return
}
appendItem(item)
})
}
if len(items) == 0 {
return nil, false
}
outputJSON, err := json.Marshal(items)
if err != nil {
return nil, false
}
return outputJSON, true
}
// isResponsesCompactionItemType reports whether the item type is the Codex
// remote-compact result item ("compaction", upstream alias "compaction_summary").
func isResponsesCompactionItemType(itemType string) bool {
switch strings.TrimSpace(itemType) {
case "compaction", "compaction_summary":
return true
default:
return false
}
}
// supplementCompactionItemFromSSE 保证 compact 请求的终态 output 携带
// compaction item:终态 output 非空但缺失 compaction、而原始事件流的
// output_item.done(或 added)中存在时(上游不一致形态),以 raw JSON 补入。
// Codex remote compact v2 只从 output_item.done 收集 item 且要求恰好一个
// compaction item——纯流式透传(v0.1.146)下客户端直接读事件流天然拿得到,
// SSE→JSON 提取链路必须给出等价结果。非 compact 请求原样返回。
func supplementCompactionItemFromSSE(c *gin.Context, finalResponse []byte, bodyText string) []byte {
if !isOpenAIResponsesCompactPath(c) {
return finalResponse
}
if len(gjson.GetBytes(finalResponse, "output").Array()) == 0 {
// 空 output 由 reconstructResponseOutputFromSSE 整体修补,不在此处理。
return finalResponse
}
if responsesOutputHasCompactionItem(finalResponse) {
return finalResponse
}
item, found := findRawCompactionItemFromSSE(bodyText)
if !found {
return finalResponse
}
patched, err := sjson.SetRawBytes(finalResponse, "output.-1", item)
if err != nil {
return finalResponse
}
return patched
}
// responsesOutputHasCompactionItem reports whether the response JSON already
// carries a compaction item in its output array.
func responsesOutputHasCompactionItem(response []byte) bool {
for _, item := range gjson.GetBytes(response, "output").Array() {
if isResponsesCompactionItemType(item.Get("type").String()) {
return true
}
}
return false
}
// findRawCompactionItemFromSSE 从原始 SSE 事件流中提取第一个 compaction 类
// item 的 raw JSONoutput_item.done 优先,output_item.added 兜底。
func findRawCompactionItemFromSSE(bodyText string) (json.RawMessage, bool) {
var found json.RawMessage
pick := func(eventType string) {
forEachOpenAISSEDataPayload(bodyText, func(data []byte) {
if found != nil {
return
}
if strings.TrimSpace(gjson.GetBytes(data, "type").String()) != eventType {
return
}
item := gjson.GetBytes(data, "item")
if !item.IsObject() || !isResponsesCompactionItemType(item.Get("type").String()) {
return
}
found = json.RawMessage(item.Raw)
})
}
pick("response.output_item.done")
if found == nil {
pick("response.output_item.added")
}
return found, found != nil
}
// reconstructResponseOutputFromSSE scans raw SSE body text and returns a
// JSON-encoded output array for a terminal event whose output is empty.
// Raw output_item.done items are preferred: per the Responses protocol they
// are the authoritative final form of each item. Delta accumulation only
// covers text/function_call/reasoning content and silently drops unknown
// item types such as compaction — Codex remote compact v2 then fails with
// "expected exactly one compaction output item, got 0" (#3887).
// Returns (nil, false) if nothing could be reconstructed.
func reconstructResponseOutputFromSSE(bodyText string) ([]byte, bool) {
if outputJSON, ok := collectRawResponsesOutputItemsFromSSE(bodyText); ok {
return outputJSON, true
}
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")
}