792 lines
24 KiB
Go
792 lines
24 KiB
Go
package service
|
||||
|
|
|
|||
|
|
import (
|
|||
|
|
"bytes"
|
|||
|
|
"context"
|
|||
|
|
"encoding/json"
|
|||
|
|
"errors"
|
|||
|
|
"fmt"
|
|||
|
|
"net/http"
|
|||
|
|
"net/url"
|
|||
|
|
"strings"
|
|||
|
|
|
|||
|
|
"github.com/gin-gonic/gin"
|
|||
|
|
"github.com/tidwall/gjson"
|
|||
|
|
"github.com/tidwall/sjson"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
func validateOpenAIWSBearerToken(account *Account, token string) error {
|
|||
|
|
if account == nil {
|
|||
|
|
return errors.New("account is nil")
|
|||
|
|
}
|
|||
|
|
if strings.TrimSpace(token) == "" && !account.IsOpenAIAgentIdentity() {
|
|||
|
|
return errors.New("token is empty")
|
|||
|
|
}
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func (s *OpenAIGatewayService) buildOpenAIResponsesWSURL(account *Account) (string, error) {
|
|||
|
|
if account == nil {
|
|||
|
|
return "", errors.New("account is nil")
|
|||
|
|
}
|
|||
|
|
var targetURL string
|
|||
|
|
switch account.Type {
|
|||
|
|
case AccountTypeOAuth:
|
|||
|
|
targetURL = chatgptCodexURL
|
|||
|
|
case AccountTypeAPIKey:
|
|||
|
|
baseURL := account.GetOpenAIBaseURL()
|
|||
|
|
if baseURL == "" {
|
|||
|
|
targetURL = openaiPlatformAPIURL
|
|||
|
|
} else {
|
|||
|
|
validatedURL, err := s.validateUpstreamBaseURL(baseURL)
|
|||
|
|
if err != nil {
|
|||
|
|
return "", err
|
|||
|
|
}
|
|||
|
|
targetURL = buildOpenAIResponsesURLForPlatform(account.Platform, validatedURL)
|
|||
|
|
}
|
|||
|
|
default:
|
|||
|
|
targetURL = openaiPlatformAPIURL
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
parsed, err := url.Parse(strings.TrimSpace(targetURL))
|
|||
|
|
if err != nil {
|
|||
|
|
return "", fmt.Errorf("invalid target url: %w", err)
|
|||
|
|
}
|
|||
|
|
switch strings.ToLower(parsed.Scheme) {
|
|||
|
|
case "https":
|
|||
|
|
parsed.Scheme = "wss"
|
|||
|
|
case "http":
|
|||
|
|
parsed.Scheme = "ws"
|
|||
|
|
case "wss", "ws":
|
|||
|
|
// 保持不变
|
|||
|
|
default:
|
|||
|
|
return "", fmt.Errorf("unsupported scheme for ws: %s", parsed.Scheme)
|
|||
|
|
}
|
|||
|
|
return parsed.String(), nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func (s *OpenAIGatewayService) buildOpenAIWSHeaders(
|
|||
|
|
ctx context.Context,
|
|||
|
|
c *gin.Context,
|
|||
|
|
account *Account,
|
|||
|
|
token string,
|
|||
|
|
decision OpenAIWSProtocolDecision,
|
|||
|
|
isCodexCLI bool,
|
|||
|
|
turnState string,
|
|||
|
|
turnMetadata string,
|
|||
|
|
promptCacheKey string,
|
|||
|
|
routingModel string,
|
|||
|
|
routingServiceTier string,
|
|||
|
|
) (http.Header, openAIWSSessionHeaderResolution, error) {
|
|||
|
|
headers := make(http.Header)
|
|||
|
|
if account == nil || !account.IsOpenAIAgentIdentity() {
|
|||
|
|
headers.Set("authorization", "Bearer "+token)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
sessionResolution := resolveOpenAIWSSessionHeaders(c, promptCacheKey)
|
|||
|
|
if c != nil && c.Request != nil {
|
|||
|
|
if v := strings.TrimSpace(c.Request.Header.Get("accept-language")); v != "" {
|
|||
|
|
headers.Set("accept-language", v)
|
|||
|
|
}
|
|||
|
|
for _, value := range c.Request.Header.Values("x-codex-beta-features") {
|
|||
|
|
if value = strings.TrimSpace(value); value != "" {
|
|||
|
|
headers.Add("x-codex-beta-features", value)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
for _, name := range [...]string{
|
|||
|
|
"x-codex-window-id",
|
|||
|
|
"x-codex-installation-id",
|
|||
|
|
"session-id",
|
|||
|
|
"thread-id",
|
|||
|
|
"x-client-request-id",
|
|||
|
|
} {
|
|||
|
|
if value := c.Request.Header.Get(name); strings.TrimSpace(value) != "" {
|
|||
|
|
headers.Set(name, value)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
// 真实 Codex 的 WS 握手同样携带会话级 x-codex-beta-features
|
|||
|
|
// (client.rs build_websocket_headers 复用 build_responses_headers),
|
|||
|
|
// 客户端未声明时补成默认形态,与 HTTP 出站保持一致。放在客户端头拷贝
|
|||
|
|
// 之外:该头是账号/会话级属性,不依赖入站请求是否存在,也避免预热与
|
|||
|
|
// 实际请求因头差异落进不同的连接池兼容分桶。
|
|||
|
|
applyOpenAICodexBetaFeatures(c, account, headers)
|
|||
|
|
// OAuth 账号:将 apiKeyID 混入 session 标识符,防止跨用户会话碰撞。
|
|||
|
|
if account != nil && account.Type == AccountTypeOAuth {
|
|||
|
|
apiKeyID := getAPIKeyIDFromContext(c)
|
|||
|
|
if sessionResolution.SessionID != "" {
|
|||
|
|
headers.Set("session_id", isolateOpenAISessionID(apiKeyID, sessionResolution.SessionID))
|
|||
|
|
}
|
|||
|
|
if sessionResolution.ConversationID != "" {
|
|||
|
|
headers.Set("conversation_id", isolateOpenAISessionID(apiKeyID, sessionResolution.ConversationID))
|
|||
|
|
}
|
|||
|
|
} else {
|
|||
|
|
if sessionResolution.SessionID != "" {
|
|||
|
|
headers.Set("session_id", sessionResolution.SessionID)
|
|||
|
|
}
|
|||
|
|
if sessionResolution.ConversationID != "" {
|
|||
|
|
headers.Set("conversation_id", sessionResolution.ConversationID)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
if state := strings.TrimSpace(turnState); state != "" {
|
|||
|
|
headers.Set(openAIWSTurnStateHeader, state)
|
|||
|
|
}
|
|||
|
|
if metadata := strings.TrimSpace(turnMetadata); metadata != "" {
|
|||
|
|
headers.Set(openAIWSTurnMetadataHeader, metadata)
|
|||
|
|
}
|
|||
|
|
applyStagedCodexFingerprintHeaders(c, account, headers)
|
|||
|
|
|
|||
|
|
if account != nil && account.Type == AccountTypeOAuth {
|
|||
|
|
if err := resolveAndSetOpenAIChatGPTAccountHeaders(ctx, s.accountRepo, headers, account); err != nil {
|
|||
|
|
return nil, sessionResolution, fmt.Errorf("resolve chatgpt account headers: %w", err)
|
|||
|
|
}
|
|||
|
|
headers.Set("originator", resolveOpenAIUpstreamOriginator(c, isCodexCLI))
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
betaValue := openAIWSBetaV2Value
|
|||
|
|
if decision.Transport == OpenAIUpstreamTransportResponsesWebsocket {
|
|||
|
|
betaValue = openAIWSBetaV1Value
|
|||
|
|
}
|
|||
|
|
headers.Set("OpenAI-Beta", betaValue)
|
|||
|
|
|
|||
|
|
customUA := ""
|
|||
|
|
if account != nil {
|
|||
|
|
customUA = account.GetOpenAIUserAgent()
|
|||
|
|
}
|
|||
|
|
if strings.TrimSpace(customUA) != "" {
|
|||
|
|
headers.Set("user-agent", customUA)
|
|||
|
|
} else if c != nil {
|
|||
|
|
if ua := strings.TrimSpace(c.GetHeader("User-Agent")); ua != "" {
|
|||
|
|
headers.Set("user-agent", ua)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
if s != nil && s.cfg != nil && s.cfg.Gateway.ForceCodexCLI {
|
|||
|
|
headers.Set("user-agent", CodexCanonicalUserAgent())
|
|||
|
|
}
|
|||
|
|
// 终态收口:WS 握手与 HTTP 出站共用同一套身份语义,账号级自定义 UA 同样作为
|
|||
|
|
// 管理员显式配置传入(上面写进 headers 的值只在强制统一被关闭时才参与配对)。
|
|||
|
|
if account != nil && account.Type == AccountTypeOAuth {
|
|||
|
|
enforceCodexIdentityHeadersWithUA(headers, s.codexIdentityOverrideUA(account))
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// 账号级请求头覆写(仅 openai api_key 账号启用时生效;OAuth 路径 no-op)。
|
|||
|
|
// 覆盖所有 WS 模式(ctx_pool/dedicated/passthrough)的握手头。
|
|||
|
|
account.ApplyHeaderOverrides(headers)
|
|||
|
|
setOpenAICodexRoutingHint(headers, account, routingModel, routingServiceTier)
|
|||
|
|
logOpenAIRoutingDiagnostics(
|
|||
|
|
ctx,
|
|||
|
|
account,
|
|||
|
|
string(decision.Transport),
|
|||
|
|
routingModel,
|
|||
|
|
routingServiceTier,
|
|||
|
|
strings.TrimSpace(headers.Get(openAICodexRoutingHintHeader)) != "",
|
|||
|
|
"soft_routing_hint",
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
return headers, sessionResolution, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func (s *OpenAIGatewayService) buildOpenAIWSCreatePayload(reqBody map[string]any, account *Account) map[string]any {
|
|||
|
|
// OpenAI WS Mode 协议:response.create 字段与 HTTP /responses 基本一致。
|
|||
|
|
// 保留 stream 字段(与 Codex CLI 一致),仅移除 background。
|
|||
|
|
payload := make(map[string]any, len(reqBody)+1)
|
|||
|
|
for k, v := range reqBody {
|
|||
|
|
payload[k] = v
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
delete(payload, "background")
|
|||
|
|
if _, exists := payload["stream"]; !exists {
|
|||
|
|
payload["stream"] = true
|
|||
|
|
}
|
|||
|
|
payload["type"] = "response.create"
|
|||
|
|
|
|||
|
|
// OAuth 默认保持 store=false,避免误依赖服务端历史。
|
|||
|
|
if account != nil && account.Type == AccountTypeOAuth && !s.isOpenAIWSStoreRecoveryAllowed(account) {
|
|||
|
|
payload["store"] = false
|
|||
|
|
}
|
|||
|
|
return payload
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func setOpenAIWSTurnMetadata(payload map[string]any, turnMetadata string) {
|
|||
|
|
if len(payload) == 0 {
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
metadata := strings.TrimSpace(turnMetadata)
|
|||
|
|
if metadata == "" {
|
|||
|
|
return
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
switch existing := payload["client_metadata"].(type) {
|
|||
|
|
case map[string]any:
|
|||
|
|
existing[openAIWSTurnMetadataHeader] = metadata
|
|||
|
|
payload["client_metadata"] = existing
|
|||
|
|
case map[string]string:
|
|||
|
|
next := make(map[string]any, len(existing)+1)
|
|||
|
|
for k, v := range existing {
|
|||
|
|
next[k] = v
|
|||
|
|
}
|
|||
|
|
next[openAIWSTurnMetadataHeader] = metadata
|
|||
|
|
payload["client_metadata"] = next
|
|||
|
|
default:
|
|||
|
|
payload["client_metadata"] = map[string]any{
|
|||
|
|
openAIWSTurnMetadataHeader: metadata,
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func (s *OpenAIGatewayService) isOpenAIWSStoreRecoveryAllowed(account *Account) bool {
|
|||
|
|
if account != nil && account.IsOpenAIWSAllowStoreRecoveryEnabled() {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
if s != nil && s.cfg != nil && s.cfg.Gateway.OpenAIWS.AllowStoreRecovery {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func (s *OpenAIGatewayService) isOpenAIWSStoreDisabledInRequest(reqBody map[string]any, account *Account) bool {
|
|||
|
|
if account != nil && account.Type == AccountTypeOAuth && !s.isOpenAIWSStoreRecoveryAllowed(account) {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
if len(reqBody) == 0 {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
rawStore, ok := reqBody["store"]
|
|||
|
|
if !ok {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
storeEnabled, ok := rawStore.(bool)
|
|||
|
|
if !ok {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
return !storeEnabled
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func (s *OpenAIGatewayService) isOpenAIWSStoreDisabledInRequestRaw(reqBody []byte, account *Account) bool {
|
|||
|
|
if account != nil && account.Type == AccountTypeOAuth && !s.isOpenAIWSStoreRecoveryAllowed(account) {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
if len(reqBody) == 0 {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
storeValue := gjson.GetBytes(reqBody, "store")
|
|||
|
|
if !storeValue.Exists() {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
if storeValue.Type != gjson.True && storeValue.Type != gjson.False {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
return !storeValue.Bool()
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func (s *OpenAIGatewayService) openAIWSStoreDisabledConnMode() string {
|
|||
|
|
if s == nil || s.cfg == nil {
|
|||
|
|
return openAIWSStoreDisabledConnModeStrict
|
|||
|
|
}
|
|||
|
|
mode := strings.ToLower(strings.TrimSpace(s.cfg.Gateway.OpenAIWS.StoreDisabledConnMode))
|
|||
|
|
switch mode {
|
|||
|
|
case openAIWSStoreDisabledConnModeStrict, openAIWSStoreDisabledConnModeAdaptive, openAIWSStoreDisabledConnModeOff:
|
|||
|
|
return mode
|
|||
|
|
case "":
|
|||
|
|
// 兼容旧配置:仅配置了布尔开关时按旧语义推导。
|
|||
|
|
if s.cfg.Gateway.OpenAIWS.StoreDisabledForceNewConn {
|
|||
|
|
return openAIWSStoreDisabledConnModeStrict
|
|||
|
|
}
|
|||
|
|
return openAIWSStoreDisabledConnModeOff
|
|||
|
|
default:
|
|||
|
|
return openAIWSStoreDisabledConnModeStrict
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func shouldForceNewConnOnStoreDisabled(mode, lastFailureReason string) bool {
|
|||
|
|
switch mode {
|
|||
|
|
case openAIWSStoreDisabledConnModeOff:
|
|||
|
|
return false
|
|||
|
|
case openAIWSStoreDisabledConnModeAdaptive:
|
|||
|
|
reason := strings.TrimPrefix(strings.TrimSpace(lastFailureReason), "prewarm_")
|
|||
|
|
switch reason {
|
|||
|
|
case "policy_violation", "message_too_big", "auth_failed", "write_request", "write":
|
|||
|
|
return true
|
|||
|
|
default:
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
default:
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func dropPreviousResponseIDFromRawPayload(payload []byte) ([]byte, bool, error) {
|
|||
|
|
return dropPreviousResponseIDFromRawPayloadWithDeleteFn(payload, sjson.DeleteBytes)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func dropPreviousResponseIDFromRawPayloadWithDeleteFn(
|
|||
|
|
payload []byte,
|
|||
|
|
deleteFn func([]byte, string) ([]byte, error),
|
|||
|
|
) ([]byte, bool, error) {
|
|||
|
|
if len(payload) == 0 {
|
|||
|
|
return payload, false, nil
|
|||
|
|
}
|
|||
|
|
if !gjson.GetBytes(payload, "previous_response_id").Exists() {
|
|||
|
|
return payload, false, nil
|
|||
|
|
}
|
|||
|
|
if deleteFn == nil {
|
|||
|
|
deleteFn = sjson.DeleteBytes
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
updated := payload
|
|||
|
|
for i := 0; i < openAIWSMaxPrevResponseIDDeletePasses &&
|
|||
|
|
gjson.GetBytes(updated, "previous_response_id").Exists(); i++ {
|
|||
|
|
next, err := deleteFn(updated, "previous_response_id")
|
|||
|
|
if err != nil {
|
|||
|
|
return payload, false, err
|
|||
|
|
}
|
|||
|
|
updated = next
|
|||
|
|
}
|
|||
|
|
return updated, !gjson.GetBytes(updated, "previous_response_id").Exists(), nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func setPreviousResponseIDToRawPayload(payload []byte, previousResponseID string) ([]byte, error) {
|
|||
|
|
normalizedPrevID := strings.TrimSpace(previousResponseID)
|
|||
|
|
if len(payload) == 0 || normalizedPrevID == "" {
|
|||
|
|
return payload, nil
|
|||
|
|
}
|
|||
|
|
updated, err := sjson.SetBytes(payload, "previous_response_id", normalizedPrevID)
|
|||
|
|
if err == nil {
|
|||
|
|
return updated, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
var reqBody map[string]any
|
|||
|
|
if unmarshalErr := json.Unmarshal(payload, &reqBody); unmarshalErr != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
reqBody["previous_response_id"] = normalizedPrevID
|
|||
|
|
rebuilt, marshalErr := json.Marshal(reqBody)
|
|||
|
|
if marshalErr != nil {
|
|||
|
|
return nil, marshalErr
|
|||
|
|
}
|
|||
|
|
return rebuilt, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func shouldInferIngressFunctionCallOutputPreviousResponseID(
|
|||
|
|
storeDisabled bool,
|
|||
|
|
turn int,
|
|||
|
|
signals ToolContinuationSignals,
|
|||
|
|
currentPreviousResponseID string,
|
|||
|
|
expectedPreviousResponseID string,
|
|||
|
|
) bool {
|
|||
|
|
if !storeDisabled || turn <= 1 || !signals.HasFunctionCallOutput {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
if strings.TrimSpace(currentPreviousResponseID) != "" {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
if signals.HasFunctionCallOutputMissingCallID {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
// If the client already sent the actual tool-call context, treat this as
|
|||
|
|
// a full replay / self-contained continuation payload rather than
|
|||
|
|
// downgrading it into an inferred delta continuation. item_reference alone
|
|||
|
|
// is not enough on the store=false WS path: it still needs a valid prior
|
|||
|
|
// response anchor so upstream can resolve the referenced function_call.
|
|||
|
|
if signals.HasToolCallContext {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
return strings.TrimSpace(expectedPreviousResponseID) != ""
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func alignStoreDisabledPreviousResponseID(
|
|||
|
|
payload []byte,
|
|||
|
|
expectedPreviousResponseID string,
|
|||
|
|
) ([]byte, bool, error) {
|
|||
|
|
if len(payload) == 0 {
|
|||
|
|
return payload, false, nil
|
|||
|
|
}
|
|||
|
|
expected := strings.TrimSpace(expectedPreviousResponseID)
|
|||
|
|
if expected == "" {
|
|||
|
|
return payload, false, nil
|
|||
|
|
}
|
|||
|
|
current := openAIWSPayloadStringFromRaw(payload, "previous_response_id")
|
|||
|
|
if current == "" || current == expected {
|
|||
|
|
return payload, false, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
withoutPrev, removed, dropErr := dropPreviousResponseIDFromRawPayload(payload)
|
|||
|
|
if dropErr != nil {
|
|||
|
|
return payload, false, dropErr
|
|||
|
|
}
|
|||
|
|
if !removed {
|
|||
|
|
return payload, false, nil
|
|||
|
|
}
|
|||
|
|
updated, setErr := setPreviousResponseIDToRawPayload(withoutPrev, expected)
|
|||
|
|
if setErr != nil {
|
|||
|
|
return payload, false, setErr
|
|||
|
|
}
|
|||
|
|
return updated, true, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func cloneOpenAIWSPayloadBytes(payload []byte) []byte {
|
|||
|
|
if len(payload) == 0 {
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
cloned := make([]byte, len(payload))
|
|||
|
|
copy(cloned, payload)
|
|||
|
|
return cloned
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func cloneOpenAIWSRawMessages(items []json.RawMessage) []json.RawMessage {
|
|||
|
|
if items == nil {
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
cloned := make([]json.RawMessage, 0, len(items))
|
|||
|
|
for idx := range items {
|
|||
|
|
cloned = append(cloned, json.RawMessage(cloneOpenAIWSPayloadBytes(items[idx])))
|
|||
|
|
}
|
|||
|
|
return cloned
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func normalizeOpenAIWSJSONForCompare(raw []byte) ([]byte, error) {
|
|||
|
|
trimmed := bytes.TrimSpace(raw)
|
|||
|
|
if len(trimmed) == 0 {
|
|||
|
|
return nil, errors.New("json is empty")
|
|||
|
|
}
|
|||
|
|
var decoded any
|
|||
|
|
if err := json.Unmarshal(trimmed, &decoded); err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return json.Marshal(decoded)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func normalizeOpenAIWSJSONForCompareOrRaw(raw []byte) []byte {
|
|||
|
|
normalized, err := normalizeOpenAIWSJSONForCompare(raw)
|
|||
|
|
if err != nil {
|
|||
|
|
return bytes.TrimSpace(raw)
|
|||
|
|
}
|
|||
|
|
return normalized
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(payload []byte) ([]byte, error) {
|
|||
|
|
if len(payload) == 0 {
|
|||
|
|
return nil, errors.New("payload is empty")
|
|||
|
|
}
|
|||
|
|
var decoded map[string]any
|
|||
|
|
if err := json.Unmarshal(payload, &decoded); err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
delete(decoded, "input")
|
|||
|
|
delete(decoded, "previous_response_id")
|
|||
|
|
// Codex changes transport-only metadata for every response.create. These fields
|
|||
|
|
// do not alter the context referenced by previous_response_id and are excluded
|
|||
|
|
// from Codex's own websocket reuse comparison.
|
|||
|
|
delete(decoded, "client_metadata")
|
|||
|
|
delete(decoded, "stream_options")
|
|||
|
|
// Official Codex prewarms a connection with generate=false, then omits the
|
|||
|
|
// field on the business request that continues from the prewarm response.
|
|||
|
|
// Only normalize false so a meaningful generate=true change remains visible.
|
|||
|
|
if generate, ok := decoded["generate"].(bool); ok && !generate {
|
|||
|
|
delete(decoded, "generate")
|
|||
|
|
}
|
|||
|
|
return json.Marshal(decoded)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func openAIWSExtractNormalizedInputSequence(payload []byte) ([]json.RawMessage, bool, error) {
|
|||
|
|
if len(payload) == 0 {
|
|||
|
|
return nil, false, nil
|
|||
|
|
}
|
|||
|
|
inputValue := gjson.GetBytes(payload, "input")
|
|||
|
|
if !inputValue.Exists() {
|
|||
|
|
return nil, false, nil
|
|||
|
|
}
|
|||
|
|
if inputValue.Type == gjson.JSON {
|
|||
|
|
raw := strings.TrimSpace(inputValue.Raw)
|
|||
|
|
if strings.HasPrefix(raw, "[") {
|
|||
|
|
var items []json.RawMessage
|
|||
|
|
if err := json.Unmarshal([]byte(raw), &items); err != nil {
|
|||
|
|
return nil, true, err
|
|||
|
|
}
|
|||
|
|
return items, true, nil
|
|||
|
|
}
|
|||
|
|
return []json.RawMessage{json.RawMessage(raw)}, true, nil
|
|||
|
|
}
|
|||
|
|
if inputValue.Type == gjson.String {
|
|||
|
|
encoded, _ := json.Marshal(inputValue.String())
|
|||
|
|
return []json.RawMessage{encoded}, true, nil
|
|||
|
|
}
|
|||
|
|
return []json.RawMessage{json.RawMessage(inputValue.Raw)}, true, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func openAIWSInputIsPrefixExtended(previousPayload, currentPayload []byte) (bool, error) {
|
|||
|
|
previousItems, previousExists, prevErr := openAIWSExtractNormalizedInputSequence(previousPayload)
|
|||
|
|
if prevErr != nil {
|
|||
|
|
return false, prevErr
|
|||
|
|
}
|
|||
|
|
currentItems, currentExists, currentErr := openAIWSExtractNormalizedInputSequence(currentPayload)
|
|||
|
|
if currentErr != nil {
|
|||
|
|
return false, currentErr
|
|||
|
|
}
|
|||
|
|
if !previousExists && !currentExists {
|
|||
|
|
return true, nil
|
|||
|
|
}
|
|||
|
|
if !previousExists {
|
|||
|
|
return len(currentItems) == 0, nil
|
|||
|
|
}
|
|||
|
|
if !currentExists {
|
|||
|
|
return len(previousItems) == 0, nil
|
|||
|
|
}
|
|||
|
|
if len(currentItems) < len(previousItems) {
|
|||
|
|
return false, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
for idx := range previousItems {
|
|||
|
|
previousNormalized := normalizeOpenAIWSJSONForCompareOrRaw(previousItems[idx])
|
|||
|
|
currentNormalized := normalizeOpenAIWSJSONForCompareOrRaw(currentItems[idx])
|
|||
|
|
if !bytes.Equal(previousNormalized, currentNormalized) {
|
|||
|
|
return false, nil
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return true, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func openAIWSRawItemsHasPrefix(items []json.RawMessage, prefix []json.RawMessage) bool {
|
|||
|
|
if len(prefix) == 0 {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
if len(items) < len(prefix) {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
for idx := range prefix {
|
|||
|
|
previousNormalized := normalizeOpenAIWSJSONForCompareOrRaw(prefix[idx])
|
|||
|
|
currentNormalized := normalizeOpenAIWSJSONForCompareOrRaw(items[idx])
|
|||
|
|
if !bytes.Equal(previousNormalized, currentNormalized) {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func openAIWSRawItemsHasFunctionCallOutput(items []json.RawMessage) bool {
|
|||
|
|
for _, item := range items {
|
|||
|
|
if isCodexToolCallOutputItemType(gjson.GetBytes(item, "type").String()) {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func openAIWSRawItemsHaveToolCallContextForOutputs(items []json.RawMessage) bool {
|
|||
|
|
if len(items) == 0 {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
contextCallIDs := make(map[string]struct{})
|
|||
|
|
outputCallIDs := make(map[string]struct{})
|
|||
|
|
for _, item := range items {
|
|||
|
|
itemType := gjson.GetBytes(item, "type").String()
|
|||
|
|
callID := strings.TrimSpace(gjson.GetBytes(item, "call_id").String())
|
|||
|
|
switch {
|
|||
|
|
case isCodexToolCallContextItemType(itemType):
|
|||
|
|
if callID != "" {
|
|||
|
|
contextCallIDs[callID] = struct{}{}
|
|||
|
|
}
|
|||
|
|
case isCodexToolCallOutputItemType(itemType):
|
|||
|
|
if callID == "" {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
outputCallIDs[callID] = struct{}{}
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
if len(outputCallIDs) == 0 || len(contextCallIDs) == 0 {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
for callID := range outputCallIDs {
|
|||
|
|
if _, ok := contextCallIDs[callID]; !ok {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func openAIWSRawPayloadHasToolCallOutput(payload []byte) bool {
|
|||
|
|
if len(payload) == 0 {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
input := gjson.GetBytes(payload, "input")
|
|||
|
|
if !input.Exists() {
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
if input.IsArray() {
|
|||
|
|
for _, item := range input.Array() {
|
|||
|
|
if isCodexToolCallOutputItemType(item.Get("type").String()) {
|
|||
|
|
return true
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
if input.Type == gjson.JSON {
|
|||
|
|
return isCodexToolCallOutputItemType(input.Get("type").String())
|
|||
|
|
}
|
|||
|
|
return false
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func buildOpenAIWSReplayInputSequence(
|
|||
|
|
previousFullInput []json.RawMessage,
|
|||
|
|
previousFullInputExists bool,
|
|||
|
|
currentPayload []byte,
|
|||
|
|
hasPreviousResponseID bool,
|
|||
|
|
) ([]json.RawMessage, bool, error) {
|
|||
|
|
currentItems, currentExists, currentErr := openAIWSExtractNormalizedInputSequence(currentPayload)
|
|||
|
|
if currentErr != nil {
|
|||
|
|
return nil, false, currentErr
|
|||
|
|
}
|
|||
|
|
if !hasPreviousResponseID {
|
|||
|
|
return cloneOpenAIWSRawMessages(currentItems), currentExists, nil
|
|||
|
|
}
|
|||
|
|
if !previousFullInputExists {
|
|||
|
|
return cloneOpenAIWSRawMessages(currentItems), currentExists, nil
|
|||
|
|
}
|
|||
|
|
if !currentExists || len(currentItems) == 0 {
|
|||
|
|
return cloneOpenAIWSRawMessages(previousFullInput), true, nil
|
|||
|
|
}
|
|||
|
|
if openAIWSRawItemsHasPrefix(currentItems, previousFullInput) {
|
|||
|
|
return cloneOpenAIWSRawMessages(currentItems), true, nil
|
|||
|
|
}
|
|||
|
|
merged := make([]json.RawMessage, 0, len(previousFullInput)+len(currentItems))
|
|||
|
|
merged = append(merged, cloneOpenAIWSRawMessages(previousFullInput)...)
|
|||
|
|
merged = append(merged, cloneOpenAIWSRawMessages(currentItems)...)
|
|||
|
|
return merged, true, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func setOpenAIWSPayloadInputSequence(
|
|||
|
|
payload []byte,
|
|||
|
|
fullInput []json.RawMessage,
|
|||
|
|
fullInputExists bool,
|
|||
|
|
) ([]byte, error) {
|
|||
|
|
if !fullInputExists {
|
|||
|
|
return payload, nil
|
|||
|
|
}
|
|||
|
|
// Preserve [] vs null semantics when input exists but is empty.
|
|||
|
|
inputForMarshal := fullInput
|
|||
|
|
if inputForMarshal == nil {
|
|||
|
|
inputForMarshal = []json.RawMessage{}
|
|||
|
|
}
|
|||
|
|
inputRaw, marshalErr := json.Marshal(inputForMarshal)
|
|||
|
|
if marshalErr != nil {
|
|||
|
|
return nil, marshalErr
|
|||
|
|
}
|
|||
|
|
return sjson.SetRawBytes(payload, "input", inputRaw)
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func buildOpenAIWSCurrentTurnRetryPayload(
|
|||
|
|
payload []byte,
|
|||
|
|
fullInput []json.RawMessage,
|
|||
|
|
fullInputExists bool,
|
|||
|
|
originalModel string,
|
|||
|
|
) ([]byte, bool, error) {
|
|||
|
|
if !fullInputExists {
|
|||
|
|
return nil, false, nil
|
|||
|
|
}
|
|||
|
|
retryPayload, err := setOpenAIWSPayloadInputSequence(payload, fullInput, true)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, false, err
|
|||
|
|
}
|
|||
|
|
retryPayload = RemovePreviousResponseIDFromBody(retryPayload)
|
|||
|
|
if model := strings.TrimSpace(originalModel); model != "" {
|
|||
|
|
retryPayload, err = sjson.SetBytes(retryPayload, "model", model)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, false, err
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
coverage := AnalyzeToolCallOutputContextCoverageBytes(retryPayload)
|
|||
|
|
if coverage.HasFunctionCallOutput && !coverage.ContextCoversAllCallIDs {
|
|||
|
|
return nil, false, nil
|
|||
|
|
}
|
|||
|
|
return retryPayload, true, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func shouldKeepIngressPreviousResponseID(
|
|||
|
|
previousPayload []byte,
|
|||
|
|
currentPayload []byte,
|
|||
|
|
lastTurnResponseID string,
|
|||
|
|
hasFunctionCallOutput bool,
|
|||
|
|
) (bool, string, error) {
|
|||
|
|
if hasFunctionCallOutput {
|
|||
|
|
return true, "has_function_call_output", nil
|
|||
|
|
}
|
|||
|
|
currentPreviousResponseID := strings.TrimSpace(openAIWSPayloadStringFromRaw(currentPayload, "previous_response_id"))
|
|||
|
|
if currentPreviousResponseID == "" {
|
|||
|
|
return false, "missing_previous_response_id", nil
|
|||
|
|
}
|
|||
|
|
expectedPreviousResponseID := strings.TrimSpace(lastTurnResponseID)
|
|||
|
|
if expectedPreviousResponseID == "" {
|
|||
|
|
return false, "missing_last_turn_response_id", nil
|
|||
|
|
}
|
|||
|
|
if currentPreviousResponseID != expectedPreviousResponseID {
|
|||
|
|
return false, "previous_response_id_mismatch", nil
|
|||
|
|
}
|
|||
|
|
if len(previousPayload) == 0 {
|
|||
|
|
return false, "missing_previous_turn_payload", nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
previousComparable, previousComparableErr := normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(previousPayload)
|
|||
|
|
if previousComparableErr != nil {
|
|||
|
|
return false, "non_input_compare_error", previousComparableErr
|
|||
|
|
}
|
|||
|
|
currentComparable, currentComparableErr := normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(currentPayload)
|
|||
|
|
if currentComparableErr != nil {
|
|||
|
|
return false, "non_input_compare_error", currentComparableErr
|
|||
|
|
}
|
|||
|
|
if !bytes.Equal(previousComparable, currentComparable) {
|
|||
|
|
return false, "non_input_changed", nil
|
|||
|
|
}
|
|||
|
|
return true, "strict_incremental_ok", nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
type openAIWSIngressPreviousTurnStrictState struct {
|
|||
|
|
nonInputComparable []byte
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func buildOpenAIWSIngressPreviousTurnStrictState(payload []byte) (*openAIWSIngressPreviousTurnStrictState, error) {
|
|||
|
|
if len(payload) == 0 {
|
|||
|
|
return nil, nil
|
|||
|
|
}
|
|||
|
|
nonInputComparable, nonInputErr := normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(payload)
|
|||
|
|
if nonInputErr != nil {
|
|||
|
|
return nil, nonInputErr
|
|||
|
|
}
|
|||
|
|
return &openAIWSIngressPreviousTurnStrictState{
|
|||
|
|
nonInputComparable: nonInputComparable,
|
|||
|
|
}, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
func shouldKeepIngressPreviousResponseIDWithStrictState(
|
|||
|
|
previousState *openAIWSIngressPreviousTurnStrictState,
|
|||
|
|
currentPayload []byte,
|
|||
|
|
lastTurnResponseID string,
|
|||
|
|
hasFunctionCallOutput bool,
|
|||
|
|
) (bool, string, error) {
|
|||
|
|
if hasFunctionCallOutput {
|
|||
|
|
return true, "has_function_call_output", nil
|
|||
|
|
}
|
|||
|
|
currentPreviousResponseID := strings.TrimSpace(openAIWSPayloadStringFromRaw(currentPayload, "previous_response_id"))
|
|||
|
|
if currentPreviousResponseID == "" {
|
|||
|
|
return false, "missing_previous_response_id", nil
|
|||
|
|
}
|
|||
|
|
expectedPreviousResponseID := strings.TrimSpace(lastTurnResponseID)
|
|||
|
|
if expectedPreviousResponseID == "" {
|
|||
|
|
return false, "missing_last_turn_response_id", nil
|
|||
|
|
}
|
|||
|
|
if currentPreviousResponseID != expectedPreviousResponseID {
|
|||
|
|
return false, "previous_response_id_mismatch", nil
|
|||
|
|
}
|
|||
|
|
if previousState == nil {
|
|||
|
|
return false, "missing_previous_turn_payload", nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
currentComparable, currentComparableErr := normalizeOpenAIWSPayloadWithoutInputAndPreviousResponseID(currentPayload)
|
|||
|
|
if currentComparableErr != nil {
|
|||
|
|
return false, "non_input_compare_error", currentComparableErr
|
|||
|
|
}
|
|||
|
|
if !bytes.Equal(previousState.nonInputComparable, currentComparable) {
|
|||
|
|
return false, "non_input_changed", nil
|
|||
|
|
}
|
|||
|
|
return true, "strict_incremental_ok", nil
|
|||
|
|
}
|