Files
sub2api/backend/internal/service/openai_ws_http_bridge.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

706 lines
25 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"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/Wei-Shaw/sub2api/internal/pkg/apicompat"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
)
const (
openAIWSClientReadLimitBytesDefault int64 = 64 * 1024 * 1024
openAIWSHTTPBridgeThresholdBytesDefault int64 = 15 * 1024 * 1024
openAIWSHTTPBridgeErrorBodyLimitBytes = 64 * 1024
)
const openAIWSHTTPBridgeToolStateContextKey = "openai_ws_http_bridge_tool_state"
type openAIWSHTTPBridgeToolState struct {
ClientMapping apicompat.ResponsesClientToolMapping
LoweredTools json.RawMessage
}
func openAIWSHTTPBridgeToolStateFromContext(c *gin.Context) (openAIWSHTTPBridgeToolState, bool) {
if c == nil {
return openAIWSHTTPBridgeToolState{}, false
}
value, ok := c.Get(openAIWSHTTPBridgeToolStateContextKey)
state, typed := value.(openAIWSHTTPBridgeToolState)
return state, ok && typed
}
func setOpenAIWSHTTPBridgeToolState(c *gin.Context, state openAIWSHTTPBridgeToolState) {
if c == nil {
return
}
state.LoweredTools = append(json.RawMessage(nil), state.LoweredTools...)
c.Set(openAIWSHTTPBridgeToolStateContextKey, state)
}
func decodeOpenAIWSHTTPBridgeLoweredTools(raw json.RawMessage) []any {
if len(raw) == 0 {
return nil
}
var tools []any
if err := json.Unmarshal(raw, &tools); err != nil {
return nil
}
return tools
}
func openAIWSHTTPBridgeRawField(body []byte, name string) (json.RawMessage, bool) {
var fields map[string]json.RawMessage
if err := json.Unmarshal(body, &fields); err != nil {
return nil, false
}
raw, present := fields[name]
return append(json.RawMessage(nil), raw...), present
}
func openAIWSHTTPBridgeToolUpstreamName(account *Account) string {
if account != nil && account.Platform == PlatformGrok {
return "Grok WS HTTP bridge"
}
return "OpenAI WS HTTP bridge"
}
// ResolveOpenAIWSClientFirstMessageTimeout returns the effective client ingress deadline.
func ResolveOpenAIWSClientFirstMessageTimeout(cfg *config.Config) time.Duration {
seconds := config.DefaultOpenAIWSClientFirstMessageTimeoutSeconds
if cfg != nil && cfg.Gateway.OpenAIWS.ClientFirstMessageTimeoutSeconds > 0 {
seconds = cfg.Gateway.OpenAIWS.ClientFirstMessageTimeoutSeconds
}
return time.Duration(seconds) * time.Second
}
func ResolveOpenAIWSClientReadLimitBytes(cfg *config.Config) int64 {
if cfg == nil || cfg.Gateway.OpenAIWS.ClientReadLimitBytes <= 0 {
return openAIWSClientReadLimitBytesDefault
}
return cfg.Gateway.OpenAIWS.ClientReadLimitBytes
}
func (s *OpenAIGatewayService) openAIWSHTTPBridgeEnabled() bool {
return s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.HTTPBridgeEnabled
}
func (s *OpenAIGatewayService) openAIWSHTTPBridgeThresholdBytes() int64 {
if s == nil || s.cfg == nil || s.cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes <= 0 {
return openAIWSHTTPBridgeThresholdBytesDefault
}
return s.cfg.Gateway.OpenAIWS.HTTPBridgeThresholdBytes
}
func (s *OpenAIGatewayService) shouldBridgeOpenAIWSHTTP(account *Account, payloadBytes int, previousResponseID string) bool {
if account != nil && account.Platform == PlatformGrok {
return true
}
if !s.openAIWSHTTPBridgeEnabled() {
return false
}
if strings.TrimSpace(previousResponseID) != "" {
return false
}
threshold := s.openAIWSHTTPBridgeThresholdBytes()
return threshold > 0 && int64(payloadBytes) >= threshold
}
func prepareOpenAIWSHTTPBridgeBody(payload []byte) ([]byte, error) {
var body map[string]any
if err := json.Unmarshal(payload, &body); err != nil {
return nil, err
}
if body == nil {
return nil, errors.New("response.create payload must be a JSON object")
}
delete(body, "type")
delete(body, "generate")
delete(body, "previous_response_id")
body["stream"] = true
return json.Marshal(body)
}
type openAIWSToolCallReplayCollector struct {
items []json.RawMessage
seen map[string]struct{}
allItems []json.RawMessage
allSeen map[string]struct{}
}
func (c *openAIWSToolCallReplayCollector) AddEvent(eventType string, message []byte) {
switch strings.TrimSpace(eventType) {
case "response.output_item.done":
item := gjson.GetBytes(message, "item")
c.addAllItem(item)
c.addItem(item)
case "response.completed", "response.done":
output := gjson.GetBytes(message, "response.output")
if !output.IsArray() {
return
}
for _, item := range output.Array() {
c.addAllItem(item)
c.addItem(item)
}
}
}
func (c *openAIWSToolCallReplayCollector) Items() []json.RawMessage {
return cloneOpenAIWSRawMessages(c.items)
}
func (c *openAIWSToolCallReplayCollector) AllItems() []json.RawMessage {
return cloneOpenAIWSRawMessages(c.allItems)
}
func (c *openAIWSToolCallReplayCollector) addAllItem(item gjson.Result) {
if !item.Exists() || item.Type != gjson.JSON {
return
}
raw := strings.TrimSpace(item.Raw)
if raw == "" || !strings.HasPrefix(raw, "{") || strings.TrimSpace(item.Get("type").String()) == "" {
return
}
key := strings.TrimSpace(item.Get("id").String())
if key == "" {
key = strings.TrimSpace(item.Get("call_id").String())
}
if key == "" {
key = raw
}
if c.allSeen == nil {
c.allSeen = make(map[string]struct{})
}
if _, ok := c.allSeen[key]; ok {
return
}
c.allSeen[key] = struct{}{}
c.allItems = append(c.allItems, json.RawMessage(raw))
}
func (c *openAIWSToolCallReplayCollector) addItem(item gjson.Result) {
if !item.Exists() || item.Type != gjson.JSON {
return
}
raw := strings.TrimSpace(item.Raw)
if raw == "" || !strings.HasPrefix(raw, "{") {
return
}
if !isCodexToolCallContextItemType(item.Get("type").String()) {
return
}
key := strings.TrimSpace(item.Get("id").String())
if key == "" {
key = strings.TrimSpace(item.Get("call_id").String())
}
if key == "" {
key = raw
}
if c.seen == nil {
c.seen = make(map[string]struct{})
}
if _, ok := c.seen[key]; ok {
return
}
c.seen[key] = struct{}{}
c.items = append(c.items, json.RawMessage(raw))
}
func buildOpenAIWSHTTPBridgeErrorEvent(statusCode int, message string) []byte {
message = strings.TrimSpace(message)
if message == "" {
message = http.StatusText(statusCode)
}
if message == "" {
message = "upstream request failed"
}
event := map[string]any{
"type": "error",
"status": statusCode,
"error": map[string]any{
"type": "upstream_error",
"message": message,
},
}
body, err := json.Marshal(event)
if err != nil {
return []byte(`{"type":"error","error":{"type":"upstream_error","message":"upstream request failed"}}`)
}
return body
}
func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
ctx context.Context,
c *gin.Context,
account *Account,
token string,
payload []byte,
payloadBytes int,
originalModel string,
imageBillingModel string,
imageSizeTier string,
imageInputSize string,
grokCacheIdentity string,
turn int,
writeClientMessage func([]byte) error,
) (*OpenAIForwardResult, error) {
if s == nil {
return nil, errors.New("service is nil")
}
if s.httpUpstream == nil {
return nil, errors.New("openai http upstream is nil")
}
if account == nil {
return nil, errors.New("account is nil")
}
if writeClientMessage == nil {
return nil, errors.New("client websocket writer is nil")
}
responseModelObserver := &upstreamResponseModelObserver{}
body, err := prepareOpenAIWSHTTPBridgeBody(payload)
if err != nil {
return nil, fmt.Errorf("prepare http bridge body: %w", err)
}
grokIntentSourceBody := append([]byte(nil), body...)
_, grokExplicitToolsField := openAIWSHTTPBridgeRawField(grokIntentSourceBody, "tools")
grokExplicitToolIntent := account.Platform == PlatformGrok && hasGrokResponsesToolIntent(grokIntentSourceBody)
var clientToolMapping apicompat.ResponsesClientToolMapping
functionToolUpstream := (account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey) || account.Platform == PlatformGrok
if functionToolUpstream {
if account.Platform == PlatformGrok {
body, err = sanitizeGrokResponsesInput(body)
if err != nil {
return nil, fmt.Errorf("sanitize Grok WS HTTP bridge input: %w", err)
}
}
inheritedState, _ := openAIWSHTTPBridgeToolStateFromContext(c)
inheritedLoweredTools := decodeOpenAIWSHTTPBridgeLoweredTools(inheritedState.LoweredTools)
body, clientToolMapping, err = adaptResponsesClientToolsForFunctionUpstreamWithMapping(
body,
openAIWSHTTPBridgeToolUpstreamName(account),
inheritedState.ClientMapping,
inheritedLoweredTools,
)
if err != nil {
return nil, fmt.Errorf("adapt %s client tools: %w", openAIWSHTTPBridgeToolUpstreamName(account), err)
}
if account.Platform == PlatformGrok && !grokExplicitToolsField && !grokExplicitToolIntent && len(inheritedLoweredTools) > 0 && hasGrokResponsesToolIntent(body) {
// This continuation omitted tools, so the pre-adapter source cannot
// represent the effective inherited declarations. Cache routing must
// see the rehydrated tool intent or it will replace client functions
// with the native-search tool-free route. Explicit current-turn tool
// intent still uses the original pre-sanitization source above.
grokIntentSourceBody = append(grokIntentSourceBody[:0], body...)
}
loweredTools := inheritedState.LoweredTools
if currentTools, present := openAIWSHTTPBridgeRawField(body, "tools"); present {
loweredTools = currentTools
}
setOpenAIWSHTTPBridgeToolState(c, openAIWSHTTPBridgeToolState{
ClientMapping: clientToolMapping,
LoweredTools: loweredTools,
})
}
upstreamCtx, releaseUpstreamCtx := detachUpstreamContext(ctx)
var upstreamReq *http.Request
if account.Platform == PlatformGrok {
upstreamModel := resolveGrokWSUpstreamModel(account, body, originalModel)
body, err = patchGrokResponsesBody(body, upstreamModel)
if err != nil {
releaseUpstreamCtx()
return nil, err
}
grokMixedCacheIntentBody := append([]byte(nil), body...)
body, err = applyGrokResponsesCacheIdentity(body, grokIntentSourceBody, grokCacheIdentity, account.IsGrokOAuth())
if err != nil {
releaseUpstreamCtx()
return nil, fmt.Errorf("apply grok prompt cache identity: %w", err)
}
body, err = applyGrokFreeRequestToolCacheRoute(c, body, grokMixedCacheIntentBody, account, grokCacheIdentity)
if err != nil {
releaseUpstreamCtx()
return nil, fmt.Errorf("apply grok Free function-tool cache route: %w", err)
}
upstreamReq, err = buildGrokResponsesRequest(upstreamCtx, c, account, body, token, grokCacheIdentity, s.cfg, s.settingService)
} else {
upstreamReq, err = s.buildUpstreamRequestOpenAIPassthrough(upstreamCtx, c, account, body, token)
}
releaseUpstreamCtx()
if err != nil {
return nil, err
}
if account.Platform != PlatformGrok && isOpenAIResponsesLiteWebSocketPayload(payload) {
upstreamReq.Header.Set(responsesLiteHeader, "true")
}
proxyURL := ""
if account.ProxyID != nil && account.Proxy != nil {
proxyURL = account.Proxy.URL()
}
if c != nil {
c.Set("openai_passthrough", true)
c.Set("openai_ws_http_bridge", true)
}
turnStart := time.Now()
resp, err := s.httpUpstream.Do(upstreamReq, proxyURL, account.ID, account.Concurrency)
if err != nil {
if turn == 1 {
return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, err, true)
}
safeErr := sanitizeUpstreamErrorMessage(err.Error())
_ = writeClientMessage(buildOpenAIWSHTTPBridgeErrorEvent(http.StatusBadGateway, "Upstream request failed"))
return nil, fmt.Errorf("upstream http bridge request failed: %s", safeErr)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode >= 400 {
respBody, _ := io.ReadAll(io.LimitReader(resp.Body, openAIWSHTTPBridgeErrorBodyLimitBytes))
upstreamMsg := sanitizeUpstreamErrorMessage(strings.TrimSpace(extractUpstreamErrorMessage(respBody)))
if upstreamMsg == "" {
upstreamMsg = http.StatusText(resp.StatusCode)
}
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(resp.StatusCode, upstreamMsg, respBody)
if account.Platform == PlatformGrok {
shouldFailover = s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody)
s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, resolveGrokWSUpstreamModel(account, body, originalModel)), account, resp.StatusCode, resp.Header, respBody)
if shouldFailover && (turn == 1 || resp.StatusCode == http.StatusTooManyRequests) {
return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, false)
}
} else if shouldFailover && (turn == 1 || resp.StatusCode == http.StatusTooManyRequests) {
return nil, s.handleFailoverErrorResponsePassthrough(ctx, resp, c, account, body, respBody)
}
if account.Platform != PlatformGrok && (shouldFailover || shouldCooldownOpenAITransientUpstreamError(resp.StatusCode, respBody)) {
canonicalModel := canonicalOpenAIAccountSchedulingModel(account, originalModel)
s.handleOpenAIAccountUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, canonicalModel)
}
_ = writeClientMessage(buildOpenAIWSHTTPBridgeErrorEvent(resp.StatusCode, upstreamMsg))
return nil, fmt.Errorf("upstream http bridge error: status=%d message=%s", resp.StatusCode, upstreamMsg)
}
if account.Platform == PlatformGrok {
s.updateGrokUsageFromResponse(withGrokTeamRateLimitModel(ctx, resolveGrokWSUpstreamModel(account, body, originalModel)), account, resp.Header, resp.StatusCode)
}
responseID := ""
usage := OpenAIUsage{}
imageCounter := newOpenAIImageOutputCounter()
var firstTokenMs *int
reqStream := openAIWSPayloadBoolFromRaw(body, "stream", true)
eventCount := 0
tokenEventCount := 0
terminalEventCount := 0
replayCollector := &openAIWSToolCallReplayCollector{}
firstEventType := ""
lastEventType := ""
upstreamTerminalEvent := ""
sawDone := false
wroteDownstream := false
pendingClientMessages := make([][]byte, 0, 4)
pendingClientMessageBytes := int64(0)
capacityFailoverSuppressedLogged := false
clientDisconnected := false
mappedModel := ""
needModelReplace := false
var mappedModelBytes []byte
if originalModel != "" {
mappedModel = strings.TrimSpace(gjson.GetBytes(body, "model").String())
if mappedModel == "" {
mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
}
needModelReplace = mappedModel != "" && mappedModel != originalModel
if needModelReplace {
mappedModelBytes = []byte(mappedModel)
}
}
resultWithUsage := func() *OpenAIForwardResult {
imageCount := imageCounter.Count()
result := &OpenAIForwardResult{
RequestID: responseID,
Usage: usage,
Model: originalModel,
UpstreamModel: mappedModel,
UpstreamResponseModel: responseModelObserver.Model(),
UpstreamResponseModelConflict: responseModelObserver.Conflict(),
ServiceTier: extractOpenAIServiceTierFromBody(body),
ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(body, mappedModel, originalModel), body, mappedModel),
Stream: reqStream,
OpenAIWSMode: true,
UpstreamTerminalEvent: upstreamTerminalEvent,
ResponseHeaders: cloneHeader(resp.Header),
Duration: time.Since(turnStart),
FirstTokenMs: firstTokenMs,
}
if replayInput := replayCollector.Items(); len(replayInput) > 0 {
result.wsReplayInput = replayInput
result.wsReplayInputExists = true
}
result.wsAccountFailoverReplayInput = replayCollector.AllItems()
if imageCount > 0 {
result.ImageCount = imageCount
result.ImageSize = imageSizeTier
result.ImageInputSize = imageInputSize
result.ImageOutputSizes = imageCounter.Sizes()
result.BillingModel = imageBillingModel
}
return result
}
maxLineSize := defaultMaxLineSize
if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 {
maxLineSize = s.cfg.Gateway.MaxLineSize
}
if hasResponsesClientToolMapping(clientToolMapping) {
resp.Body = newResponsesClientToolStreamBody(resp.Body, clientToolMapping, maxLineSize)
}
scanner := bufio.NewScanner(resp.Body)
scanBuf := getSSEScannerBuf64K()
scanner.Buffer(scanBuf[:0], maxLineSize)
defer putSSEScannerBuf64K(scanBuf)
for scanner.Scan() {
line := scanner.Text()
data, ok := extractOpenAISSEDataLine(line)
if !ok {
continue
}
trimmedData := strings.TrimSpace(data)
if trimmedData == "" {
continue
}
if trimmedData == "[DONE]" {
sawDone = true
continue
}
upstreamMessage := []byte(trimmedData)
if normalized, changed := normalizeCompletedImageGenerationStatus(upstreamMessage); changed {
upstreamMessage = normalized
}
eventType, eventResponseID, _ := parseOpenAIWSEventEnvelope(upstreamMessage)
responseModelObserver.ObserveOpenAI(upstreamMessage, eventType)
if responseID == "" && eventResponseID != "" {
responseID = eventResponseID
}
if eventType != "" {
eventCount++
if firstEventType == "" {
firstEventType = eventType
}
lastEventType = eventType
}
if isOpenAIWSTokenEvent(eventType) {
tokenEventCount++
if firstTokenMs == nil {
ms := int(time.Since(turnStart).Milliseconds())
firstTokenMs = &ms
}
}
if openAIWSEventShouldParseUsage(eventType) {
parseOpenAIWSResponseUsageFromCompletedEvent(upstreamMessage, &usage)
}
imageCounter.AddSSEData(upstreamMessage)
if needModelReplace && len(mappedModelBytes) > 0 && openAIWSEventMayContainModel(eventType) && strings.Contains(trimmedData, mappedModel) {
upstreamMessage = replaceOpenAIWSMessageModel(upstreamMessage, mappedModel, originalModel)
}
if s.toolCorrector != nil && openAIWSEventMayContainToolCalls(eventType) && openAIWSMessageLikelyContainsToolCalls(upstreamMessage) {
if corrected, changed := s.toolCorrector.CorrectToolCallsInSSEBytes(upstreamMessage); changed {
upstreamMessage = corrected
}
}
replayCollector.AddEvent(eventType, upstreamMessage)
var upstreamEventErr error
if eventType == "error" || eventType == "response.failed" {
errMessage := extractOpenAISSEErrorMessage(upstreamMessage)
if errMessage == "" {
errMessage = "upstream error event"
}
statusCode := openAIStreamFailureStatus(upstreamMessage, errMessage)
shouldFailover := openAIStreamFailedEventShouldFailover(upstreamMessage, errMessage)
if eventType == "error" {
errCodeRaw, errTypeRaw, _ := parseOpenAIWSErrorEventFields(upstreamMessage)
statusCode = openAIWSErrorHTTPStatusFromRaw(errCodeRaw, errTypeRaw)
shouldFailover = s.shouldFailoverOpenAIUpstreamResponse(statusCode, errMessage, upstreamMessage)
}
requestScopedCapacity := isOpenAIUpstreamCapacityShedEvent(upstreamMessage)
if account.Platform == PlatformGrok && eventType == "error" {
// SSE error events do not carry an HTTP status. The local status
// mapper therefore defaults unknown xAI codes (for example
// new_sensitive) to 502; classify the body as a request-scoped
// 403 before applying status-based failover or account state.
if isGrokContentPolicyRejection(http.StatusForbidden, upstreamMessage) {
shouldFailover = false
} else {
shouldFailover = s.shouldFailoverGrokUpstreamError(statusCode, upstreamMessage)
s.handleGrokAccountUpstreamError(ctx, account, statusCode, resp.Header, upstreamMessage)
}
} else if eventType == "error" && shouldFailover && !requestScopedCapacity {
accountStatus := statusCode
if transientStatus := openAIWSPayloadTransientStatus(upstreamMessage); transientStatus != 0 {
accountStatus = transientStatus
}
canonicalModel := canonicalOpenAIAccountSchedulingModel(account, originalModel)
s.handleOpenAIAccountUpstreamError(ctx, account, accountStatus, resp.Header, upstreamMessage, canonicalModel)
}
if !wroteDownstream && shouldFailover && (turn == 1 || statusCode == http.StatusTooManyRequests) {
if account.Platform == PlatformGrok {
return nil, newOpenAIUpstreamFailoverError(statusCode, resp.Header, upstreamMessage, errMessage, false)
}
return nil, s.newOpenAIStreamFailoverError(c, account, true, resp.Header.Get("x-request-id"), upstreamMessage, errMessage, resp.Header)
}
if wroteDownstream && requestScopedCapacity && !capacityFailoverSuppressedLogged {
logOpenAICapacityFailoverSuppressed(ctx, account, "ws_http_bridge", resp.Header.Get("x-request-id"), eventType)
capacityFailoverSuppressedLogged = true
}
if eventType == "error" {
upstreamEventErr = errors.New(errMessage)
}
}
// 客户端写出副本改写容量降载码:Codex 对 error/response.failed 中的
// server_is_overloaded / slow_down 判致命并终止会话,改写后走客户端内置
// 重试。账号状态与终止事件判定(下方 handleOpenAIWSTerminalTransientFailure
// 仍使用未改写的 upstreamMessage。
clientMessage := upstreamMessage
if eventType == "error" || eventType == "response.failed" {
if rewritten, changed := sanitizeOpenAICapacityShedErrorCodeForClient(clientMessage); changed {
clientMessage = rewritten
}
}
if !clientDisconnected {
stageBeforeSemanticOutput := turn == 1 && account.Platform == PlatformOpenAI && !wroteDownstream
commitStagedMessages := !stageBeforeSemanticOutput ||
openAIStreamDataStartsClientOutput(string(clientMessage), eventType) ||
isOpenAIWSTerminalEvent(eventType)
if stageBeforeSemanticOutput && !commitStagedMessages {
if pendingClientMessageBytes+int64(len(clientMessage)) > openAIFirstOutputStageMaxBytes {
return nil, s.newOpenAIStreamFailoverError(
c,
account,
true,
resp.Header.Get("x-request-id"),
nil,
"OpenAI WS HTTP bridge first-output staging limit exceeded",
resp.Header,
)
}
pendingClientMessages = append(pendingClientMessages, append([]byte(nil), clientMessage...))
pendingClientMessageBytes += int64(len(clientMessage))
} else {
messages := append(pendingClientMessages, clientMessage)
pendingClientMessages = nil
pendingClientMessageBytes = 0
for _, message := range messages {
if err := writeClientMessage(message); err != nil {
if isOpenAIWSClientDisconnectError(err) {
clientDisconnected = true
closeStatus, closeReason := summarizeOpenAIWSReadCloseError(err)
logOpenAIWSModeInfo(
"ingress_ws_http_bridge_client_disconnected_drain account_id=%d turn=%d close_status=%s close_reason=%s",
account.ID,
turn,
closeStatus,
truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
)
break
}
return nil, wrapOpenAIWSIngressTurnError(
"write_client",
fmt.Errorf("write client websocket event: %w", err),
wroteDownstream,
)
}
wroteDownstream = true
}
}
}
if upstreamEventErr != nil {
return resultWithUsage(), upstreamEventErr
}
if isOpenAIWSTerminalEvent(eventType) {
upstreamTerminalEvent = s.handleOpenAIWSTerminalTransientFailure(ctx, account, canonicalOpenAIAccountSchedulingModel(account, originalModel), resp.Header, upstreamMessage)
terminalEventCount++
firstTokenMsValue := -1
if firstTokenMs != nil {
firstTokenMsValue = *firstTokenMs
}
logOpenAIWSModeInfo(
"ingress_ws_http_bridge_turn_completed account_id=%d turn=%d response_id=%s payload_bytes=%d duration_ms=%d events=%d token_events=%d terminal_events=%d first_event=%s last_event=%s first_token_ms=%d client_disconnected=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen),
payloadBytes,
time.Since(turnStart).Milliseconds(),
eventCount,
tokenEventCount,
terminalEventCount,
truncateOpenAIWSLogValue(firstEventType, openAIWSLogValueMaxLen),
truncateOpenAIWSLogValue(lastEventType, openAIWSLogValueMaxLen),
firstTokenMsValue,
clientDisconnected,
)
return resultWithUsage(), nil
}
}
if err := scanner.Err(); err != nil {
streamErr := fmt.Errorf("read upstream http bridge stream: %w", err)
if turn == 1 && !wroteDownstream {
return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, streamErr, true)
}
return resultWithUsage(), streamErr
}
terminalErr := errors.New("upstream http bridge stream ended before terminal event")
if sawDone {
terminalErr = errors.New("upstream http bridge stream sent [DONE] before terminal event")
}
if turn == 1 && !wroteDownstream {
return nil, s.handleOpenAIUpstreamTransportError(ctx, c, account, terminalErr, true)
}
return resultWithUsage(), terminalErr
}
func resolveGrokWSCacheIdentity(c *gin.Context, account *Account, seedPayload, currentPayload []byte, originalModel string) (string, error) {
body, err := prepareOpenAIWSHTTPBridgeBody(seedPayload)
if err != nil {
return "", err
}
upstreamModel := resolveGrokWSUpstreamModel(account, currentPayload, originalModel)
body, err = patchGrokResponsesBody(body, upstreamModel)
if err != nil {
return "", err
}
return resolveGrokCacheIdentity(c, body, "", upstreamModel), nil
}
func resolveGrokWSUpstreamModel(account *Account, body []byte, originalModel string) string {
upstreamModel := strings.TrimSpace(gjson.GetBytes(body, "model").String())
originalModel = strings.TrimSpace(originalModel)
// Shared ingress has already applied channel and account mappings when the
// body model differs from the client-facing model. Only resolve from the
// original model when the body still carries that original value.
if account != nil && originalModel != "" && (upstreamModel == "" || upstreamModel == originalModel) {
if mappedModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel)); mappedModel != "" {
upstreamModel = mappedModel
}
}
if upstreamModel == "" {
upstreamModel = grokDefaultResponsesModel
}
return upstreamModel
}