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
1757 lines
69 KiB
Go
1757 lines
69 KiB
Go
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, ¤tReqBody); 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++
|
||
}
|
||
}
|