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

1757 lines
69 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 (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/url"
"strings"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
coderws "github.com/coder/websocket"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
"github.com/tidwall/sjson"
)
func (s *OpenAIGatewayService) openAIWSIngressInterTurnIdleTimeout() time.Duration {
if s == nil || s.cfg == nil || s.cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds <= 0 {
return 0
}
return time.Duration(s.cfg.Gateway.OpenAIWS.IngressInterTurnIdleTimeoutSeconds) * time.Second
}
// newOpenAIWSDownstreamWriteContext binds writes directly to the client
// lifecycle while excluding the separate ingress-lease cancellation signal.
// This lets a lease-loss path finish its current client write before
// ReadOpenAIWSClientMessage sends the retryable close frame.
func newOpenAIWSDownstreamWriteContext(controlCtx context.Context, hooks *OpenAIWSIngressHooks, timeout time.Duration) (context.Context, context.CancelFunc) {
writeParent := controlCtx
if hooks != nil && hooks.ClientLifecycleContext != nil {
writeParent = hooks.ClientLifecycleContext
}
if writeParent == nil {
writeParent = context.Background()
}
return context.WithTimeout(writeParent, timeout)
}
func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
ctx context.Context,
c *gin.Context,
clientConn *coderws.Conn,
account *Account,
token string,
firstClientMessage []byte,
hooks *OpenAIWSIngressHooks,
) error {
if s == nil {
return errors.New("service is nil")
}
if c == nil {
return errors.New("gin context is nil")
}
if clientConn == nil {
return errors.New("client websocket is nil")
}
if account == nil {
return errors.New("account is nil")
}
if err := validateOpenAIWSBearerToken(account, token); err != nil {
return err
}
// 预取一次 OpenAI Fast Policy settings,绑定到 ctx,让该 WS session
// 内所有帧的 evaluateOpenAIFastPolicy 调用复用同一份快照,避免每帧
// 进入 DB / settingRepo。Trade-off 见 withOpenAIFastPolicyContext 注释。
if s.settingService != nil {
if settings, err := s.settingService.GetOpenAIFastPolicySettings(ctx); err == nil && settings != nil {
ctx = withOpenAIFastPolicyContext(ctx, settings)
}
}
wsDecision := s.getOpenAIWSProtocolResolver().Resolve(account)
forceHTTPBridge := account.Platform == PlatformGrok
modeRouterV2Enabled := s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.ModeRouterV2Enabled
ingressMode := OpenAIWSIngressModeCtxPool
if modeRouterV2Enabled && !forceHTTPBridge {
ingressMode = account.ResolveOpenAIResponsesWebSocketV2Mode(s.cfg.Gateway.OpenAIWS.IngressModeDefault)
if ingressMode == OpenAIWSIngressModeOff {
return NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
"websocket mode is disabled for this account",
nil,
)
}
switch ingressMode {
case OpenAIWSIngressModePassthrough:
if wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
return fmt.Errorf("websocket ingress requires ws_v2 transport, got=%s", wsDecision.Transport)
}
// 透传 relay 通过 TurnStarted 记录每个 turn 的开始时刻,但不触发
// BeforeTurn;因此仍只有建连时的利润准入门,没有 turn 级复核。
// handler 计费在 turn 定价未冻结时回退到对应的 turn 开始时刻。
return s.proxyResponsesWebSocketV2Passthrough(
ctx,
c,
clientConn,
account,
token,
firstClientMessage,
hooks,
wsDecision,
)
case OpenAIWSIngressModeHTTPBridge:
forceHTTPBridge = true
case OpenAIWSIngressModeCtxPool, OpenAIWSIngressModeShared, OpenAIWSIngressModeDedicated:
// continue
default:
return NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
"websocket mode only supports ctx_pool/passthrough/http_bridge",
nil,
)
}
}
if !forceHTTPBridge && wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 {
return fmt.Errorf("websocket ingress requires ws_v2 transport, got=%s", wsDecision.Transport)
}
dedicatedMode := modeRouterV2Enabled && ingressMode == OpenAIWSIngressModeDedicated
wsURL := ""
wsHost := "-"
wsPath := "-"
if forceHTTPBridge {
wsHost = "xai-http-bridge"
wsPath = "/v1/responses"
} else {
var err error
wsURL, err = s.buildOpenAIResponsesWSURL(account)
if err != nil {
return fmt.Errorf("build ws url: %w", err)
}
if parsedURL, parseErr := url.Parse(wsURL); parseErr == nil && parsedURL != nil {
wsHost = normalizeOpenAIWSLogValue(parsedURL.Host)
wsPath = normalizeOpenAIWSLogValue(parsedURL.Path)
}
}
debugEnabled := isOpenAIWSModeDebugEnabled()
isCodexCLI := openai.IsCodexOfficialClientByHeaders(c.GetHeader("User-Agent"), c.GetHeader("originator")) || (s.cfg != nil && s.cfg.Gateway.ForceCodexCLI)
type openAIWSClientPayload struct {
payloadRaw []byte
rawForHash []byte
promptCacheKey string
previousResponseID string
originalModel string
imageBillingModel string
imageSizeTier string
imageInputSize string
payloadBytes int
}
ingressSessionOriginalModel := ""
applyPayloadMutation := func(current []byte, path string, value any) ([]byte, error) {
next, err := sjson.SetBytes(current, path, value)
if err == nil {
return next, nil
}
// 仅在确实需要修改 payload 且 sjson 失败时,退回 map 路径确保兼容性。
payload := make(map[string]any)
if unmarshalErr := json.Unmarshal(current, &payload); unmarshalErr != nil {
return nil, err
}
switch path {
case "type", "model":
payload[path] = value
case "client_metadata." + openAIWSTurnMetadataHeader:
setOpenAIWSTurnMetadata(payload, fmt.Sprintf("%v", value))
default:
return nil, err
}
rebuilt, marshalErr := json.Marshal(payload)
if marshalErr != nil {
return nil, marshalErr
}
return rebuilt, nil
}
parseClientPayload := func(turn int, raw []byte) (openAIWSClientPayload, error) {
trimmed := bytes.TrimSpace(raw)
if len(trimmed) == 0 {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "empty websocket request payload", nil)
}
if !gjson.ValidBytes(trimmed) {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", errors.New("invalid json"))
}
values := gjson.GetManyBytes(trimmed, "type", "model", "prompt_cache_key", "previous_response_id")
eventType := strings.TrimSpace(values[0].String())
normalized := trimmed
switch eventType {
case "":
eventType = "response.create"
next, setErr := applyPayloadMutation(normalized, "type", eventType)
if setErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", setErr)
}
normalized = next
case "response.create":
case "response.append":
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
"response.append is not supported in ws v2; use response.create with previous_response_id",
nil,
)
default:
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
fmt.Sprintf("unsupported websocket request type: %s", eventType),
nil,
)
}
if hooks != nil && (hooks.MaxReasoningEffort != "" || len(hooks.ReasoningEffortMappings) > 0) {
if capped, changed := ApplyOpenAIReasoningEffortPolicy(normalized, hooks.MaxReasoningEffort, hooks.ReasoningEffortMappings); changed {
normalized = capped
}
}
originalModel := strings.TrimSpace(values[1].String())
modelMissing := originalModel == ""
if originalModel == "" {
// 入站 WS 长会话里,部分客户端只在第一轮 response.create 上声明
// model,后续 turn 复用同一 session-level model。为避免因省略
// model 直接断开用户连接,这里回落到上一轮已通过校验的客户端模型,
// 并在下方写回上游 payload,保证账号模型映射/fast policy/图片权限
// 仍按同一模型执行。
originalModel = ingressSessionOriginalModel
if originalModel == "" {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
"model is required in response.create payload",
nil,
)
}
}
promptCacheKey := strings.TrimSpace(values[2].String())
previousResponseID := strings.TrimSpace(values[3].String())
previousResponseIDKind := ClassifyOpenAIPreviousResponseIDKind(previousResponseID)
if previousResponseID != "" && previousResponseIDKind == OpenAIPreviousResponseIDKindMessageID {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
"previous_response_id must be a response.id (resp_*), not a message id",
nil,
)
}
if turnMetadata := strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)); turnMetadata != "" {
next, setErr := applyPayloadMutation(normalized, "client_metadata."+openAIWSTurnMetadataHeader, turnMetadata)
if setErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", setErr)
}
normalized = next
}
if account.IsOpenAIOAuth() && isOpenAIResponsesLiteWebSocketPayload(normalized) {
litePayload, _, liteErr := normalizeOpenAIResponsesLiteToolsPayload(normalized)
if liteErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
liteErr.Error(),
liteErr,
)
}
normalized = litePayload
}
apiKey := getAPIKeyFromContext(c)
imageGenerationAllowed := GroupAllowsImageGeneration(apiKeyGroup(apiKey))
codexImageGenerationExplicitToolPolicy := codexImageGenerationExplicitToolPolicyAllow
if isCodexCLI {
codexImageGenerationExplicitToolPolicy = account.CodexImageGenerationExplicitToolPolicy()
}
codexBridgeEnabled := isCodexCLI &&
!isOpenAIResponsesLiteWebSocketPayload(normalized) &&
imageGenerationAllowed &&
codexImageGenerationExplicitToolPolicy != codexImageGenerationExplicitToolPolicyStrip &&
s.isCodexImageGenerationBridgeEnabled(ctx, account, apiKey)
if codexBridgeEnabled {
payloadMap := make(map[string]any)
if err := json.Unmarshal(normalized, &payloadMap); err != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", err)
}
bridgeModified := false
if ensureOpenAIResponsesImageGenerationTool(payloadMap) {
bridgeModified = true
logOpenAIWSModeInfo("ingress_ws_codex_image_tool_injected account_id=%d", account.ID)
}
if ensureOpenAIResponsesImageGenerationToolChoiceAuto(payloadMap) {
bridgeModified = true
logOpenAIWSModeInfo("ingress_ws_codex_image_tool_choice_auto account_id=%d", account.ID)
}
if normalizeOpenAIResponsesImageGenerationTools(payloadMap) {
bridgeModified = true
}
if applyCodexImageGenerationBridgeInstructions(payloadMap) {
bridgeModified = true
logOpenAIWSModeInfo("ingress_ws_codex_image_bridge_instructions_added account_id=%d", account.ID)
}
if bridgeModified {
rebuilt, marshalErr := json.Marshal(payloadMap)
if marshalErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", marshalErr)
}
normalized = rebuilt
}
}
requestModel := originalModel
if hooks != nil && hooks.MapRequestModel != nil {
mappedModel, mapErr := hooks.MapRequestModel(turn, originalModel)
if mapErr != nil {
return openAIWSClientPayload{}, mapErr
}
if mappedModel = strings.TrimSpace(mappedModel); mappedModel != "" {
requestModel = mappedModel
}
}
upstreamModel := normalizeOpenAIModelForUpstream(account, account.GetMappedModel(requestModel))
if modelMissing || upstreamModel != originalModel {
next, setErr := applyPayloadMutation(normalized, "model", upstreamModel)
if setErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", setErr)
}
normalized = next
}
if isCodexCLI && codexImageGenerationExplicitToolPolicy == codexImageGenerationExplicitToolPolicyStrip {
if stripped, changed, stripErr := stripOpenAIImageGenerationToolsFromRawPayload(normalized); stripErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", stripErr)
} else if changed {
normalized = stripped
logOpenAIWSModeInfo("ingress_ws_codex_image_tool_stripped_by_policy account_id=%d", account.ID)
}
}
if stripped, changed, stripErr := stripCodexSparkImageGenerationToolFromRawPayload(normalized, upstreamModel); stripErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", stripErr)
} else if changed {
normalized = stripped
logOpenAIWSModeInfo("ingress_ws_codex_spark_image_tool_stripped account_id=%d", account.ID)
}
imageIntent := IsImageGenerationIntentForPlatform(openAIResponsesEndpoint, originalModel, normalized, account.Platform)
if imageIntent && !imageGenerationAllowed {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, ImageGenerationPermissionMessage(), nil)
}
imageBillingModel := ""
imageSizeTier := ""
imageInputSize := ""
if imageIntent {
var imageCfgErr error
imageCfg, imageCfgErr := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(normalized, originalModel)
if imageCfgErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, imageCfgErr.Error(), imageCfgErr)
}
imageBillingModel = imageCfg.Model
imageSizeTier = imageCfg.SizeTier
imageInputSize = imageCfg.InputSize
}
// Apply OpenAI Fast Policy on the response.create frame using the same
// evaluator/normalize/scope rules as the HTTP entrypoints. This is the
// single integration point for all WS ingress turns (first + follow-up
// frames flow through here).
//
// Model fallback: first turn still requires model at the handler layer
// follow-up response.create frames may omit it and then reuse
// ingressSessionOriginalModel. We always write a concrete upstream model
// before evaluating policy, so whitelist / filter behavior remains stable.
policyApplied, blocked, policyErr := s.applyOpenAIFastPolicyToWSResponseCreate(ctx, account, upstreamModel, normalized)
if policyErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, "invalid websocket request payload", policyErr)
}
if blocked != nil {
MarkOpsClientBusinessLimited(c, OpsClientBusinessLimitedReasonLocalPolicyDenied)
// Send a Realtime-style error event to the client first, then
// signal the handler to close the connection with PolicyViolation.
// We intentionally do NOT forward this frame upstream.
//
// coder/websocket@v1.8.14 Conn.Write is synchronous and flushes
// the underlying bufio writer before returning (write.go:42 →
// 307-311), and the subsequent close handshake re-acquires the
// same writeFrameMu, so the error event is guaranteed to reach
// the kernel send buffer before any close frame is queued.
eventBytes := buildOpenAIFastPolicyBlockedWSEvent(blocked)
if eventBytes != nil {
writeCtx, cancel := newOpenAIWSDownstreamWriteContext(ctx, hooks, s.openAIWSWriteTimeout())
_ = clientConn.Write(writeCtx, coderws.MessageText, eventBytes)
cancel()
}
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
blocked.Message,
blocked,
)
}
normalized = policyApplied
ingressSessionOriginalModel = originalModel
return openAIWSClientPayload{
payloadRaw: normalized,
rawForHash: trimmed,
promptCacheKey: promptCacheKey,
previousResponseID: previousResponseID,
originalModel: originalModel,
imageBillingModel: imageBillingModel,
imageSizeTier: imageSizeTier,
imageInputSize: imageInputSize,
payloadBytes: len(normalized),
}, nil
}
writeClientMessage := func(message []byte) error {
writeCtx, cancel := newOpenAIWSDownstreamWriteContext(ctx, hooks, s.openAIWSWriteTimeout())
defer cancel()
return clientConn.Write(writeCtx, coderws.MessageText, message)
}
readClientMessage := func() ([]byte, error) {
idleTimeout := s.openAIWSIngressInterTurnIdleTimeout()
msgType, payload, readErr := ReadOpenAIWSClientMessage(
ctx,
clientConn,
idleTimeout,
coderws.StatusNormalClosure,
"websocket idle timeout",
)
if readErr != nil {
var closeErr *OpenAIWSClientCloseError
if errors.As(readErr, &closeErr) && closeErr.StatusCode() == coderws.StatusNormalClosure {
logOpenAIWSModeInfo("ingress_ws_inter_turn_idle_timeout account_id=%d timeout_seconds=%d", account.ID, int(idleTimeout.Seconds()))
}
return nil, readErr
}
if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
return nil, NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
fmt.Sprintf("unsupported websocket client message type: %s", msgType.String()),
nil,
)
}
return payload, nil
}
firstPayload, err := parseClientPayload(1, firstClientMessage)
if err != nil {
return err
}
turnState := strings.TrimSpace(c.GetHeader(openAIWSTurnStateHeader))
stateStore := s.getOpenAIWSStateStore()
groupID := getOpenAIGroupIDFromContext(c)
storeDisabledConnMode := s.openAIWSStoreDisabledConnMode()
sessionHash := ""
preferredConnID := ""
storeDisabled := false
refreshIngressRouteState := func(payload openAIWSClientPayload) {
sessionHash = s.GenerateSessionHash(c, payload.rawForHash)
if turnState == "" && stateStore != nil && sessionHash != "" {
if savedTurnState, ok := stateStore.GetSessionTurnState(groupID, sessionHash); ok {
turnState = savedTurnState
}
}
preferredConnID = ""
if stateStore != nil && payload.previousResponseID != "" {
if connID, ok := stateStore.GetResponseConn(payload.previousResponseID); ok {
preferredConnID = connID
}
}
storeDisabled = s.isOpenAIWSStoreDisabledInRequestRaw(payload.payloadRaw, account)
if stateStore != nil && storeDisabled && payload.previousResponseID == "" && sessionHash != "" {
if connID, ok := stateStore.GetSessionConn(groupID, sessionHash); ok {
preferredConnID = connID
}
}
}
refreshIngressRouteState(firstPayload)
if forceHTTPBridge || s.shouldBridgeOpenAIWSHTTP(account, firstPayload.payloadBytes, firstPayload.previousResponseID) {
logOpenAIWSModeInfo(
"ingress_ws_http_bridge_start account_id=%d account_type=%s payload_bytes=%d threshold_bytes=%d has_session_hash=%v store_disabled=%v",
account.ID,
account.Type,
firstPayload.payloadBytes,
s.openAIWSHTTPBridgeThresholdBytes(),
sessionHash != "",
storeDisabled,
)
currentBridgePayload := firstPayload
// Keep the first turn as the stable conversation seed. The mapped model
// is resolved again for each turn below so an in-connection model switch
// cannot reuse another model's upstream cache identity.
grokCacheSeedPayload := firstPayload.payloadRaw
var bridgeReplayInput []json.RawMessage
bridgeReplayInputExists := false
var bridgeAccountFailoverInput []json.RawMessage
bridgeAccountFailoverInputExists := false
for turn := 1; ; turn++ {
if turn > 1 && hooks != nil && hooks.BeforeRequest != nil {
if err := hooks.BeforeRequest(turn, currentBridgePayload.payloadRaw, currentBridgePayload.originalModel); err != nil {
return err
}
}
if hooks != nil && hooks.BeforeTurn != nil {
if err := hooks.BeforeTurn(turn); err != nil {
return err
}
}
if turnState != "" && c != nil && c.Request != nil {
c.Request.Header.Set(openAIWSTurnStateHeader, turnState)
}
bridgePayloadRaw := currentBridgePayload.payloadRaw
bridgePayloadBytes := currentBridgePayload.payloadBytes
needsBridgeReplay := currentBridgePayload.previousResponseID != "" || openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw)
turnReplayInput, turnReplayInputExists, replayInputErr := buildOpenAIWSReplayInputSequence(
bridgeReplayInput,
bridgeReplayInputExists,
currentBridgePayload.payloadRaw,
needsBridgeReplay,
)
if replayInputErr != nil {
return fmt.Errorf("build websocket http bridge replay input: %w", replayInputErr)
}
turnAccountFailoverInput, turnAccountFailoverInputExists, failoverInputErr := buildOpenAIWSReplayInputSequence(
bridgeAccountFailoverInput,
bridgeAccountFailoverInputExists,
currentBridgePayload.payloadRaw,
needsBridgeReplay,
)
if failoverInputErr != nil {
return fmt.Errorf("build websocket account failover input: %w", failoverInputErr)
}
if needsBridgeReplay && turnReplayInputExists {
updatedPayload, setInputErr := setOpenAIWSPayloadInputSequence(
currentBridgePayload.payloadRaw,
turnReplayInput,
true,
)
if setInputErr != nil {
return fmt.Errorf("set websocket http bridge replay input: %w", setInputErr)
}
bridgePayloadRaw = updatedPayload
bridgePayloadBytes = len(updatedPayload)
logOpenAIWSModeInfo(
"ingress_ws_http_bridge_replay_input account_id=%d turn=%d input_items=%d previous_response_id_present=%v has_tool_output=%v",
account.ID,
turn,
len(turnReplayInput),
currentBridgePayload.previousResponseID != "",
openAIWSRawPayloadHasToolCallOutput(currentBridgePayload.payloadRaw),
)
}
grokCacheIdentity := ""
if account.Platform == PlatformGrok {
grokCacheIdentity, err = resolveGrokWSCacheIdentity(
c,
account,
grokCacheSeedPayload,
currentBridgePayload.payloadRaw,
currentBridgePayload.originalModel,
)
if err != nil {
return fmt.Errorf("resolve Grok websocket cache identity: %w", err)
}
}
result, bridgeErr := s.proxyOpenAIWSHTTPBridgeTurn(
ctx,
c,
account,
token,
bridgePayloadRaw,
bridgePayloadBytes,
currentBridgePayload.originalModel,
currentBridgePayload.imageBillingModel,
currentBridgePayload.imageSizeTier,
currentBridgePayload.imageInputSize,
grokCacheIdentity,
turn,
writeClientMessage,
)
if hooks != nil && hooks.AfterTurn != nil {
hooks.AfterTurn(turn, result, bridgeErr)
}
if bridgeErr != nil {
var failoverErr *UpstreamFailoverError
if turn > 1 && errors.As(bridgeErr, &failoverErr) && failoverErr != nil {
retryPayload, retrySafe, retryPayloadErr := buildOpenAIWSCurrentTurnRetryPayload(
bridgePayloadRaw,
turnAccountFailoverInput,
turnAccountFailoverInputExists,
currentBridgePayload.originalModel,
)
if retryPayloadErr != nil {
return fmt.Errorf("build websocket current-turn failover payload: %w", retryPayloadErr)
}
if !retrySafe {
retryPayload = nil
}
return newOpenAIWSCurrentTurnFailoverError(bridgeErr, retryPayload)
}
return bridgeErr
}
if result == nil {
return errors.New("websocket http bridge turn result is nil")
}
bridgeReplayInput = cloneOpenAIWSRawMessages(turnReplayInput)
bridgeReplayInputExists = turnReplayInputExists
if result.wsReplayInputExists {
bridgeReplayInput = append(bridgeReplayInput, cloneOpenAIWSRawMessages(result.wsReplayInput)...)
bridgeReplayInputExists = true
}
bridgeAccountFailoverInput = cloneOpenAIWSRawMessages(turnAccountFailoverInput)
bridgeAccountFailoverInputExists = turnAccountFailoverInputExists
if len(result.wsAccountFailoverReplayInput) > 0 {
bridgeAccountFailoverInput = append(
bridgeAccountFailoverInput,
cloneOpenAIWSRawMessages(result.wsAccountFailoverReplayInput)...,
)
bridgeAccountFailoverInputExists = true
}
if bridgeTurnState := strings.TrimSpace(result.ResponseHeaders.Get(openAIWSTurnStateHeader)); bridgeTurnState != "" {
turnState = bridgeTurnState
if stateStore != nil && sessionHash != "" {
stateStore.BindSessionTurnState(groupID, sessionHash, bridgeTurnState, s.openAIWSSessionStickyTTL())
}
}
responseID := strings.TrimSpace(result.RequestID)
if responseID != "" && stateStore != nil {
ttl := s.openAIWSResponseStickyTTL()
logOpenAIWSBindResponseAccountWarn(groupID, account.ID, responseID, stateStore.BindResponseAccount(ctx, groupID, responseID, account.ID, ttl))
}
nextClientMessage, readErr := readClientMessage()
if readErr != nil {
if isOpenAIWSClientDisconnectError(readErr) {
closeStatus, closeReason := summarizeOpenAIWSReadCloseError(readErr)
logOpenAIWSModeInfo(
"ingress_ws_http_bridge_client_closed account_id=%d close_status=%s close_reason=%s",
account.ID,
closeStatus,
truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
)
return nil
}
return fmt.Errorf("read client websocket request: %w", readErr)
}
nextPayload, parseErr := parseClientPayload(turn+1, nextClientMessage)
if parseErr != nil {
return parseErr
}
currentBridgePayload = nextPayload
}
}
firstRoutingFields := gjson.GetManyBytes(firstPayload.payloadRaw, "model", "service_tier")
wsHeaders, _, buildHdrErr := s.buildOpenAIWSHeaders(
ctx,
c,
account,
token,
wsDecision,
isCodexCLI,
turnState,
strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)),
firstPayload.promptCacheKey,
firstRoutingFields[0].String(),
firstRoutingFields[1].String(),
)
if buildHdrErr != nil {
return fmt.Errorf("build ws headers: %w", buildHdrErr)
}
baseAcquireReq := openAIWSAcquireRequest{
Account: account,
WSURL: wsURL,
Headers: wsHeaders,
HeadersFactory: func(factoryCtx context.Context, headers http.Header) (http.Header, error) {
return s.refreshOpenAIAgentIdentityHeaders(factoryCtx, account, headers)
},
ProxyURL: func() string {
if account.ProxyID != nil && account.Proxy != nil {
return account.Proxy.URL()
}
return ""
}(),
ForceNewConn: false,
}
pool := s.getOpenAIWSConnPool()
if pool == nil {
return errors.New("openai ws conn pool is nil")
}
logOpenAIWSModeInfo(
"ingress_ws_protocol_confirm account_id=%d account_type=%s transport=%s ws_host=%s ws_path=%s ws_mode=%s store_disabled=%v has_session_hash=%v has_previous_response_id=%v",
account.ID,
account.Type,
normalizeOpenAIWSLogValue(string(wsDecision.Transport)),
wsHost,
wsPath,
normalizeOpenAIWSLogValue(ingressMode),
storeDisabled,
sessionHash != "",
firstPayload.previousResponseID != "",
)
if debugEnabled {
logOpenAIWSModeDebug(
"ingress_ws_start account_id=%d account_type=%s transport=%s ws_host=%s preferred_conn_id=%s has_session_hash=%v has_previous_response_id=%v store_disabled=%v",
account.ID,
account.Type,
normalizeOpenAIWSLogValue(string(wsDecision.Transport)),
wsHost,
truncateOpenAIWSLogValue(preferredConnID, openAIWSIDValueMaxLen),
sessionHash != "",
firstPayload.previousResponseID != "",
storeDisabled,
)
}
if firstPayload.previousResponseID != "" {
firstPreviousResponseIDKind := ClassifyOpenAIPreviousResponseIDKind(firstPayload.previousResponseID)
logOpenAIWSModeInfo(
"ingress_ws_continuation_probe account_id=%d turn=%d previous_response_id=%s previous_response_id_kind=%s preferred_conn_id=%s session_hash=%s header_session_id=%s header_conversation_id=%s has_turn_state=%v turn_state_len=%d has_prompt_cache_key=%v store_disabled=%v",
account.ID,
1,
truncateOpenAIWSLogValue(firstPayload.previousResponseID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(firstPreviousResponseIDKind),
truncateOpenAIWSLogValue(preferredConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(sessionHash, 12),
openAIWSHeaderValueForLog(baseAcquireReq.Headers, "session_id"),
openAIWSHeaderValueForLog(baseAcquireReq.Headers, "conversation_id"),
turnState != "",
len(turnState),
firstPayload.promptCacheKey != "",
storeDisabled,
)
}
acquireTimeout := s.openAIWSAcquireTimeout()
if acquireTimeout <= 0 {
acquireTimeout = 30 * time.Second
}
agentTaskRecoveryTried := false
var acquireTurnLease func(int, string, bool) (*openAIWSConnLease, error)
acquireTurnLease = func(turn int, preferred string, forcePreferredConn bool) (*openAIWSConnLease, error) {
req := cloneOpenAIWSAcquireRequest(baseAcquireReq)
req.PreferredConnID = strings.TrimSpace(preferred)
req.ForcePreferredConn = forcePreferredConn
// dedicated 模式下每次获取均新建连接,避免跨会话复用残留上下文。
req.ForceNewConn = dedicatedMode
acquireCtx, acquireCancel := context.WithTimeout(ctx, acquireTimeout)
lease, acquireErr := pool.Acquire(acquireCtx, req)
acquireCancel()
var dialErr *openAIWSDialError
if acquireErr != nil && s.isAgentIdentityAccount(ctx, account) && errors.As(acquireErr, &dialErr) && isAgentIdentityTaskInvalidWSDialError(dialErr) && !agentTaskRecoveryTried {
agentTaskRecoveryTried = true
if recoveryErr := s.recoverAgentIdentityTask(ctx, account, account.GetCredential("task_id")); recoveryErr != nil {
return nil, fmt.Errorf("agent identity task recovery failed: %w", recoveryErr)
}
return acquireTurnLease(turn, preferred, forcePreferredConn)
}
if acquireErr != nil {
canonicalModel := canonicalOpenAIAccountSchedulingModel(account, ingressSessionOriginalModel)
s.handleOpenAIWSDialTransientFailure(ctx, account, canonicalModel, acquireErr)
dialStatus, dialClass, dialCloseStatus, dialCloseReason, dialRespServer, dialRespVia, dialRespCFRay, dialRespReqID := summarizeOpenAIWSDialError(acquireErr)
logOpenAIWSModeInfo(
"ingress_ws_upstream_acquire_fail account_id=%d turn=%d reason=%s dial_status=%d dial_class=%s dial_close_status=%s dial_close_reason=%s dial_resp_server=%s dial_resp_via=%s dial_resp_cf_ray=%s dial_resp_x_request_id=%s cause=%s preferred_conn_id=%s force_preferred_conn=%v ws_host=%s ws_path=%s proxy_enabled=%v",
account.ID,
turn,
normalizeOpenAIWSLogValue(classifyOpenAIWSAcquireError(acquireErr)),
dialStatus,
dialClass,
dialCloseStatus,
truncateOpenAIWSLogValue(dialCloseReason, openAIWSHeaderValueMaxLen),
dialRespServer,
dialRespVia,
dialRespCFRay,
dialRespReqID,
truncateOpenAIWSLogValue(acquireErr.Error(), openAIWSLogValueMaxLen),
truncateOpenAIWSLogValue(preferred, openAIWSIDValueMaxLen),
forcePreferredConn,
wsHost,
wsPath,
account.ProxyID != nil && account.Proxy != nil,
)
var dialErr *openAIWSDialError
if errors.As(acquireErr, &dialErr) && dialErr != nil && dialErr.StatusCode == http.StatusTooManyRequests {
s.persistOpenAIWSRateLimitSignal(ctx, account, dialErr.ResponseHeaders, nil, "rate_limit_exceeded", "rate_limit_error", strings.TrimSpace(acquireErr.Error()))
return nil, &UpstreamFailoverError{
StatusCode: http.StatusTooManyRequests,
ResponseHeaders: cloneHeader(dialErr.ResponseHeaders),
}
}
if errors.Is(acquireErr, errOpenAIWSPreferredConnUnavailable) {
return nil, NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
"upstream continuation connection is unavailable; please restart the conversation",
acquireErr,
)
}
if errors.Is(acquireErr, context.DeadlineExceeded) || errors.Is(acquireErr, errOpenAIWSConnQueueFull) {
return nil, NewOpenAIWSClientCloseError(
coderws.StatusTryAgainLater,
"upstream websocket is busy, please retry later",
acquireErr,
)
}
return nil, acquireErr
}
connID := strings.TrimSpace(lease.ConnID())
if handshakeTurnState := strings.TrimSpace(lease.HandshakeHeader(openAIWSTurnStateHeader)); handshakeTurnState != "" {
turnState = handshakeTurnState
if stateStore != nil && sessionHash != "" {
stateStore.BindSessionTurnState(groupID, sessionHash, handshakeTurnState, s.openAIWSSessionStickyTTL())
}
updatedHeaders := cloneHeader(baseAcquireReq.Headers)
if updatedHeaders == nil {
updatedHeaders = make(http.Header)
}
updatedHeaders.Set(openAIWSTurnStateHeader, handshakeTurnState)
baseAcquireReq.Headers = updatedHeaders
}
logOpenAIWSModeInfo(
"ingress_ws_upstream_connected account_id=%d turn=%d conn_id=%s conn_reused=%v conn_pick_ms=%d queue_wait_ms=%d preferred_conn_id=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
lease.Reused(),
lease.ConnPickDuration().Milliseconds(),
lease.QueueWaitDuration().Milliseconds(),
truncateOpenAIWSLogValue(preferred, openAIWSIDValueMaxLen),
)
return lease, nil
}
sendAndRelay := func(turn int, lease *openAIWSConnLease, payload []byte, payloadBytes int, originalModel string, imageBillingModel string, imageSizeTier string, imageInputSize string) (*OpenAIForwardResult, error) {
responseModelObserver := &upstreamResponseModelObserver{}
if lease == nil {
return nil, errors.New("upstream websocket lease is nil")
}
turnStart := time.Now()
wroteDownstream := false
if err := lease.WriteJSONWithContextTimeout(ctx, json.RawMessage(payload), s.openAIWSWriteTimeout()); err != nil {
return nil, wrapOpenAIWSIngressTurnError(
"write_upstream",
fmt.Errorf("write upstream websocket request: %w", err),
false,
)
}
if debugEnabled {
logOpenAIWSModeDebug(
"ingress_ws_turn_request_sent account_id=%d turn=%d conn_id=%s payload_bytes=%d",
account.ID,
turn,
truncateOpenAIWSLogValue(lease.ConnID(), openAIWSIDValueMaxLen),
payloadBytes,
)
}
responseID := ""
usage := OpenAIUsage{}
imageCounter := newOpenAIImageOutputCounter()
var firstTokenMs *int
reqStream := openAIWSPayloadBoolFromRaw(payload, "stream", true)
turnPreviousResponseID := openAIWSPayloadStringFromRaw(payload, "previous_response_id")
turnPreviousResponseIDKind := ClassifyOpenAIPreviousResponseIDKind(turnPreviousResponseID)
turnPromptCacheKey := openAIWSPayloadStringFromRaw(payload, "prompt_cache_key")
turnStoreDisabled := s.isOpenAIWSStoreDisabledInRequestRaw(payload, account)
turnHasFunctionCallOutput := openAIWSRawPayloadHasToolCallOutput(payload)
eventCount := 0
tokenEventCount := 0
terminalEventCount := 0
replayCollector := &openAIWSToolCallReplayCollector{}
firstEventType := ""
lastEventType := ""
needModelReplace := false
clientDisconnected := false
mappedModel := ""
var mappedModelBytes []byte
if originalModel != "" {
mappedModel = strings.TrimSpace(gjson.GetBytes(payload, "model").String())
if mappedModel == "" {
mappedModel = normalizeOpenAIModelForUpstream(account, account.GetMappedModel(originalModel))
}
needModelReplace = mappedModel != "" && mappedModel != originalModel
if needModelReplace {
mappedModelBytes = []byte(mappedModel)
}
}
for {
upstreamMessage, readErr := lease.ReadMessageWithContextTimeout(ctx, s.openAIWSReadTimeout())
if readErr != nil {
lease.MarkBroken()
return nil, wrapOpenAIWSIngressTurnError(
"read_upstream",
fmt.Errorf("read upstream websocket event: %w", readErr),
wroteDownstream,
)
}
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 eventType == "error" {
canonicalModel := canonicalOpenAIAccountSchedulingModel(account, originalModel)
s.handleOpenAIWSErrorEventTransientFailure(ctx, account, canonicalModel, lease.HandshakeHeaders(), upstreamMessage)
errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(upstreamMessage)
s.persistOpenAIWSRateLimitSignal(ctx, account, lease.HandshakeHeaders(), upstreamMessage, errCodeRaw, errTypeRaw, errMsgRaw)
fallbackReason, _ := classifyOpenAIWSErrorEventFromRaw(errCodeRaw, errTypeRaw, errMsgRaw)
errCode, errType, errMessage := summarizeOpenAIWSErrorEventFieldsFromRaw(errCodeRaw, errTypeRaw, errMsgRaw)
recoverablePrevNotFound := fallbackReason == openAIWSIngressStagePreviousResponseNotFound &&
turnPreviousResponseID != "" &&
!turnHasFunctionCallOutput &&
s.openAIWSIngressPreviousResponseRecoveryEnabled() &&
!wroteDownstream
if recoverablePrevNotFound {
// 可恢复场景使用非 error 关键字日志,避免被 LegacyPrintf 误判为 ERROR 级别。
logOpenAIWSModeInfo(
"ingress_ws_prev_response_recoverable account_id=%d turn=%d conn_id=%s idx=%d reason=%s code=%s type=%s message=%s previous_response_id=%s previous_response_id_kind=%s response_id=%s store_disabled=%v has_prompt_cache_key=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(lease.ConnID(), openAIWSIDValueMaxLen),
eventCount,
truncateOpenAIWSLogValue(fallbackReason, openAIWSLogValueMaxLen),
errCode,
errType,
errMessage,
truncateOpenAIWSLogValue(turnPreviousResponseID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(turnPreviousResponseIDKind),
truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen),
turnStoreDisabled,
turnPromptCacheKey != "",
)
} else {
logOpenAIWSModeInfo(
"ingress_ws_error_event account_id=%d turn=%d conn_id=%s idx=%d fallback_reason=%s err_code=%s err_type=%s err_message=%s previous_response_id=%s previous_response_id_kind=%s response_id=%s store_disabled=%v has_prompt_cache_key=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(lease.ConnID(), openAIWSIDValueMaxLen),
eventCount,
truncateOpenAIWSLogValue(fallbackReason, openAIWSLogValueMaxLen),
errCode,
errType,
errMessage,
truncateOpenAIWSLogValue(turnPreviousResponseID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(turnPreviousResponseIDKind),
truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen),
turnStoreDisabled,
turnPromptCacheKey != "",
)
}
// previous_response_not_found 在 ingress 模式支持单次恢复重试:
// 不把该 error 直接下发客户端,而是由上层去掉 previous_response_id 后重放当前 turn。
if recoverablePrevNotFound {
lease.MarkBroken()
errMsg := strings.TrimSpace(errMsgRaw)
if errMsg == "" {
errMsg = "previous response not found"
}
return nil, wrapOpenAIWSIngressTurnError(
openAIWSIngressStagePreviousResponseNotFound,
errors.New(errMsg),
false,
)
}
if !wroteDownstream && isOpenAIWSRateLimitError(errCodeRaw, errTypeRaw, errMsgRaw) {
lease.MarkBroken()
return nil, &UpstreamFailoverError{
StatusCode: http.StatusTooManyRequests,
ResponseBody: append([]byte(nil), upstreamMessage...),
ResponseHeaders: cloneHeader(lease.HandshakeHeaders()),
}
}
}
isTokenEvent := isOpenAIWSTokenEvent(eventType)
if isTokenEvent {
tokenEventCount++
}
isTerminalEvent := isOpenAIWSTerminalEvent(eventType)
if isTerminalEvent {
terminalEventCount++
}
if firstTokenMs == nil && isTokenEvent {
ms := int(time.Since(turnStart).Milliseconds())
firstTokenMs = &ms
}
if openAIWSEventShouldParseUsage(eventType) {
parseOpenAIWSResponseUsageFromCompletedEvent(upstreamMessage, &usage)
}
imageCounter.AddSSEData(upstreamMessage)
if eventType == "response.failed" {
if hit, code, msg := detectOpenAICyberPolicy(upstreamMessage); hit {
MarkOpsCyberPolicy(c, CyberPolicyMark{
Code: code,
Message: msg,
Body: truncateString(string(upstreamMessage), 4096),
UpstreamStatus: http.StatusOK,
UpstreamInTok: usage.InputTokens,
UpstreamOutTok: usage.OutputTokens,
})
}
}
if !clientDisconnected {
if needModelReplace && len(mappedModelBytes) > 0 && openAIWSEventMayContainModel(eventType) && bytes.Contains(upstreamMessage, mappedModelBytes) {
upstreamMessage = replaceOpenAIWSMessageModel(upstreamMessage, mappedModel, originalModel)
}
if openAIWSEventMayContainToolCalls(eventType) && openAIWSMessageLikelyContainsToolCalls(upstreamMessage) {
if corrected, changed := s.toolCorrector.CorrectToolCallsInSSEBytes(upstreamMessage); changed {
upstreamMessage = corrected
}
}
replayCollector.AddEvent(eventType, upstreamMessage)
if err := writeClientMessage(upstreamMessage); err != nil {
if isOpenAIWSClientDisconnectError(err) {
clientDisconnected = true
closeStatus, closeReason := summarizeOpenAIWSReadCloseError(err)
logOpenAIWSModeInfo(
"ingress_ws_client_disconnected_drain account_id=%d turn=%d conn_id=%s close_status=%s close_reason=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(lease.ConnID(), openAIWSIDValueMaxLen),
closeStatus,
truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
)
} else {
return nil, wrapOpenAIWSIngressTurnError(
"write_client",
fmt.Errorf("write client websocket event: %w", err),
wroteDownstream,
)
}
} else {
wroteDownstream = true
}
}
if isTerminalEvent {
canonicalModel := canonicalOpenAIAccountSchedulingModel(account, originalModel)
terminalEvent := s.handleOpenAIWSTerminalTransientFailure(ctx, account, canonicalModel, lease.HandshakeHeaders(), upstreamMessage)
// 客户端已断连时,上游连接的 session 状态不可信,标记 broken 避免回池复用。
if clientDisconnected {
lease.MarkBroken()
}
firstTokenMsValue := -1
if firstTokenMs != nil {
firstTokenMsValue = *firstTokenMs
}
if debugEnabled {
logOpenAIWSModeDebug(
"ingress_ws_turn_completed account_id=%d turn=%d conn_id=%s response_id=%s 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(lease.ConnID(), openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen),
time.Since(turnStart).Milliseconds(),
eventCount,
tokenEventCount,
terminalEventCount,
truncateOpenAIWSLogValue(firstEventType, openAIWSLogValueMaxLen),
truncateOpenAIWSLogValue(lastEventType, openAIWSLogValueMaxLen),
firstTokenMsValue,
clientDisconnected,
)
}
imageCount := imageCounter.Count()
result := &OpenAIForwardResult{
RequestID: responseID,
Usage: usage,
Model: originalModel,
UpstreamModel: mappedModel,
UpstreamResponseModel: responseModelObserver.Model(),
UpstreamResponseModelConflict: responseModelObserver.Conflict(),
ServiceTier: extractOpenAIServiceTierFromBody(payload),
ReasoningEffort: ApplyThinkingEnabledFallback(extractOpenAIReasoningEffortFromBody(payload, mappedModel, originalModel), payload, mappedModel),
Stream: reqStream,
OpenAIWSMode: true,
UpstreamTerminalEvent: terminalEvent,
ResponseHeaders: lease.HandshakeHeaders(),
Duration: time.Since(turnStart),
FirstTokenMs: firstTokenMs,
}
if replayInput := replayCollector.Items(); len(replayInput) > 0 {
result.wsReplayInput = replayInput
result.wsReplayInputExists = true
}
if imageCount > 0 {
result.ImageCount = imageCount
result.ImageSize = imageSizeTier
result.ImageInputSize = imageInputSize
result.ImageOutputSizes = imageCounter.Sizes()
result.BillingModel = imageBillingModel
}
return result, nil
}
}
}
currentPayload := firstPayload.payloadRaw
currentOriginalModel := firstPayload.originalModel
currentImageBillingModel := firstPayload.imageBillingModel
currentImageSizeTier := firstPayload.imageSizeTier
currentImageInputSize := firstPayload.imageInputSize
currentPayloadBytes := firstPayload.payloadBytes
isStrictAffinityTurn := func(payload []byte) bool {
if !storeDisabled {
return false
}
return strings.TrimSpace(openAIWSPayloadStringFromRaw(payload, "previous_response_id")) != ""
}
var sessionLease *openAIWSConnLease
sessionConnID := ""
pinnedSessionConnID := ""
unpinSessionConn := func(connID string) {
connID = strings.TrimSpace(connID)
if connID == "" || pinnedSessionConnID != connID {
return
}
pool.UnpinConn(account.ID, connID)
pinnedSessionConnID = ""
}
pinSessionConn := func(connID string) {
if !storeDisabled {
return
}
connID = strings.TrimSpace(connID)
if connID == "" || pinnedSessionConnID == connID {
return
}
if pinnedSessionConnID != "" {
pool.UnpinConn(account.ID, pinnedSessionConnID)
pinnedSessionConnID = ""
}
if pool.PinConn(account.ID, connID) {
pinnedSessionConnID = connID
}
}
// lastTurnClean 标记最后一轮 sendAndRelay 是否正常完成(收到终端事件且客户端未断连)。
// 所有异常路径(读写错误、error 事件、客户端断连)已在各自分支或上层(L3403)中 MarkBroken
// 因此 releaseSessionLease 中只需在非正常结束时 MarkBroken。
lastTurnClean := false
releaseSessionLease := func() {
if sessionLease == nil {
return
}
if !lastTurnClean {
sessionLease.MarkBroken()
}
unpinSessionConn(sessionConnID)
sessionLease.Release()
if debugEnabled {
logOpenAIWSModeDebug(
"ingress_ws_upstream_released account_id=%d conn_id=%s",
account.ID,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
)
}
}
defer releaseSessionLease()
turn := 1
turnRetry := 0
turnPrevRecoveryTried := false
lastTurnFinishedAt := time.Time{}
lastTurnResponseID := ""
lastTurnPayload := []byte(nil)
var lastTurnStrictState *openAIWSIngressPreviousTurnStrictState
lastTurnReplayInput := []json.RawMessage(nil)
lastTurnReplayInputExists := false
currentTurnReplayInput := []json.RawMessage(nil)
currentTurnReplayInputExists := false
skipBeforeTurn := false
hasCurrentOrReplayFunctionCallOutput := func(payload []byte) bool {
if openAIWSRawPayloadHasToolCallOutput(payload) {
return true
}
return currentTurnReplayInputExists && openAIWSRawItemsHasFunctionCallOutput(currentTurnReplayInput)
}
resetSessionLease := func(markBroken bool) {
if sessionLease == nil {
return
}
if markBroken {
sessionLease.MarkBroken()
}
releaseSessionLease()
sessionLease = nil
sessionConnID = ""
preferredConnID = ""
}
recoverIngressPrevResponseNotFound := func(relayErr error, turn int, connID string) bool {
if !isOpenAIWSIngressPreviousResponseNotFound(relayErr) {
return false
}
if turnPrevRecoveryTried || !s.openAIWSIngressPreviousResponseRecoveryEnabled() {
return false
}
// 携带 function_call_output 的请求不能丢弃 previous_response_id
// 上游 API 需要 response chain 来匹配 tool_result 与之前的 tool_use
// 丢弃后会导致 "No tool call found for function call output" 400 错误。
if hasCurrentOrReplayFunctionCallOutput(currentPayload) {
return false
}
if isStrictAffinityTurn(currentPayload) {
// Layer 2:严格亲和链路命中 previous_response_not_found 时,降级为“去掉 previous_response_id 后重放一次”。
// 该错误说明续链锚点已失效,继续 strict fail-close 只会直接中断本轮请求。
logOpenAIWSModeInfo(
"ingress_ws_prev_response_recovery_layer2 account_id=%d turn=%d conn_id=%s store_disabled_conn_mode=%s action=drop_previous_response_id_retry",
account.ID,
turn,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(storeDisabledConnMode),
)
}
turnPrevRecoveryTried = true
updatedPayload, removed, dropErr := dropPreviousResponseIDFromRawPayload(currentPayload)
if dropErr != nil || !removed {
reason := "not_removed"
if dropErr != nil {
reason = "drop_error"
}
logOpenAIWSModeInfo(
"ingress_ws_prev_response_recovery_skip account_id=%d turn=%d conn_id=%s reason=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(reason),
)
return false
}
updatedWithInput, setInputErr := setOpenAIWSPayloadInputSequence(
updatedPayload,
currentTurnReplayInput,
currentTurnReplayInputExists,
)
if setInputErr != nil {
logOpenAIWSModeInfo(
"ingress_ws_prev_response_recovery_skip account_id=%d turn=%d conn_id=%s reason=set_full_input_error cause=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(setInputErr.Error(), openAIWSLogValueMaxLen),
)
return false
}
logOpenAIWSModeInfo(
"ingress_ws_prev_response_recovery account_id=%d turn=%d conn_id=%s action=drop_previous_response_id retry=1",
account.ID,
turn,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
)
currentPayload = updatedWithInput
currentPayloadBytes = len(updatedWithInput)
resetSessionLease(true)
skipBeforeTurn = true
return true
}
retryIngressTurn := func(relayErr error, turn int, connID string) bool {
if !isOpenAIWSIngressTurnRetryable(relayErr) || turnRetry >= 1 {
return false
}
if isStrictAffinityTurn(currentPayload) {
logOpenAIWSModeInfo(
"ingress_ws_turn_retry_skip account_id=%d turn=%d conn_id=%s reason=strict_affinity",
account.ID,
turn,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
)
return false
}
turnRetry++
logOpenAIWSModeInfo(
"ingress_ws_turn_retry account_id=%d turn=%d retry=%d reason=%s conn_id=%s",
account.ID,
turn,
turnRetry,
truncateOpenAIWSLogValue(openAIWSIngressTurnRetryReason(relayErr), openAIWSLogValueMaxLen),
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
)
resetSessionLease(true)
skipBeforeTurn = true
return true
}
for {
if turn > 1 && !skipBeforeTurn && hooks != nil && hooks.BeforeRequest != nil {
if err := hooks.BeforeRequest(turn, currentPayload, currentOriginalModel); err != nil {
return err
}
}
if !skipBeforeTurn && hooks != nil && hooks.BeforeTurn != nil {
if err := hooks.BeforeTurn(turn); err != nil {
return err
}
}
skipBeforeTurn = false
currentPreviousResponseID := openAIWSPayloadStringFromRaw(currentPayload, "previous_response_id")
expectedPrev := strings.TrimSpace(lastTurnResponseID)
toolSignals := ToolContinuationSignals{
HasFunctionCallOutput: openAIWSRawPayloadHasToolCallOutput(currentPayload),
}
if toolSignals.HasFunctionCallOutput {
var currentReqBody map[string]any
if err := json.Unmarshal(currentPayload, &currentReqBody); err == nil {
toolSignals = AnalyzeToolContinuationSignals(currentReqBody)
}
}
hasFunctionCallOutput := toolSignals.HasFunctionCallOutput
// store=false + function_call_output 场景必须有续链锚点。
// 若客户端未传 previous_response_id,优先回填上一轮响应 ID,避免上游报 call_id 无法关联。
if shouldInferIngressFunctionCallOutputPreviousResponseID(
storeDisabled,
turn,
toolSignals,
currentPreviousResponseID,
expectedPrev,
) {
updatedPayload, setPrevErr := setPreviousResponseIDToRawPayload(currentPayload, expectedPrev)
if setPrevErr != nil {
logOpenAIWSModeInfo(
"ingress_ws_function_call_output_prev_infer_skip account_id=%d turn=%d conn_id=%s reason=set_previous_response_id_error cause=%s expected_previous_response_id=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(setPrevErr.Error(), openAIWSLogValueMaxLen),
truncateOpenAIWSLogValue(expectedPrev, openAIWSIDValueMaxLen),
)
} else {
currentPayload = updatedPayload
currentPayloadBytes = len(updatedPayload)
currentPreviousResponseID = expectedPrev
logOpenAIWSModeInfo(
"ingress_ws_function_call_output_prev_infer account_id=%d turn=%d conn_id=%s action=set_previous_response_id previous_response_id=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(expectedPrev, openAIWSIDValueMaxLen),
)
}
}
nextReplayInput, nextReplayInputExists, replayInputErr := buildOpenAIWSReplayInputSequence(
lastTurnReplayInput,
lastTurnReplayInputExists,
currentPayload,
currentPreviousResponseID != "",
)
if replayInputErr != nil {
logOpenAIWSModeInfo(
"ingress_ws_replay_input_skip account_id=%d turn=%d conn_id=%s reason=build_error cause=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(replayInputErr.Error(), openAIWSLogValueMaxLen),
)
currentTurnReplayInput = nil
currentTurnReplayInputExists = false
} else {
currentTurnReplayInput = nextReplayInput
currentTurnReplayInputExists = nextReplayInputExists
}
replayHasFunctionCallOutput := currentTurnReplayInputExists &&
openAIWSRawItemsHasFunctionCallOutput(currentTurnReplayInput)
hasFunctionCallOutput = hasFunctionCallOutput || replayHasFunctionCallOutput
if storeDisabled && turn > 1 && currentPreviousResponseID != "" {
shouldKeepPreviousResponseID := false
strictReason := ""
var strictErr error
if lastTurnStrictState != nil {
shouldKeepPreviousResponseID, strictReason, strictErr = shouldKeepIngressPreviousResponseIDWithStrictState(
lastTurnStrictState,
currentPayload,
lastTurnResponseID,
hasFunctionCallOutput,
)
} else {
shouldKeepPreviousResponseID, strictReason, strictErr = shouldKeepIngressPreviousResponseID(
lastTurnPayload,
currentPayload,
lastTurnResponseID,
hasFunctionCallOutput,
)
}
if strictErr != nil {
logOpenAIWSModeInfo(
"ingress_ws_prev_response_strict_eval account_id=%d turn=%d conn_id=%s action=keep_previous_response_id reason=%s cause=%s previous_response_id=%s expected_previous_response_id=%s has_function_call_output=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(strictReason),
truncateOpenAIWSLogValue(strictErr.Error(), openAIWSLogValueMaxLen),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(expectedPrev, openAIWSIDValueMaxLen),
hasFunctionCallOutput,
)
} else if !shouldKeepPreviousResponseID {
updatedPayload, removed, dropErr := dropPreviousResponseIDFromRawPayload(currentPayload)
if dropErr != nil || !removed {
dropReason := "not_removed"
if dropErr != nil {
dropReason = "drop_error"
}
logOpenAIWSModeInfo(
"ingress_ws_prev_response_strict_eval account_id=%d turn=%d conn_id=%s action=keep_previous_response_id reason=%s drop_reason=%s previous_response_id=%s expected_previous_response_id=%s has_function_call_output=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(strictReason),
normalizeOpenAIWSLogValue(dropReason),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(expectedPrev, openAIWSIDValueMaxLen),
hasFunctionCallOutput,
)
} else {
updatedWithInput, setInputErr := setOpenAIWSPayloadInputSequence(
updatedPayload,
currentTurnReplayInput,
currentTurnReplayInputExists,
)
if setInputErr != nil {
logOpenAIWSModeInfo(
"ingress_ws_prev_response_strict_eval account_id=%d turn=%d conn_id=%s action=keep_previous_response_id reason=%s drop_reason=set_full_input_error previous_response_id=%s expected_previous_response_id=%s cause=%s has_function_call_output=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(strictReason),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(expectedPrev, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(setInputErr.Error(), openAIWSLogValueMaxLen),
hasFunctionCallOutput,
)
} else {
currentPayload = updatedWithInput
currentPayloadBytes = len(updatedWithInput)
logOpenAIWSModeInfo(
"ingress_ws_prev_response_strict_eval account_id=%d turn=%d conn_id=%s action=drop_previous_response_id_full_create reason=%s previous_response_id=%s expected_previous_response_id=%s has_function_call_output=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(strictReason),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(expectedPrev, openAIWSIDValueMaxLen),
hasFunctionCallOutput,
)
currentPreviousResponseID = ""
}
}
}
}
forcePreferredConn := isStrictAffinityTurn(currentPayload)
if sessionLease == nil {
acquiredLease, acquireErr := acquireTurnLease(turn, preferredConnID, forcePreferredConn)
if acquireErr != nil {
return fmt.Errorf("acquire upstream websocket: %w", acquireErr)
}
sessionLease = acquiredLease
sessionConnID = strings.TrimSpace(sessionLease.ConnID())
if storeDisabled {
pinSessionConn(sessionConnID)
} else {
unpinSessionConn(sessionConnID)
}
}
shouldPreflightPing := turn > 1 && sessionLease != nil && sessionLease.SupportsIdlePingWithoutReader() && turnRetry == 0
if shouldPreflightPing && openAIWSIngressPreflightPingIdle > 0 && !lastTurnFinishedAt.IsZero() {
if time.Since(lastTurnFinishedAt) < openAIWSIngressPreflightPingIdle {
shouldPreflightPing = false
}
}
if shouldPreflightPing {
if pingErr := sessionLease.PingWithTimeout(openAIWSConnHealthCheckTO); pingErr != nil {
logOpenAIWSModeInfo(
"ingress_ws_upstream_preflight_ping_fail account_id=%d turn=%d conn_id=%s cause=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(pingErr.Error(), openAIWSLogValueMaxLen),
)
if forcePreferredConn {
// 携带 function_call_output 的请求不能丢弃 previous_response_id
// 上游 API 需要 response chain 来匹配 tool_result 与之前的 tool_use
// 除非 replay input 已经包含与每个 tool_result 匹配的 tool_use 上下文。
hasFCOutput := hasFunctionCallOutput
hasReplayToolContext := hasFCOutput &&
currentTurnReplayInputExists &&
openAIWSRawItemsHaveToolCallContextForOutputs(currentTurnReplayInput)
if !turnPrevRecoveryTried && currentPreviousResponseID != "" && (!hasFCOutput || hasReplayToolContext) {
updatedPayload, removed, dropErr := dropPreviousResponseIDFromRawPayload(currentPayload)
if dropErr != nil || !removed {
reason := "not_removed"
if dropErr != nil {
reason = "drop_error"
}
logOpenAIWSModeInfo(
"ingress_ws_preflight_ping_recovery_skip account_id=%d turn=%d conn_id=%s reason=%s previous_response_id=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(reason),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
)
} else {
updatedWithInput, setInputErr := setOpenAIWSPayloadInputSequence(
updatedPayload,
currentTurnReplayInput,
currentTurnReplayInputExists,
)
if setInputErr != nil {
logOpenAIWSModeInfo(
"ingress_ws_preflight_ping_recovery_skip account_id=%d turn=%d conn_id=%s reason=set_full_input_error previous_response_id=%s cause=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(setInputErr.Error(), openAIWSLogValueMaxLen),
)
} else {
logOpenAIWSModeInfo(
"ingress_ws_preflight_ping_recovery account_id=%d turn=%d conn_id=%s action=drop_previous_response_id_retry previous_response_id=%s has_function_call_output=%v has_replay_tool_context=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
hasFCOutput,
hasReplayToolContext,
)
turnPrevRecoveryTried = true
currentPayload = updatedWithInput
currentPayloadBytes = len(updatedWithInput)
resetSessionLease(true)
skipBeforeTurn = true
continue
}
}
}
if hasFCOutput && currentPreviousResponseID != "" {
reason := "function_call_output_missing_replay_context"
if hasReplayToolContext {
reason = "function_call_output_replay_not_applied"
}
logOpenAIWSModeInfo(
"ingress_ws_preflight_ping_recovery_skip account_id=%d turn=%d conn_id=%s reason=%s action=fail_close previous_response_id=%s has_replay_tool_context=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
reason,
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
hasReplayToolContext,
)
}
resetSessionLease(true)
return NewOpenAIWSClientCloseError(
coderws.StatusPolicyViolation,
"upstream continuation connection is unavailable; please restart the conversation",
pingErr,
)
}
resetSessionLease(true)
acquiredLease, acquireErr := acquireTurnLease(turn, preferredConnID, forcePreferredConn)
if acquireErr != nil {
return fmt.Errorf("acquire upstream websocket after preflight ping fail: %w", acquireErr)
}
sessionLease = acquiredLease
sessionConnID = strings.TrimSpace(sessionLease.ConnID())
if storeDisabled {
pinSessionConn(sessionConnID)
}
}
}
connID := sessionConnID
if currentPreviousResponseID != "" {
chainedFromLast := expectedPrev != "" && currentPreviousResponseID == expectedPrev
currentPreviousResponseIDKind := ClassifyOpenAIPreviousResponseIDKind(currentPreviousResponseID)
logOpenAIWSModeInfo(
"ingress_ws_turn_chain account_id=%d turn=%d conn_id=%s previous_response_id=%s previous_response_id_kind=%s last_turn_response_id=%s chained_from_last=%v preferred_conn_id=%s header_session_id=%s header_conversation_id=%s has_turn_state=%v turn_state_len=%d has_prompt_cache_key=%v store_disabled=%v",
account.ID,
turn,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(currentPreviousResponseID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(currentPreviousResponseIDKind),
truncateOpenAIWSLogValue(expectedPrev, openAIWSIDValueMaxLen),
chainedFromLast,
truncateOpenAIWSLogValue(preferredConnID, openAIWSIDValueMaxLen),
openAIWSHeaderValueForLog(baseAcquireReq.Headers, "session_id"),
openAIWSHeaderValueForLog(baseAcquireReq.Headers, "conversation_id"),
turnState != "",
len(turnState),
openAIWSPayloadStringFromRaw(currentPayload, "prompt_cache_key") != "",
storeDisabled,
)
}
result, relayErr := sendAndRelay(turn, sessionLease, currentPayload, currentPayloadBytes, currentOriginalModel, currentImageBillingModel, currentImageSizeTier, currentImageInputSize)
if relayErr != nil {
lastTurnClean = false
if recoverIngressPrevResponseNotFound(relayErr, turn, connID) {
continue
}
if retryIngressTurn(relayErr, turn, connID) {
continue
}
finalErr := relayErr
if unwrapped := errors.Unwrap(relayErr); unwrapped != nil {
finalErr = unwrapped
}
if hooks != nil && hooks.AfterTurn != nil {
hooks.AfterTurn(turn, nil, finalErr)
}
sessionLease.MarkBroken()
return finalErr
}
turnRetry = 0
turnPrevRecoveryTried = false
lastTurnFinishedAt = time.Now()
lastTurnClean = true
if hooks != nil && hooks.AfterTurn != nil {
hooks.AfterTurn(turn, result, nil)
}
if result == nil {
return errors.New("websocket turn result is nil")
}
responseID := strings.TrimSpace(result.RequestID)
lastTurnResponseID = responseID
lastTurnPayload = cloneOpenAIWSPayloadBytes(currentPayload)
lastTurnReplayInput = cloneOpenAIWSRawMessages(currentTurnReplayInput)
lastTurnReplayInputExists = currentTurnReplayInputExists
if result.wsReplayInputExists {
lastTurnReplayInput = append(lastTurnReplayInput, cloneOpenAIWSRawMessages(result.wsReplayInput)...)
lastTurnReplayInputExists = true
}
nextStrictState, strictStateErr := buildOpenAIWSIngressPreviousTurnStrictState(currentPayload)
if strictStateErr != nil {
lastTurnStrictState = nil
logOpenAIWSModeInfo(
"ingress_ws_prev_response_strict_state_skip account_id=%d turn=%d conn_id=%s reason=build_error cause=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(strictStateErr.Error(), openAIWSLogValueMaxLen),
)
} else {
lastTurnStrictState = nextStrictState
}
if responseID != "" && stateStore != nil {
ttl := s.openAIWSResponseStickyTTL()
logOpenAIWSBindResponseAccountWarn(groupID, account.ID, responseID, stateStore.BindResponseAccount(ctx, groupID, responseID, account.ID, ttl))
stateStore.BindResponseConn(responseID, connID, ttl)
}
if stateStore != nil && storeDisabled && sessionHash != "" {
stateStore.BindSessionConn(groupID, sessionHash, connID, s.openAIWSSessionStickyTTL())
}
if connID != "" {
preferredConnID = connID
}
nextClientMessage, readErr := readClientMessage()
if readErr != nil {
if isOpenAIWSClientDisconnectError(readErr) {
closeStatus, closeReason := summarizeOpenAIWSReadCloseError(readErr)
logOpenAIWSModeInfo(
"ingress_ws_client_closed account_id=%d conn_id=%s close_status=%s close_reason=%s",
account.ID,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
closeStatus,
truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
)
return nil
}
return fmt.Errorf("read client websocket request: %w", readErr)
}
nextPayload, parseErr := parseClientPayload(turn+1, nextClientMessage)
if parseErr != nil {
return parseErr
}
nextRoutingFields := gjson.GetManyBytes(nextPayload.payloadRaw, "model", "service_tier")
if nextPayload.promptCacheKey != "" {
// ingress 会话在整个客户端 WS 生命周期内复用同一上游连接;
// prompt_cache_key 对握手头的更新仅在未来需要重新建连时生效。
updatedHeaders, _, updHdrErr := s.buildOpenAIWSHeaders(
ctx,
c,
account,
token,
wsDecision,
isCodexCLI,
turnState,
strings.TrimSpace(c.GetHeader(openAIWSTurnMetadataHeader)),
nextPayload.promptCacheKey,
nextRoutingFields[0].String(),
nextRoutingFields[1].String(),
)
if updHdrErr != nil {
logOpenAIWSModeInfo("ingress_ws_update_headers_failed account_id=%d err=%v", account.ID, updHdrErr)
} else {
baseAcquireReq.Headers = updatedHeaders
}
}
setOpenAICodexRoutingHint(baseAcquireReq.Headers, account, nextRoutingFields[0].String(), nextRoutingFields[1].String())
if nextPayload.previousResponseID != "" {
expectedPrev := strings.TrimSpace(lastTurnResponseID)
chainedFromLast := expectedPrev != "" && nextPayload.previousResponseID == expectedPrev
nextPreviousResponseIDKind := ClassifyOpenAIPreviousResponseIDKind(nextPayload.previousResponseID)
logOpenAIWSModeInfo(
"ingress_ws_next_turn_chain account_id=%d turn=%d next_turn=%d conn_id=%s previous_response_id=%s previous_response_id_kind=%s last_turn_response_id=%s chained_from_last=%v has_prompt_cache_key=%v store_disabled=%v",
account.ID,
turn,
turn+1,
truncateOpenAIWSLogValue(connID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(nextPayload.previousResponseID, openAIWSIDValueMaxLen),
normalizeOpenAIWSLogValue(nextPreviousResponseIDKind),
truncateOpenAIWSLogValue(expectedPrev, openAIWSIDValueMaxLen),
chainedFromLast,
nextPayload.promptCacheKey != "",
storeDisabled,
)
}
if stateStore != nil && nextPayload.previousResponseID != "" {
if stickyConnID, ok := stateStore.GetResponseConn(nextPayload.previousResponseID); ok {
if sessionConnID != "" && stickyConnID != "" && stickyConnID != sessionConnID {
logOpenAIWSModeInfo(
"ingress_ws_keep_session_conn account_id=%d turn=%d conn_id=%s sticky_conn_id=%s previous_response_id=%s",
account.ID,
turn,
truncateOpenAIWSLogValue(sessionConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(stickyConnID, openAIWSIDValueMaxLen),
truncateOpenAIWSLogValue(nextPayload.previousResponseID, openAIWSIDValueMaxLen),
)
} else {
preferredConnID = stickyConnID
}
}
}
currentPayload = nextPayload.payloadRaw
currentOriginalModel = nextPayload.originalModel
currentImageBillingModel = nextPayload.imageBillingModel
currentImageSizeTier = nextPayload.imageSizeTier
currentImageInputSize = nextPayload.imageInputSize
currentPayloadBytes = nextPayload.payloadBytes
storeDisabled = s.isOpenAIWSStoreDisabledInRequestRaw(currentPayload, account)
if !storeDisabled {
unpinSessionConn(sessionConnID)
}
turn++
}
}