1167 lines
39 KiB
Go
1167 lines
39 KiB
Go
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")
|
||
}
|