diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index 79709cabb7..a0a96713ae 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -55,6 +55,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) { h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit)) return } + logRequestBodyReadFailure(reqLog, c.Request, err) h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body") return } diff --git a/backend/internal/handler/openai_gateway_credential_failover_test.go b/backend/internal/handler/openai_gateway_credential_failover_test.go index 046b7ec52a..3bb86898de 100644 --- a/backend/internal/handler/openai_gateway_credential_failover_test.go +++ b/backend/internal/handler/openai_gateway_credential_failover_test.go @@ -197,7 +197,7 @@ func TestOpsClassificationTreatsCredentialFailureAsAuthNotInference(t *testing.T require.Equal(t, http.StatusForbidden, entry.UpstreamErrors[0].UpstreamStatusCode) } -func TestOpsRecoveredCredentialFailoverUsesAccountAuthAttribution(t *testing.T) { +func TestOpsRecoveredCredentialFailoverDoesNotCreateRequestError(t *testing.T) { setupOpsErrorLogTestQueue(t, 2) gin.SetMode(gin.TestMode) ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) @@ -217,21 +217,72 @@ func TestOpsRecoveredCredentialFailoverUsesAccountAuthAttribution(t *testing.T) recorder := httptest.NewRecorder() router.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/openai/v1/responses", nil)) + require.Equal(t, http.StatusOK, recorder.Code) + require.Zero(t, OpsErrorLogQueueLength()) + select { + case job := <-opsErrorLogQueue: + t.Fatalf("successful failover must not create ops error row: %+v", job.entry) + default: + } +} + +func TestOpsWebSocketCredentialFailoverSuccessDoesNotCreateRequestError(t *testing.T) { + setupOpsErrorLogTestQueue(t, 2) + gin.SetMode(gin.TestMode) + ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.Use(OpsErrorLoggerMiddleware(ops)) + router.GET("/openai/v1/responses", func(c *gin.Context) { + c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{ + Stage: string(service.GatewayFailureStageAccountAuth), Scope: string(service.GatewayFailureScopeAccount), + Reason: string(service.GrokCredentialReasonRevoked), Message: "Grok OAuth credentials require account action", + }}) + }) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, "/openai/v1/responses", nil) + request.Header.Set("Connection", "Upgrade") + request.Header.Set("Upgrade", "websocket") + router.ServeHTTP(recorder, request) + + require.Equal(t, http.StatusOK, recorder.Code) + require.Zero(t, OpsErrorLogQueueLength()) + select { + case job := <-opsErrorLogQueue: + t.Fatalf("successful websocket failover must not create ops error row: %+v", job.entry) + default: + } +} + +func TestOpsWebSocketCredentialFailoverExhaustedIsRecorded(t *testing.T) { + setupOpsErrorLogTestQueue(t, 2) + gin.SetMode(gin.TestMode) + ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.Use(OpsErrorLoggerMiddleware(ops)) + router.GET("/openai/v1/responses", func(c *gin.Context) { + c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{ + Stage: string(service.GatewayFailureStageAccountAuth), Scope: string(service.GatewayFailureScopeAccount), + Reason: string(service.GrokCredentialReasonRevoked), Message: "Grok OAuth credentials require account action", + }}) + closeOpenAIWSFailoverExhausted(c, nil, &service.UpstreamFailoverError{ + Stage: service.GatewayFailureStageAccountAuth, + Scope: service.GatewayFailureScopeAccount, + Reason: service.GrokCredentialReasonRevoked, + NextAccountAction: service.NextAccountStop, + }) + }) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodGet, "/openai/v1/responses", nil) + request.Header.Set("Connection", "Upgrade") + request.Header.Set("Upgrade", "websocket") + router.ServeHTTP(recorder, request) + require.Equal(t, http.StatusOK, recorder.Code) require.Equal(t, int64(1), OpsErrorLogQueueLength()) job := <-opsErrorLogQueue require.Equal(t, "account_auth", job.entry.ErrorPhase) - require.Equal(t, "provider", job.entry.ErrorOwner) - require.Equal(t, "gateway", job.entry.ErrorSource) - require.Contains(t, job.entry.ErrorMessage, "Recovered account authentication failure") - require.NotContains(t, job.entry.ErrorMessage, "403") - require.NotContains(t, job.entry.ErrorMessage, "earlier inference failure") - require.NotNil(t, job.entry.UpstreamStatusCode) - require.Zero(t, *job.entry.UpstreamStatusCode) - require.Nil(t, job.entry.UpstreamErrors) - require.NotNil(t, job.entry.UpstreamErrorsJSON) - events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON) - require.NoError(t, err) - require.Len(t, events, 2) - require.Equal(t, http.StatusForbidden, events[0].UpstreamStatusCode) + require.Equal(t, http.StatusServiceUnavailable, job.entry.StatusCode) + require.Equal(t, service.GrokCredentialUnavailableClientMessage, job.entry.ErrorMessage) } diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index aebb68f543..4549d29697 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -320,6 +320,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { h.errorResponse(c, http.StatusRequestEntityTooLarge, "invalid_request_error", buildBodyTooLargeMessage(maxErr.Limit)) return } + logRequestBodyReadFailure(reqLog, c.Request, err) h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "Failed to read request body") return } @@ -839,6 +840,11 @@ func isOpenAIRemoteCompactionV2Request(body []byte) bool { func (h *OpenAIGatewayHandler) normalizeOpenAIResponsesCompactRequest(c *gin.Context, reqLog *zap.Logger, body []byte) ([]byte, bool) { isCompactRequest := isOpenAILegacyCompactPath(c) if !isCompactRequest && isBareOpenAIResponsesPath(c) && service.HasCompactionTriggerInInput(body) { + if normalized, changed, err := service.NormalizeCompactionTriggerInputOrder(body); err != nil { + reqLog.Warn("codex.remote_compact.trigger_order_normalization_failed", zap.Error(err)) + } else if changed { + body = normalized + } if isOpenAIRemoteCompactionV2Request(body) { return body, true } @@ -1908,7 +1914,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { } releaseAccountSlot() if !failoverErr.ShouldRetryNextAccount() { - closeOpenAIWSFailoverExhausted(wsConn, failoverErr) + closeOpenAIWSFailoverExhausted(c, wsConn, failoverErr) return false } if ctx.Err() != nil { @@ -1918,12 +1924,12 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { failedAccountIDs[account.ID] = struct{}{} lastFailoverErr = failoverErr if switchCount >= maxAccountSwitches { - closeOpenAIWSFailoverExhausted(wsConn, failoverErr) + closeOpenAIWSFailoverExhausted(c, wsConn, failoverErr) return false } switchCount++ if h.gatewayService.ShouldStopOpenAIOAuth429Failover(account, failoverErr.StatusCode, switchCount, &oauth429FailoverState) { - closeOpenAIWSFailoverExhausted(wsConn, failoverErr) + closeOpenAIWSFailoverExhausted(c, wsConn, failoverErr) return false } reqLog.Warn("openai.websocket_upstream_failover_switching", @@ -1979,7 +1985,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { zap.Int("excluded_account_count", len(failedAccountIDs)), ) if lastFailoverErr != nil { - closeOpenAIWSFailoverExhausted(wsConn, lastFailoverErr) + closeOpenAIWSFailoverExhausted(c, wsConn, lastFailoverErr) } else { closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "no available account") } @@ -1987,7 +1993,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { } if selection == nil || selection.Account == nil { if lastFailoverErr != nil { - closeOpenAIWSFailoverExhausted(wsConn, lastFailoverErr) + closeOpenAIWSFailoverExhausted(c, wsConn, lastFailoverErr) } else { closeOpenAIClientWS(wsConn, coderws.StatusTryAgainLater, "no available account") } @@ -2329,7 +2335,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { retryPayload, retryCurrentTurn := service.OpenAIWSCurrentTurnRetryPayload(err) nextAttemptMessage, retrySafe := openAIWSNextAttemptMessage(wsAttemptMessage, retryPayload, retryCurrentTurn) if !retrySafe { - closeOpenAIWSFailoverExhausted(wsConn, failoverErr) + closeOpenAIWSFailoverExhausted(c, wsConn, failoverErr) return } wsAttemptMessage = nextAttemptMessage @@ -2993,25 +2999,44 @@ func openAIWSNextAttemptMessage(current, retryPayload []byte, retryCurrentTurn b return append([]byte(nil), retryPayload...), true } -func closeOpenAIWSFailoverExhausted(conn *coderws.Conn, failoverErr *service.UpstreamFailoverError) { - if failoverErr == nil { - closeOpenAIClientWS(conn, coderws.StatusInternalError, "upstream websocket proxy failed") - return - } - if failoverErr.Stage == service.GatewayFailureStageAccountAuth { - closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, service.GrokCredentialUnavailableClientMessage) - return - } - switch failoverErr.StatusCode { - case http.StatusTooManyRequests: - closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, "upstream rate limit exceeded, please retry later") - case 529, http.StatusInternalServerError, http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout: - closeOpenAIClientWS(conn, coderws.StatusTryAgainLater, "upstream service temporarily unavailable") - case http.StatusUnauthorized, http.StatusForbidden: - closeOpenAIClientWS(conn, coderws.StatusPolicyViolation, "upstream websocket authentication failed") - default: - closeOpenAIClientWS(conn, coderws.StatusInternalError, "upstream websocket proxy failed") +func closeOpenAIWSFailoverExhausted(c *gin.Context, conn *coderws.Conn, failoverErr *service.UpstreamFailoverError) { + intendedStatus := http.StatusBadGateway + errorType := "upstream_error" + errorCode := "upstream_ws_failover_exhausted" + message := "upstream websocket proxy failed" + closeStatus := coderws.StatusInternalError + + if failoverErr != nil { + if reason := strings.TrimSpace(string(failoverErr.Reason)); reason != "" { + errorCode = reason + } + if failoverErr.Stage == service.GatewayFailureStageAccountAuth { + intendedStatus = http.StatusServiceUnavailable + errorType = "api_error" + message = service.GrokCredentialUnavailableClientMessage + closeStatus = coderws.StatusTryAgainLater + } else { + switch failoverErr.StatusCode { + case http.StatusTooManyRequests: + intendedStatus = http.StatusTooManyRequests + errorType = "rate_limit_error" + message = "upstream rate limit exceeded, please retry later" + closeStatus = coderws.StatusTryAgainLater + case 529, http.StatusInternalServerError, http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout: + intendedStatus = failoverErr.StatusCode + message = "upstream service temporarily unavailable" + closeStatus = coderws.StatusTryAgainLater + case http.StatusUnauthorized, http.StatusForbidden: + intendedStatus = failoverErr.StatusCode + errorType = "authentication_error" + message = "upstream websocket authentication failed" + closeStatus = coderws.StatusPolicyViolation + } + } } + + service.MarkOpsStreamFailure(c, errorType, errorCode, message, intendedStatus) + closeOpenAIClientWS(conn, closeStatus, message) } func writeContentModerationWSError(ctx context.Context, conn *coderws.Conn, decision *service.ContentModerationDecision) { diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go index a5343d9077..59e1c80cd7 100644 --- a/backend/internal/handler/ops_error_logger.go +++ b/backend/internal/handler/ops_error_logger.go @@ -518,9 +518,11 @@ func isOpsNoAvailableAccountError(err error) bool { type opsCaptureWriter struct { gin.ResponseWriter - limit int - buf bytes.Buffer - ctx *gin.Context + limit int + buf bytes.Buffer + probe []byte + sseCapturing bool + ctx *gin.Context } const opsCaptureWriterLimit = service.OpsErrorLogQueueBodyMaxBytes @@ -541,6 +543,8 @@ func acquireOpsCaptureWriter(rw gin.ResponseWriter) *opsCaptureWriter { w.ResponseWriter = rw w.limit = opsCaptureWriterLimit w.buf.Reset() + w.probe = w.probe[:0] + w.sseCapturing = false return w } @@ -551,6 +555,8 @@ func releaseOpsCaptureWriter(w *opsCaptureWriter) { w.ResponseWriter = nil w.ctx = nil w.limit = opsCaptureWriterLimit + w.probe = w.probe[:0] + w.sseCapturing = false if !shouldPoolOpsCaptureWriter(w) { return } @@ -638,13 +644,8 @@ func (w *opsCaptureWriter) Write(b []byte) (int, error) { if w.ResponseWriter == nil { return 0, nil } - if w.shouldCapture() && w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit { - remaining := w.limit - w.buf.Len() - if len(b) > remaining { - _, _ = w.buf.Write(b[:remaining]) - } else { - _, _ = w.buf.Write(b) - } + if w.shouldCapture() { + w.captureResponseChunk(b, w.Status()) } return w.ResponseWriter.Write(b) } @@ -653,17 +654,94 @@ func (w *opsCaptureWriter) WriteString(s string) (int, error) { if w.ResponseWriter == nil { return 0, nil } - if w.shouldCapture() && w.Status() >= 400 && w.limit > 0 && w.buf.Len() < w.limit { - remaining := w.limit - w.buf.Len() - if len(s) > remaining { - _, _ = w.buf.WriteString(s[:remaining]) - } else { - _, _ = w.buf.WriteString(s) - } + if w.shouldCapture() { + w.captureResponseChunk([]byte(s), w.Status()) } return w.ResponseWriter.WriteString(s) } +var opsTerminalSSEMarkers = [][]byte{ + []byte("event: response.failed"), + []byte("event: error"), + []byte(`data: {"type":"response.failed"`), + []byte(`data:{"type":"response.failed"`), + []byte(`data: {"type":"error"`), + []byte(`data:{"type":"error"`), +} + +func (w *opsCaptureWriter) captureResponseChunk(chunk []byte, status int) { + if w == nil || w.limit <= 0 || w.buf.Len() >= w.limit || len(chunk) == 0 { + return + } + if status >= 400 || w.sseCapturing { + w.appendCapturedResponse(chunk) + return + } + + combined := make([]byte, 0, len(w.probe)+len(chunk)) + combined = append(combined, w.probe...) + combined = append(combined, chunk...) + if start := findOpsTerminalSSEStart(combined); start >= 0 { + w.sseCapturing = true + w.probe = w.probe[:0] + w.appendCapturedResponse(combined[start:]) + return + } + + // Retain one full marker width so a marker split across writes still has + // its preceding byte available for the SSE line-boundary check. + keep := opsTerminalSSEProbeSize + if keep > len(combined) { + keep = len(combined) + } + w.probe = append(w.probe[:0], combined[len(combined)-keep:]...) +} + +func (w *opsCaptureWriter) appendCapturedResponse(chunk []byte) { + remaining := w.limit - w.buf.Len() + if remaining <= 0 { + return + } + if len(chunk) > remaining { + chunk = chunk[:remaining] + } + _, _ = w.buf.Write(chunk) +} + +func findOpsTerminalSSEStart(data []byte) int { + earliest := -1 + for _, marker := range opsTerminalSSEMarkers { + searchFrom := 0 + for searchFrom < len(data) { + idx := bytes.Index(data[searchFrom:], marker) + if idx < 0 { + break + } + idx += searchFrom + if idx == 0 || data[idx-1] == '\n' { + if earliest < 0 || idx < earliest { + earliest = idx + } + break + } + searchFrom = idx + 1 + } + } + return earliest +} + +func maxOpsTerminalSSEMarkerSize() int { + maxSize := 0 + for _, marker := range opsTerminalSSEMarkers { + if len(marker) > maxSize { + maxSize = len(marker) + } + } + return maxSize +} + +var opsTerminalSSEProbeSize = maxOpsTerminalSSEMarkerSize() + func (w *opsCaptureWriter) shouldCapture() bool { if w.ctx == nil { return true @@ -675,7 +753,7 @@ func (w *opsCaptureWriter) shouldCapture() bool { // OpsErrorLoggerMiddleware records error responses (status >= 400) into ops_error_logs. // // Notes: -// - It buffers response bodies only when status >= 400 to avoid overhead for successful traffic. +// - It buffers response bodies only for status >= 400 or terminal SSE frames. // - Streaming errors after the response has started (SSE) may still need explicit logging. func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { return func(c *gin.Context) { @@ -709,268 +787,20 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { } status := c.Writer.Status() + body := w.buf.Bytes() + parsed := parseOpsErrorResponse(body) if status < 400 { - // Even when the client request succeeds, we still want to persist upstream error attempts - // (retries/failover) so ops can observe upstream instability that gets "covered" by retries. - var events []*service.OpsUpstreamErrorEvent - if v, ok := c.Get(service.OpsUpstreamErrorsKey); ok { - if arr, ok := v.([]*service.OpsUpstreamErrorEvent); ok && len(arr) > 0 { - events = arr - } - } - // Also accept single upstream fields set by gateway services (rare for successful requests). - hasUpstreamContext := len(events) > 0 - if !hasUpstreamContext { - if v, ok := c.Get(service.OpsUpstreamStatusCodeKey); ok { - switch t := v.(type) { - case int: - hasUpstreamContext = t > 0 - case int64: - hasUpstreamContext = t > 0 - } - } - } - if !hasUpstreamContext { - if v, ok := c.Get(service.OpsUpstreamErrorMessageKey); ok { - if s, ok := v.(string); ok && strings.TrimSpace(s) != "" { - hasUpstreamContext = true - } - } - } - if !hasUpstreamContext { - if v, ok := c.Get(service.OpsUpstreamErrorDetailKey); ok { - if s, ok := v.(string); ok && strings.TrimSpace(s) != "" { - hasUpstreamContext = true - } - } - } - if !hasUpstreamContext { - // 没有上游错误上下文,但网关可能在已固化的 200 流上就地补发了 SSE 错误帧 - // (如 ping 等待后并发超限、Wait 后二次计费校验失败)。这类失败若不在此补记, - // 会因 wire 状态码为 200 而在错误看板里彻底隐形。 + if parsed.StreamFailure { + status = inferStreamFailureStatus(c, parsed) + } else { + // Locally generated in-band errors use an explicit context marker and + // may not have a capturable terminal frame. Preserve that fallback, + // but never turn recovered upstream attempts into request errors. logOpsStreamError(c, ops, status) return } - - apiKey := getOpsAPIKey(c) - clientRequestID, _ := c.Request.Context().Value(ctxkey.ClientRequestID).(string) - - model, _ := c.Get(opsModelKey) - streamV, _ := c.Get(opsStreamKey) - accountIDV, _ := c.Get(opsAccountIDKey) - - var modelName string - if s, ok := model.(string); ok { - modelName = s - } - stream := false - if b, ok := streamV.(bool); ok { - stream = b - } - - // Prefer showing the account that experienced the upstream error (if we have events), - // otherwise fall back to the final selected account (best-effort). - var accountID *int64 - if len(events) > 0 { - if last := events[len(events)-1]; last != nil && last.AccountID > 0 { - v := last.AccountID - accountID = &v - } - } - if accountID == nil { - if v, ok := accountIDV.(int64); ok && v > 0 { - accountID = &v - } - } - - fallbackPlatform := guessPlatformFromPath(c.Request.URL.Path) - platform := resolveOpsPlatform(c.Request.Context(), apiKey, fallbackPlatform) - - requestID := c.Writer.Header().Get("X-Request-Id") - if requestID == "" { - requestID = c.Writer.Header().Get("x-request-id") - } - - // Best-effort backfill single upstream fields from the last event (if present). - var upstreamStatusCode *int - var upstreamErrorMessage *string - var upstreamErrorDetail *string - finalAccountAuth := false - if len(events) > 0 { - last := events[len(events)-1] - if last != nil { - finalAccountAuth = last.Stage == string(service.GatewayFailureStageAccountAuth) - if finalAccountAuth { - code := 0 - upstreamStatusCode = &code - } else if last.UpstreamStatusCode > 0 { - code := last.UpstreamStatusCode - upstreamStatusCode = &code - } - if msg := strings.TrimSpace(last.Message); msg != "" { - upstreamErrorMessage = &msg - } - if detail := strings.TrimSpace(last.Detail); detail != "" { - upstreamErrorDetail = &detail - } - } - } - - if !finalAccountAuth && upstreamStatusCode == nil { - if v, ok := c.Get(service.OpsUpstreamStatusCodeKey); ok { - switch t := v.(type) { - case int: - if t > 0 { - code := t - upstreamStatusCode = &code - } - case int64: - if t > 0 { - code := int(t) - upstreamStatusCode = &code - } - } - } - } - if !finalAccountAuth && upstreamErrorMessage == nil { - if v, ok := c.Get(service.OpsUpstreamErrorMessageKey); ok { - if s, ok := v.(string); ok && strings.TrimSpace(s) != "" { - msg := strings.TrimSpace(s) - upstreamErrorMessage = &msg - } - } - } - if !finalAccountAuth && upstreamErrorDetail == nil { - if v, ok := c.Get(service.OpsUpstreamErrorDetailKey); ok { - if s, ok := v.(string); ok && strings.TrimSpace(s) != "" { - detail := strings.TrimSpace(s) - upstreamErrorDetail = &detail - } - } - } - - // If we still have nothing meaningful, skip. - if upstreamStatusCode == nil && upstreamErrorMessage == nil && upstreamErrorDetail == nil && len(events) == 0 { - return - } - - effectiveUpstreamStatus := 0 - if upstreamStatusCode != nil { - effectiveUpstreamStatus = *upstreamStatusCode - } - - recoveredMsg := "Recovered upstream error" - if finalAccountAuth { - recoveredMsg = "Recovered account authentication failure" - } else if effectiveUpstreamStatus > 0 { - recoveredMsg += " " + strconvItoa(effectiveUpstreamStatus) - } - if upstreamErrorMessage != nil && strings.TrimSpace(*upstreamErrorMessage) != "" { - recoveredMsg += ": " + strings.TrimSpace(*upstreamErrorMessage) - } - recoveredMsg = truncateString(recoveredMsg, 2048) - recoveredPhase, recoveredBusinessLimited, recoveredOwner, recoveredSource := classifyOpsErrorLog( - c, "upstream_error", recoveredMsg, "", effectiveUpstreamStatus, - ) - - entry := &service.OpsInsertErrorLogInput{ - RequestID: requestID, - ClientRequestID: clientRequestID, - - AccountID: accountID, - Platform: platform, - Model: modelName, - RequestPath: func() string { - if c.Request != nil && c.Request.URL != nil { - return c.Request.URL.Path - } - return "" - }(), - Stream: stream, - InboundEndpoint: GetInboundEndpoint(c), - UpstreamEndpoint: GetUpstreamEndpoint(c, platform), - RequestedModel: modelName, - UpstreamModel: func() string { - if v, ok := c.Get(opsUpstreamModelKey); ok { - if s, ok := v.(string); ok { - return strings.TrimSpace(s) - } - } - return "" - }(), - RequestType: func() *int16 { - if v, ok := c.Get(opsRequestTypeKey); ok { - switch t := v.(type) { - case int16: - return &t - case int: - v16 := int16(t) - return &v16 - } - } - return nil - }(), - UserAgent: c.GetHeader("User-Agent"), - - ErrorPhase: recoveredPhase, - ErrorType: "upstream_error", - // Severity should reflect the upstream failure, not the final client status (200). - Severity: classifyOpsSeverity("upstream_error", effectiveUpstreamStatus), - StatusCode: status, - IsBusinessLimited: recoveredBusinessLimited, - IsCountTokens: isCountTokensRequest(c), - - ErrorMessage: recoveredMsg, - ErrorBody: "", - - ErrorSource: recoveredSource, - ErrorOwner: recoveredOwner, - - UpstreamStatusCode: upstreamStatusCode, - UpstreamErrorMessage: upstreamErrorMessage, - UpstreamErrorDetail: upstreamErrorDetail, - UpstreamErrors: events, - - CreatedAt: time.Now(), - } - applyOpsLatencyFieldsFromContext(c, entry) - applyOpsUpstreamFieldsFromContext(c, entry) - - if apiKey != nil { - entry.APIKeyID = &apiKey.ID - entry.APIKeyPrefix = keyPrefix(apiKey.Key, 8) - if apiKey.User != nil { - entry.UserID = &apiKey.User.ID - } - if apiKey.GroupID != nil { - entry.GroupID = apiKey.GroupID - } - // Prefer group platform if present (more stable than inferring from path). - if apiKey.Group != nil && apiKey.Group.Platform != "" { - entry.Platform = apiKey.Group.Platform - } - } - - var clientIP string - if ip := strings.TrimSpace(ip.GetClientIP(c)); ip != "" { - clientIP = ip - entry.ClientIP = &clientIP - } - - // Skip logging if a passthrough rule with skip_monitoring=true matched. - if v, ok := c.Get(service.OpsSkipPassthroughKey); ok { - if skip, _ := v.(bool); skip { - return - } - } - - enqueueOpsErrorLog(ops, entry) - return } - body := w.buf.Bytes() - parsed := parseOpsErrorResponse(body) - // Skip logging if a passthrough rule with skip_monitoring=true matched. if v, ok := c.Get(service.OpsSkipPassthroughKey); ok { if skip, _ := v.(bool); skip { @@ -1007,9 +837,13 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { fallbackPlatform := guessPlatformFromPath(c.Request.URL.Path) platform := resolveOpsPlatform(c.Request.Context(), apiKey, fallbackPlatform) - requestID := c.Writer.Header().Get("X-Request-Id") + requestID, _ := c.Request.Context().Value(ctxkey.RequestID).(string) + requestID = strings.TrimSpace(requestID) if requestID == "" { - requestID = c.Writer.Header().Get("x-request-id") + requestID = c.Writer.Header().Get("X-Request-Id") + if requestID == "" { + requestID = c.Writer.Header().Get("x-request-id") + } } normalizedType := normalizeOpsErrorType(parsed.ErrorType, parsed.Code) @@ -1073,6 +907,15 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { } applyOpsLatencyFieldsFromContext(c, entry) applyOpsUpstreamFieldsFromContext(c, entry) + if parsed.StreamFailure { + if message := strings.TrimSpace(parsed.Message); message != "" { + entry.UpstreamErrorMessage = &message + } + if status >= 400 { + finalStatus := status + entry.UpstreamStatusCode = &finalStatus + } + } if apiKey != nil { entry.APIKeyID = &apiKey.ID @@ -1162,9 +1005,13 @@ func logOpsStreamError(c *gin.Context, ops *service.OpsService, wireStatus int) fallbackPlatform := guessPlatformFromPath(c.Request.URL.Path) platform := resolveOpsPlatform(c.Request.Context(), apiKey, fallbackPlatform) - requestID := c.Writer.Header().Get("X-Request-Id") + requestID, _ := c.Request.Context().Value(ctxkey.RequestID).(string) + requestID = strings.TrimSpace(requestID) if requestID == "" { - requestID = c.Writer.Header().Get("x-request-id") + requestID = c.Writer.Header().Get("X-Request-Id") + if requestID == "" { + requestID = c.Writer.Header().Get("x-request-id") + } } entry := &service.OpsInsertErrorLogInput{ @@ -1222,6 +1069,7 @@ func logOpsStreamError(c *gin.Context, ops *service.OpsService, wireStatus int) CreatedAt: time.Now(), } applyOpsLatencyFieldsFromContext(c, entry) + applyOpsUpstreamFieldsFromContext(c, entry) if apiKey != nil { entry.APIKeyID = &apiKey.ID @@ -1366,15 +1214,19 @@ func getContextLatencyMs(c *gin.Context, key string) *int64 { } type parsedOpsError struct { - ErrorType string - Message string - Code string + ErrorType string + Message string + Code string + StreamFailure bool } func parseOpsErrorResponse(body []byte) parsedOpsError { if len(body) == 0 { return parsedOpsError{} } + if parsed, ok := parseOpsSSEFailure(body); ok { + return parsed + } // Fast path: attempt to decode into a generic map. var m map[string]any @@ -1386,17 +1238,9 @@ func parseOpsErrorResponse(body []byte) parsedOpsError { if errObj, ok := m["error"].(map[string]any); ok { t, _ := errObj["type"].(string) msg, _ := errObj["message"].(string) - // Gemini googleError also uses "error": { code, message, status } - if msg == "" { - if v, ok := errObj["message"]; ok { - msg, _ = v.(string) - } - } if t == "" { - // Gemini error does not have "type" field. t = "api_error" } - // For gemini error, capture numeric code as string for business-limited mapping if needed. var code string if v, ok := errObj["code"]; ok { switch n := v.(type) { @@ -1421,6 +1265,153 @@ func parseOpsErrorResponse(body []byte) parsedOpsError { return parsedOpsError{Message: truncateString(string(body), 1024)} } +func parseOpsSSEFailure(body []byte) (parsedOpsError, bool) { + normalized := strings.ReplaceAll(string(body), "\r\n", "\n") + if findOpsTerminalSSEStart([]byte(normalized)) < 0 { + return parsedOpsError{}, false + } + + var errorCandidate *parsedOpsError + for _, frame := range strings.Split(normalized, "\n\n") { + frame = strings.TrimSpace(frame) + if frame == "" { + continue + } + var eventType string + dataLines := make([]string, 0, 1) + for _, line := range strings.Split(frame, "\n") { + switch { + case strings.HasPrefix(line, "event:"): + eventType = strings.TrimSpace(strings.TrimPrefix(line, "event:")) + case strings.HasPrefix(line, "data:"): + dataLines = append(dataLines, strings.TrimSpace(strings.TrimPrefix(line, "data:"))) + } + } + if eventType != "response.failed" && eventType != "error" && len(dataLines) == 0 { + continue + } + + payload := strings.Join(dataLines, "\n") + var event map[string]any + if err := json.Unmarshal([]byte(payload), &event); err == nil { + if eventType == "" { + eventType, _ = event["type"].(string) + } + } + if eventType != "response.failed" && eventType != "error" { + continue + } + + parsed := parsedOpsError{ErrorType: "upstream_error", StreamFailure: true} + if eventType == "error" { + parsed.ErrorType = "api_error" + } + if errObj := opsSSEErrorObject(event); errObj != nil { + parsed.ErrorType, _ = errObj["type"].(string) + parsed.Message, _ = errObj["message"].(string) + parsed.Code, _ = errObj["code"].(string) + if parsed.ErrorType == "" { + parsed.ErrorType = inferResponsesFailedOpsErrorType(parsed.Code) + } + if parsed.ErrorType == "" { + if eventType == "error" { + parsed.ErrorType = "api_error" + } else { + parsed.ErrorType = "upstream_error" + } + } + } + if strings.TrimSpace(parsed.Message) == "" && payload != "" { + parsed.Message = truncateString(payload, 1024) + } + if eventType == "response.failed" { + return parsed, true + } + candidate := parsed + errorCandidate = &candidate + } + if errorCandidate != nil { + return *errorCandidate, true + } + return parsedOpsError{}, false +} + +func opsSSEErrorObject(event map[string]any) map[string]any { + if event == nil { + return nil + } + if errObj, ok := event["error"].(map[string]any); ok { + return errObj + } + if response, ok := event["response"].(map[string]any); ok { + if errObj, ok := response["error"].(map[string]any); ok { + return errObj + } + } + return nil +} + +func inferResponsesFailedOpsErrorType(code string) string { + switch strings.TrimSpace(code) { + case "rate_limit_exceeded": + return "rate_limit_error" + case "permission_denied", "cyber_policy", "content_policy": + return "permission_error" + case "invalid_request", "context_length_exceeded": + return "invalid_request_error" + case "server_is_overloaded": + return "overloaded_error" + case "authentication_failed": + return "authentication_error" + default: + return "" + } +} + +func inferStreamFailureStatus(c *gin.Context, parsed parsedOpsError) int { + switch strings.TrimSpace(parsed.Code) { + case "rate_limit_exceeded": + return http.StatusTooManyRequests + case "permission_denied", "cyber_policy", "content_policy": + return http.StatusForbidden + case "invalid_request", "context_length_exceeded": + return http.StatusBadRequest + case "server_is_overloaded": + return http.StatusServiceUnavailable + case "authentication_failed": + return http.StatusUnauthorized + } + + switch strings.TrimSpace(parsed.ErrorType) { + case "rate_limit_error": + return http.StatusTooManyRequests + case "permission_error", "forbidden_error": + return http.StatusForbidden + case "authentication_error": + return http.StatusUnauthorized + case "invalid_request_error": + return http.StatusBadRequest + case "overloaded_error", "service_unavailable_error": + return http.StatusServiceUnavailable + } + + if c != nil { + if v, ok := c.Get(service.OpsUpstreamStatusCodeKey); ok { + switch code := v.(type) { + case int: + if code >= 400 { + return code + } + case int64: + if code >= 400 { + return int(code) + } + } + } + } + return http.StatusBadGateway +} + // getOpsAPIKey 返回用于 Ops 错误日志的 API Key:优先取已鉴权写入的正式 key; // 鉴权早退(分组停用/删除、Key 停用/过期/额度、用户停用、IP 限制等)时, // 正式 key 尚未写入,回退到 middleware 写入的 ops fallback key @@ -1472,6 +1463,7 @@ func isKnownOpsErrorType(t string) bool { "subscription_error", "upstream_error", "overloaded_error", + "service_unavailable_error", "api_error", "not_found_error", "forbidden_error": diff --git a/backend/internal/handler/ops_error_logger_test.go b/backend/internal/handler/ops_error_logger_test.go index 3b778cdcf9..4151093ae2 100644 --- a/backend/internal/handler/ops_error_logger_test.go +++ b/backend/internal/handler/ops_error_logger_test.go @@ -9,6 +9,7 @@ import ( "testing" "unicode/utf8" + "github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey" middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" @@ -279,6 +280,168 @@ func TestOpsErrorLoggerMiddleware_HardSkipsIngressRejection(t *testing.T) { require.Zero(t, OpsErrorLogEnqueuedTotal(), "ingress rejection must not enter the error queue") } +func TestOpsErrorLoggerMiddleware_SkipsRecoveredUpstreamErrorOnSuccessfulRequest(t *testing.T) { + setupOpsErrorLogTestQueue(t, 2) + gin.SetMode(gin.TestMode) + + ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.Use(OpsErrorLoggerMiddleware(ops)) + router.POST("/v1/responses", func(c *gin.Context) { + c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{ + UpstreamStatusCode: http.StatusTooManyRequests, + Message: "earlier attempt was rate limited", + }}) + c.JSON(http.StatusOK, gin.H{"status": "completed"}) + }) + + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil)) + + require.Equal(t, http.StatusOK, recorder.Code) + require.Zero(t, OpsErrorLogQueueLength()) +} + +func TestOpsErrorLoggerMiddleware_CapturesSplitResponsesFailedSSE(t *testing.T) { + setupOpsErrorLogTestQueue(t, 2) + gin.SetMode(gin.TestMode) + + ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.Use(OpsErrorLoggerMiddleware(ops)) + router.POST("/v1/responses", func(c *gin.Context) { + setOpsRequestContext(c, "gpt-5.5", true) + c.Status(http.StatusOK) + _, _ = c.Writer.Write([]byte("event: response.")) + _, _ = c.Writer.Write([]byte("failed\n")) + _, _ = c.Writer.Write([]byte(`data: {"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded","message":"Too many pending requests"}}}`)) + _, _ = c.Writer.Write([]byte("\n\n")) + }) + + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil)) + + require.Equal(t, http.StatusOK, recorder.Code) + require.Equal(t, int64(1), OpsErrorLogQueueLength()) + job := <-opsErrorLogQueue + require.Equal(t, http.StatusTooManyRequests, job.entry.StatusCode) + require.Equal(t, "rate_limit_error", job.entry.ErrorType) + require.Contains(t, job.entry.ErrorMessage, "Too many pending requests") +} + +func TestOpsCaptureWriter_CapturesSplitDataOnlyTerminalMarkers(t *testing.T) { + tests := []struct { + name string + prefix string + suffix string + wantType string + wantCode string + wantError string + }{ + { + name: "response failed with space", + prefix: `data: {"type":"response.`, + suffix: `failed","response":{"error":{"code":"server_is_overloaded","message":"busy"}}}`, + wantType: "overloaded_error", + wantCode: "server_is_overloaded", + wantError: "busy", + }, + { + name: "response failed without space", + prefix: `data:{"type":"response.`, + suffix: `failed","error":{"code":"rate_limit_exceeded","message":"slow down"}}`, + wantType: "rate_limit_error", + wantCode: "rate_limit_exceeded", + wantError: "slow down", + }, + { + name: "error with space", + prefix: `data: {"type":"er`, + suffix: `ror","error":{"type":"invalid_request_error","code":"invalid_request","message":"bad input"}}`, + wantType: "invalid_request_error", + wantCode: "invalid_request", + wantError: "bad input", + }, + { + name: "error without space", + prefix: `data:{"type":"er`, + suffix: `ror","error":{"type":"authentication_error","code":"authentication_failed","message":"sign in"}}`, + wantType: "authentication_error", + wantCode: "authentication_failed", + wantError: "sign in", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + writer := &opsCaptureWriter{limit: opsCaptureWriterLimit} + writer.captureResponseChunk([]byte(tt.prefix), http.StatusOK) + require.Empty(t, writer.buf.Bytes(), "partial marker must remain in the bounded probe") + writer.captureResponseChunk([]byte(tt.suffix+"\n\n"), http.StatusOK) + + parsed := parseOpsErrorResponse(writer.buf.Bytes()) + require.True(t, writer.sseCapturing) + require.True(t, parsed.StreamFailure) + require.Equal(t, tt.wantType, parsed.ErrorType) + require.Equal(t, tt.wantCode, parsed.Code) + require.Equal(t, tt.wantError, parsed.Message) + require.LessOrEqual(t, len(writer.probe), opsTerminalSSEProbeSize) + }) + } +} + +func TestOpsErrorLoggerMiddleware_StreamFailureUsesTerminalErrorOverAttemptContext(t *testing.T) { + setupOpsErrorLogTestQueue(t, 2) + gin.SetMode(gin.TestMode) + + ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.Use(OpsErrorLoggerMiddleware(ops)) + router.POST("/v1/responses", func(c *gin.Context) { + service.SetOpsUpstreamError(c, http.StatusBadGateway, "Upstream transport error", "earlier attempt failed") + c.Status(http.StatusOK) + _, _ = c.Writer.WriteString("event: er") + _, _ = c.Writer.WriteString("ror\n") + _, _ = c.Writer.WriteString(`data: {"type":"error","error":{"type":"invalid_request_error","code":"context_length_exceeded","message":"input exceeds the context window"}}`) + _, _ = c.Writer.WriteString("\n\n") + }) + + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil)) + + require.Equal(t, http.StatusOK, recorder.Code) + require.Equal(t, int64(1), OpsErrorLogQueueLength()) + job := <-opsErrorLogQueue + require.Equal(t, http.StatusBadRequest, job.entry.StatusCode) + require.Equal(t, "invalid_request_error", job.entry.ErrorType) + require.NotNil(t, job.entry.UpstreamStatusCode) + require.Equal(t, http.StatusBadRequest, *job.entry.UpstreamStatusCode) + require.NotNil(t, job.entry.UpstreamErrorMessage) + require.Equal(t, "input exceeds the context window", *job.entry.UpstreamErrorMessage) +} + +func TestOpsErrorLoggerMiddleware_PrefersContextRequestID(t *testing.T) { + setupOpsErrorLogTestQueue(t, 2) + gin.SetMode(gin.TestMode) + + ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.Use(OpsErrorLoggerMiddleware(ops)) + router.POST("/v1/responses", func(c *gin.Context) { + c.Header("X-Request-Id", "response-header-id") + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": "bad input"}}) + }) + + recorder := httptest.NewRecorder() + request := httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + request = request.WithContext(context.WithValue(request.Context(), ctxkey.RequestID, "context-request-id")) + router.ServeHTTP(recorder, request) + + require.Equal(t, int64(1), OpsErrorLogQueueLength()) + job := <-opsErrorLogQueue + require.Equal(t, "context-request-id", job.entry.RequestID) +} + func TestNormalizeOpsPersistentUserAgentBoundsAndPreservesUTF8(t *testing.T) { value := strings.Repeat("a", opsErrorLogMaxUserAgentBytes-1) + "你" + strings.Repeat("b", 32) got := normalizeOpsPersistentUserAgent(" " + value + " ") diff --git a/backend/internal/handler/request_body_read_log.go b/backend/internal/handler/request_body_read_log.go new file mode 100644 index 0000000000..5c6d73d934 --- /dev/null +++ b/backend/internal/handler/request_body_read_log.go @@ -0,0 +1,79 @@ +package handler + +import ( + "context" + "errors" + "io" + "net" + "net/http" + "strings" + "syscall" + + "go.uber.org/zap" +) + +// logRequestBodyReadFailure records a bounded, payload-free reason for a body +// read failure. Clients continue to receive the stable generic error message; +// operators get enough information to distinguish compression failures from a +// disconnected/truncated upload without logging request content. +func logRequestBodyReadFailure(reqLog *zap.Logger, req *http.Request, err error) { + if reqLog == nil || err == nil { + return + } + + contentLength := int64(-1) + contentEncoding := "identity" + if req != nil { + contentLength = req.ContentLength + contentEncoding = requestContentEncodingCategory(req.Header.Get("Content-Encoding")) + } + + reqLog.Warn("read request body failed", + zap.String("error_kind", requestBodyReadErrorKind(err)), + zap.String("content_encoding", contentEncoding), + zap.Int64("content_length", contentLength), + ) +} + +func requestContentEncodingCategory(value string) string { + switch strings.ToLower(strings.TrimSpace(value)) { + case "", "identity": + return "identity" + case "gzip", "x-gzip": + return "gzip" + case "zstd": + return "zstd" + case "deflate": + return "deflate" + default: + return "other" + } +} + +func requestBodyReadErrorKind(err error) string { + if err == nil { + return "none" + } + var maxErr *http.MaxBytesError + if errors.As(err, &maxErr) { + return "max_bytes" + } + lower := strings.ToLower(err.Error()) + if strings.Contains(lower, "decode content-encoding") { + if strings.Contains(lower, "unsupported content-encoding") { + return "unsupported_content_encoding" + } + return "decode_content_encoding" + } + if errors.Is(err, context.Canceled) || errors.Is(err, syscall.ECONNRESET) || errors.Is(err, syscall.EPIPE) { + return "client_disconnect" + } + if errors.Is(err, io.ErrUnexpectedEOF) { + return "truncated_body" + } + var netErr net.Error + if errors.As(err, &netErr) { + return "transport" + } + return "io_read" +} diff --git a/backend/internal/handler/request_body_read_log_test.go b/backend/internal/handler/request_body_read_log_test.go new file mode 100644 index 0000000000..8a8d21f100 --- /dev/null +++ b/backend/internal/handler/request_body_read_log_test.go @@ -0,0 +1,40 @@ +//go:build unit + +package handler + +import ( + "errors" + "fmt" + "io" + "net/http" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestLogRequestBodyReadFailureClassifiesWithoutPayload(t *testing.T) { + log, logs := newObservedLogger(t) + req, err := http.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader("secret-payload-marker")) + require.NoError(t, err) + req.Header.Set("Content-Encoding", "gzip") + req.ContentLength = 21 + + logRequestBodyReadFailure(log, req, errors.New(`decode Content-Encoding "gzip": unexpected EOF secret-payload-marker`)) + + entries := logs.All() + require.Len(t, entries, 1) + require.Equal(t, "read request body failed", entries[0].Message) + fields := entries[0].ContextMap() + require.Equal(t, "decode_content_encoding", fields["error_kind"]) + require.Equal(t, "gzip", fields["content_encoding"]) + require.EqualValues(t, 21, fields["content_length"]) + require.NotContains(t, entries[0].Message+fmt.Sprint(fields), "secret-payload-marker") +} + +func TestRequestBodyReadErrorKind(t *testing.T) { + require.Equal(t, "unsupported_content_encoding", requestBodyReadErrorKind(errors.New(`decode Content-Encoding "br": unsupported Content-Encoding`))) + require.Equal(t, "truncated_body", requestBodyReadErrorKind(io.ErrUnexpectedEOF)) + require.Equal(t, "max_bytes", requestBodyReadErrorKind(&http.MaxBytesError{Limit: 10})) + require.Equal(t, "other", requestContentEncodingCategory("private-payload-marker")) +} diff --git a/backend/internal/service/image_generation_intent.go b/backend/internal/service/image_generation_intent.go index 30da6ab6ba..927fd08765 100644 --- a/backend/internal/service/image_generation_intent.go +++ b/backend/internal/service/image_generation_intent.go @@ -293,11 +293,16 @@ func openAIRequestBodyImageGenerationToolNeedsNormalization(body []byte) bool { if openAIJSONString(item.Get("type")) != "image_generation" { return true } - // 只有旧字段需要迁移时才进入 map 修改,纯计费读取保持 raw 路径。 + // 只有旧字段或明确的模型不兼容字段需要修正时才进入 map 修改。 if item.Get("format").Exists() || item.Get("compression").Exists() { needsNormalization = true return false } + imageModel := strings.ToLower(strings.TrimSpace(item.Get("model").String())) + if strings.HasPrefix(imageModel, "gpt-image-2") && item.Get("input_fidelity").Exists() { + needsNormalization = true + return false + } return true }) return needsNormalization diff --git a/backend/internal/service/openai_alpha_search.go b/backend/internal/service/openai_alpha_search.go index 7e7a13f3b5..c49ecead94 100644 --- a/backend/internal/service/openai_alpha_search.go +++ b/backend/internal/service/openai_alpha_search.go @@ -449,6 +449,7 @@ var openAIAlphaSearchUnsupportedBodyFields = [...]string{ // alpha/search 会对这些字段返回 Unknown parameter(例如 prompt_cache_key)。 "prompt_cache_key", "prompt_cache_retention", + "store", } func sanitizeOpenAIAlphaSearchBody(body []byte) ([]byte, error) { diff --git a/backend/internal/service/openai_alpha_search_test.go b/backend/internal/service/openai_alpha_search_test.go index c4eded7d59..ab3f9145bf 100644 --- a/backend/internal/service/openai_alpha_search_test.go +++ b/backend/internal/service/openai_alpha_search_test.go @@ -482,6 +482,16 @@ func TestShouldApplyOpenAIAlphaSearchAccountErrorSideEffects(t *testing.T) { require.True(t, shouldApplyOpenAIAlphaSearchAccountErrorSideEffects(http.StatusTooManyRequests)) } +func TestSanitizeOpenAIAlphaSearchBody_RemovesResponsesOnlyFields(t *testing.T) { + body := []byte(`{"id":"search-session","store":false,"prompt_cache_key":"cache","commands":{"search_query":[{"q":"news"}]}}`) + + normalized, err := sanitizeOpenAIAlphaSearchBody(body) + require.NoError(t, err) + require.False(t, gjson.GetBytes(normalized, "store").Exists()) + require.False(t, gjson.GetBytes(normalized, "prompt_cache_key").Exists()) + require.Equal(t, "news", gjson.GetBytes(normalized, "commands.search_query.0.q").String()) +} + func TestIsOpenAIAlphaSearchEndpointUnsupported(t *testing.T) { apiKey := &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey} oauth := &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth} diff --git a/backend/internal/service/openai_codex_function_call_id_test.go b/backend/internal/service/openai_codex_function_call_id_test.go index edfcac1d71..1d798eec5e 100644 --- a/backend/internal/service/openai_codex_function_call_id_test.go +++ b/backend/internal/service/openai_codex_function_call_id_test.go @@ -66,6 +66,91 @@ func TestFilterCodexInput_KeepsFcID_WhenPreservingReferences(t *testing.T) { require.Equal(t, "fc_validID123", fc["id"], "valid fc* id must be preserved") } +func TestFilterCodexInput_PreservesNativeCustomAndToolSearchIDs(t *testing.T) { + input := []any{ + map[string]any{"type": "custom_tool_call", "id": "ctc_valid", "call_id": "call_custom", "name": "apply_patch"}, + map[string]any{"type": "tool_search_call", "id": "tsc_valid", "call_id": "call_search"}, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true}) + + require.Equal(t, "ctc_valid", filtered[0].(map[string]any)["id"]) + require.Equal(t, "tsc_valid", filtered[1].(map[string]any)["id"]) +} + +func TestFilterCodexInput_StripsWrongCustomAndToolSearchIDs(t *testing.T) { + input := []any{ + map[string]any{"type": "custom_tool_call", "id": "fc_wrong", "call_id": "call_custom", "name": "apply_patch"}, + map[string]any{"type": "tool_search_call", "id": "fc_wrong", "call_id": "call_search"}, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true}) + + require.NotContains(t, filtered[0].(map[string]any), "id") + require.NotContains(t, filtered[1].(map[string]any), "id") +} + +func TestFilterCodexInput_MapsItemReferencesToNativeToolCallPair(t *testing.T) { + input := []any{ + map[string]any{"type": "custom_tool_call", "id": "fc_custom", "call_id": "call_custom", "name": "apply_patch"}, + map[string]any{"type": "custom_tool_call_output", "call_id": "fc_custom", "output": "done"}, + map[string]any{"type": "item_reference", "id": "call_custom"}, + map[string]any{"type": "tool_search_call", "id": "fc_search", "call_id": "call_search"}, + map[string]any{"type": "tool_search_output", "call_id": "fc_search", "output": "result"}, + map[string]any{"type": "item_reference", "id": "call_search"}, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true}) + + require.Equal(t, "ctc_custom", filtered[0].(map[string]any)["call_id"]) + require.Equal(t, "ctc_custom", filtered[1].(map[string]any)["call_id"]) + require.Equal(t, "ctc_custom", filtered[2].(map[string]any)["id"]) + require.Equal(t, "tsc_search", filtered[3].(map[string]any)["call_id"]) + require.Equal(t, "tsc_search", filtered[4].(map[string]any)["call_id"]) + require.Equal(t, "tsc_search", filtered[5].(map[string]any)["id"]) +} + +func TestFilterCodexInput_PreservesAmbiguousItemReference(t *testing.T) { + input := []any{ + map[string]any{"type": "custom_tool_call", "call_id": "call_shared", "name": "apply_patch"}, + map[string]any{"type": "tool_search_call", "call_id": "call_shared"}, + map[string]any{"type": "item_reference", "id": "call_shared"}, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true}) + + require.Equal(t, "ctc_shared", filtered[0].(map[string]any)["call_id"]) + require.Equal(t, "tsc_shared", filtered[1].(map[string]any)["call_id"]) + require.Equal(t, "call_shared", filtered[2].(map[string]any)["id"]) +} + +func TestFilterCodexInput_PreservesNativeItemIDReferenceIndependentlyFromCallID(t *testing.T) { + input := []any{ + map[string]any{"type": "custom_tool_call", "id": "ctc_item", "call_id": "call_custom", "name": "apply_patch"}, + map[string]any{"type": "item_reference", "id": "ctc_item"}, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true}) + + require.Equal(t, "ctc_item", filtered[0].(map[string]any)["id"]) + require.Equal(t, "ctc_custom", filtered[0].(map[string]any)["call_id"]) + require.Equal(t, "ctc_item", filtered[1].(map[string]any)["id"]) +} + +func TestFilterCodexInput_ExistingItemIDWinsOverLegacyCallIDMapping(t *testing.T) { + input := []any{ + map[string]any{"type": "custom_tool_call", "call_id": "call_shared", "name": "apply_patch"}, + map[string]any{"type": "function_call_output", "id": "call_shared", "call_id": "call_other", "output": "done"}, + map[string]any{"type": "item_reference", "id": "call_shared"}, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true}) + + require.Equal(t, "ctc_shared", filtered[0].(map[string]any)["call_id"]) + require.Equal(t, "call_shared", filtered[1].(map[string]any)["id"]) + require.Equal(t, "call_shared", filtered[2].(map[string]any)["id"]) +} + // TestFilterCodexInput_StripsItemIDFromAllToolCallInputTypes verifies that // item_* ids are stripped from all call-input types (not output types). func TestFilterCodexInput_StripsItemIDFromAllToolCallInputTypes(t *testing.T) { diff --git a/backend/internal/service/openai_codex_tool_names.go b/backend/internal/service/openai_codex_tool_names.go new file mode 100644 index 0000000000..243757c113 --- /dev/null +++ b/backend/internal/service/openai_codex_tool_names.go @@ -0,0 +1,244 @@ +package service + +import ( + "encoding/json" + "fmt" + "strings" + + "github.com/gin-gonic/gin" +) + +const ( + codexReservedPythonToolName = "python" + codexPythonToolAlias = "python__sub2api" + codexToolNameReverseKey = "openai_codex_tool_name_reverse" +) + +type codexToolNameField struct { + object map[string]any + key string + name string +} + +// aliasOpenAIOAuthReservedToolNames avoids names reserved by the ChatGPT +// Codex backend. It validates every declaration/reference before mutating so +// collisions cannot leave a partially rewritten request. +func aliasOpenAIOAuthReservedToolNames(reqBody map[string]any) (map[string]string, bool, error) { + if reqBody == nil { + return nil, false, nil + } + + fields := collectOpenAIResponsesToolNameFields(reqBody) + owners := make(map[string]string) + reverse := make(map[string]string) + for _, field := range fields { + normalized := aliasOpenAIOAuthReservedToolName(field.name) + original := field.name + if normalized != field.name { + original = strings.TrimSpace(field.name) + } + if previous, exists := owners[normalized]; exists && previous != original { + return nil, false, fmt.Errorf("tool names %q and %q both normalize to %q", previous, original, normalized) + } + owners[normalized] = original + if normalized != field.name { + reverse[normalized] = original + } + } + if len(reverse) == 0 { + return nil, false, nil + } + for _, field := range fields { + if aliased := aliasOpenAIOAuthReservedToolName(field.name); aliased != field.name { + field.object[field.key] = aliased + } + } + return reverse, true, nil +} + +func aliasOpenAIOAuthReservedToolName(name string) string { + trimmed := strings.TrimSpace(name) + if strings.EqualFold(trimmed, codexReservedPythonToolName) { + return codexPythonToolAlias + } + return name +} + +func collectOpenAIResponsesToolNameFields(reqBody map[string]any) []codexToolNameField { + fields := make([]codexToolNameField, 0, 8) + appendName := func(object map[string]any, key string) { + if object == nil { + return + } + name, ok := object[key].(string) + if !ok || strings.TrimSpace(name) == "" { + return + } + fields = append(fields, codexToolNameField{object: object, key: key, name: name}) + } + var collectTools func(any) + collectTools = func(rawTools any) { + tools, ok := rawTools.([]any) + if !ok { + return + } + for _, raw := range tools { + tool, ok := raw.(map[string]any) + if !ok { + continue + } + toolType := strings.ToLower(strings.TrimSpace(firstNonEmptyString(tool["type"]))) + if toolType != "namespace" { + appendName(tool, "name") + } + if function, ok := tool["function"].(map[string]any); ok { + appendName(function, "name") + } + collectTools(tool["tools"]) + } + } + collectTools(reqBody["tools"]) + collectTools(reqBody["functions"]) + if choice, ok := reqBody["tool_choice"].(map[string]any); ok { + if !strings.EqualFold(strings.TrimSpace(firstNonEmptyString(choice["type"])), "namespace") { + appendName(choice, "name") + } + if function, ok := choice["function"].(map[string]any); ok { + appendName(function, "name") + } + } + if input, ok := reqBody["input"].([]any); ok { + for _, raw := range input { + item, ok := raw.(map[string]any) + if !ok { + continue + } + typ := strings.ToLower(strings.TrimSpace(firstNonEmptyString(item["type"]))) + if typ == "additional_tools" { + collectTools(item["tools"]) + } + if strings.HasSuffix(typ, "_call") || typ == "tool_call" { + appendName(item, "name") + if function, ok := item["function"].(map[string]any); ok { + appendName(function, "name") + } + } + } + } + return fields +} + +func aliasOpenAIOAuthReservedToolNamesBody(body []byte) ([]byte, map[string]string, bool, error) { + if len(body) == 0 || !containsASCIIFold(body, []byte(codexReservedPythonToolName)) { + return body, nil, false, nil + } + var reqBody map[string]any + if err := json.Unmarshal(body, &reqBody); err != nil { + return body, nil, false, fmt.Errorf("decode OAuth reserved tool names: %w", err) + } + reverse, changed, err := aliasOpenAIOAuthReservedToolNames(reqBody) + if err != nil || !changed { + return body, reverse, false, err + } + normalized, err := json.Marshal(reqBody) + if err != nil { + return body, nil, false, fmt.Errorf("encode OAuth reserved tool names: %w", err) + } + return normalized, reverse, true, nil +} + +func containsASCIIFold(haystack, needle []byte) bool { + if len(needle) == 0 || len(haystack) < len(needle) { + return false + } + for i := 0; i <= len(haystack)-len(needle); i++ { + matched := true + for j := range needle { + a, b := haystack[i+j], needle[j] + if a >= 'A' && a <= 'Z' { + a += 'a' - 'A' + } + if b >= 'A' && b <= 'Z' { + b += 'a' - 'A' + } + if a != b { + matched = false + break + } + } + if matched { + return true + } + } + return false +} + +func setCodexToolNameReverse(c *gin.Context, reverse map[string]string) { + if c == nil { + return + } + copyMap := make(map[string]string, len(reverse)) + for aliased, original := range reverse { + copyMap[aliased] = original + } + c.Set(codexToolNameReverseKey, copyMap) +} + +func codexToolNameReverseFromContext(c *gin.Context) map[string]string { + if c == nil { + return nil + } + raw, ok := c.Get(codexToolNameReverseKey) + if !ok { + return nil + } + reverse, _ := raw.(map[string]string) + return reverse +} + +func restoreCodexToolNamesInJSON(data []byte, reverse map[string]string) []byte { + if len(data) == 0 || len(reverse) == 0 || !json.Valid(data) { + return data + } + var decoded any + if err := json.Unmarshal(data, &decoded); err != nil { + return data + } + if !restoreCodexToolNameFields(decoded, reverse) { + return data + } + restored, err := json.Marshal(decoded) + if err != nil { + return data + } + return restored +} + +func restoreCodexToolNamesFromContext(c *gin.Context, data []byte) []byte { + return restoreCodexToolNamesInJSON(data, codexToolNameReverseFromContext(c)) +} + +func restoreCodexToolNameFields(value any, reverse map[string]string) bool { + changed := false + switch typed := value.(type) { + case map[string]any: + if name, ok := typed["name"].(string); ok { + if original, exists := reverse[name]; exists { + typed["name"] = original + changed = true + } + } + for _, child := range typed { + if restoreCodexToolNameFields(child, reverse) { + changed = true + } + } + case []any: + for _, child := range typed { + if restoreCodexToolNameFields(child, reverse) { + changed = true + } + } + } + return changed +} diff --git a/backend/internal/service/openai_codex_tool_names_test.go b/backend/internal/service/openai_codex_tool_names_test.go new file mode 100644 index 0000000000..25f8248b48 --- /dev/null +++ b/backend/internal/service/openai_codex_tool_names_test.go @@ -0,0 +1,95 @@ +package service + +import ( + "encoding/json" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestAliasOpenAIOAuthReservedToolNames_RewritesDeclarationsAndReferences(t *testing.T) { + reqBody := map[string]any{ + "tools": []any{ + map[string]any{"type": "function", "name": "python"}, + map[string]any{"type": "namespace", "name": "code", "tools": []any{ + map[string]any{"type": "function", "name": "shell"}, + }}, + }, + "tool_choice": map[string]any{"type": "function", "name": "python"}, + "input": []any{ + map[string]any{"type": "function_call", "name": "python", "call_id": "fc_1"}, + map[string]any{"type": "additional_tools", "tools": []any{ + map[string]any{"type": "function", "function": map[string]any{"name": "python"}}, + }}, + }, + } + + reverse, changed, err := aliasOpenAIOAuthReservedToolNames(reqBody) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "python", reverse[codexPythonToolAlias]) + require.Equal(t, codexPythonToolAlias, reqBody["tools"].([]any)[0].(map[string]any)["name"]) + require.Equal(t, codexPythonToolAlias, reqBody["tool_choice"].(map[string]any)["name"]) + require.Equal(t, codexPythonToolAlias, reqBody["input"].([]any)[0].(map[string]any)["name"]) + nested := reqBody["input"].([]any)[1].(map[string]any)["tools"].([]any)[0].(map[string]any)["function"].(map[string]any) + require.Equal(t, codexPythonToolAlias, nested["name"]) +} + +func TestAliasOpenAIOAuthReservedToolNames_CollisionDoesNotMutate(t *testing.T) { + reqBody := map[string]any{"tools": []any{ + map[string]any{"type": "function", "name": "python"}, + map[string]any{"type": "function", "name": codexPythonToolAlias}, + }} + before, err := json.Marshal(reqBody) + require.NoError(t, err) + + reverse, changed, err := aliasOpenAIOAuthReservedToolNames(reqBody) + require.ErrorContains(t, err, `both normalize to "python__sub2api"`) + require.False(t, changed) + require.Nil(t, reverse) + after, marshalErr := json.Marshal(reqBody) + require.NoError(t, marshalErr) + require.JSONEq(t, string(before), string(after)) +} + +func TestApplyCodexOAuthTransform_ReservedPythonNameIsOAuthOnly(t *testing.T) { + reqBody := map[string]any{ + "model": "gpt-5.5", + "tools": []any{map[string]any{"type": "function", "name": "PYTHON"}}, + } + + result := applyCodexOAuthTransform(reqBody, true, false) + require.NoError(t, result.Error) + require.Equal(t, "PYTHON", result.ToolNameReverse[codexPythonToolAlias]) + require.Equal(t, codexPythonToolAlias, reqBody["tools"].([]any)[0].(map[string]any)["name"]) + + apiKeyBody := []byte(`{"type":"response.create","tools":[{"type":"function","name":"python"}]}`) + normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(apiKeyBody, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}) + require.NoError(t, err) + require.False(t, changed) + require.JSONEq(t, string(apiKeyBody), string(normalized)) +} + +func TestRestoreCodexToolNamesFromContext_HTTPAndWSPayloadShapes(t *testing.T) { + c, _ := gin.CreateTestContext(nil) + setCodexToolNameReverse(c, map[string]string{codexPythonToolAlias: "python"}) + + streamEvent := restoreCodexToolNamesFromContext(c, []byte( + `{"type":"response.output_item.done","item":{"type":"function_call","name":"python__sub2api"},"note":"python__sub2api"}`, + )) + require.Equal(t, "python", gjson.GetBytes(streamEvent, "item.name").String()) + require.Equal(t, "python__sub2api", gjson.GetBytes(streamEvent, "note").String()) + + nonStreaming := restoreCodexToolNamesFromContext(c, []byte( + `{"id":"resp_1","output":[{"type":"function_call","name":"python__sub2api"}]}`, + )) + require.Equal(t, "python", gjson.GetBytes(nonStreaming, "output.0.name").String()) + + setCodexToolNameReverse(c, nil) + require.JSONEq(t, + `{"type":"response.output_item.added","item":{"name":"python__sub2api"}}`, + string(restoreCodexToolNamesFromContext(c, []byte(`{"type":"response.output_item.added","item":{"name":"python__sub2api"}}`))), + ) +} diff --git a/backend/internal/service/openai_codex_transform.go b/backend/internal/service/openai_codex_transform.go index c89fd038d0..3b9cafd946 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -76,6 +76,8 @@ type codexTransformResult struct { Modified bool NormalizedModel string PromptCacheKey string + ToolNameReverse map[string]string + Error error } type codexOAuthTransformOptions struct { @@ -92,26 +94,45 @@ const ( ) func normalizeCodexCallID(id string) string { + return normalizeCodexCallIDForItemType("function_call", id) +} + +func normalizeCodexCallIDForItemType(itemType, id string) string { + prefix := openAIResponsesToolCallIDPrefix(itemType) + "_" candidate := id switch { case id == "": return "" - case strings.HasPrefix(id, "fc"): + case strings.HasPrefix(id, strings.TrimSuffix(prefix, "_")): case strings.HasPrefix(id, "call_"): - candidate = codexCallIDPrefix + strings.TrimPrefix(id, "call_") + candidate = prefix + strings.TrimPrefix(id, "call_") default: - candidate = codexCallIDPrefix + id + candidate = prefix + trimOpenAIResponsesKnownCallIDPrefix(id) } if len(candidate) <= codexCallIDMaxLength { return candidate } - return compactCodexCallID(candidate) + return compactCodexCallIDForItemType(itemType, candidate) } func compactCodexCallID(id string) string { + return compactCodexCallIDForItemType("function_call", id) +} + +func compactCodexCallIDForItemType(itemType, id string) string { + prefix := openAIResponsesToolCallIDPrefix(itemType) + "_" digest := sha256.Sum256([]byte("sub2api:codex-call-id:v1:" + id)) encoded := hex.EncodeToString(digest[:]) - return codexCallIDPrefix + encoded[:codexCallIDMaxLength-len(codexCallIDPrefix)] + return prefix + encoded[:codexCallIDMaxLength-len(prefix)] +} + +func trimOpenAIResponsesKnownCallIDPrefix(id string) string { + for _, prefix := range []string{"fc_", "ctc_", "tsc_"} { + if strings.HasPrefix(id, prefix) { + return strings.TrimPrefix(id, prefix) + } + } + return id } const codexImageGenerationFunctionToolName = "image_gen.imagegen" @@ -124,11 +145,14 @@ const ( ) var openAIChatGPTInternalUnsupportedFields = []string{ + "chat_template_kwargs", "user", "metadata", "prompt_cache_retention", "safety_identifier", "stream_options", + "truncation", + "stop_sequences", } var openAICodexOAuthUnsupportedFields = append([]string{ @@ -149,6 +173,18 @@ func applyCodexOAuthTransform(reqBody map[string]any, isCodexCLI bool, isCompact func applyCodexOAuthTransformWithOptions(reqBody map[string]any, opts codexOAuthTransformOptions) codexTransformResult { result := codexTransformResult{} + toolNameReverse, toolNamesChanged, err := aliasOpenAIOAuthReservedToolNames(reqBody) + if err != nil { + result.Error = err + return result + } + result.ToolNameReverse = toolNameReverse + if toolNamesChanged { + result.Modified = true + } + if normalizeOpenAIOAuthResponsesCompatibilityFields(reqBody) { + result.Modified = true + } // 工具续链需求会影响存储策略与 input 过滤逻辑。 needsToolContinuation := NeedsToolContinuation(reqBody) @@ -884,6 +920,124 @@ func normalizeOpenAIResponsesImageGenerationTools(reqBody map[string]any) bool { delete(toolMap, "compression") modified = true } + imageModel := strings.ToLower(strings.TrimSpace(firstNonEmptyString(toolMap["model"]))) + if strings.HasPrefix(imageModel, "gpt-image-2") { + if _, ok := toolMap["input_fidelity"]; ok { + delete(toolMap, "input_fidelity") + modified = true + } + } + } + return modified +} + +func normalizeOpenAIResponseFormatSchemas(reqBody map[string]any) bool { + if reqBody == nil { + return false + } + modified := false + normalizeFormat := func(format map[string]any) { + if format == nil || strings.TrimSpace(firstNonEmptyString(format["type"])) != "json_schema" { + return + } + if schema, ok := format["schema"].(map[string]any); ok && normalizeOpenAIResponseJSONSchema(schema) { + modified = true + } + if jsonSchema, ok := format["json_schema"].(map[string]any); ok { + if schema, ok := jsonSchema["schema"].(map[string]any); ok && normalizeOpenAIResponseJSONSchema(schema) { + modified = true + } + } + } + if text, ok := reqBody["text"].(map[string]any); ok { + if format, ok := text["format"].(map[string]any); ok { + normalizeFormat(format) + } + } + if responseFormat, ok := reqBody["response_format"].(map[string]any); ok { + normalizeFormat(responseFormat) + } + return modified +} + +func normalizeOpenAIResponseJSONSchema(schema map[string]any) bool { + if schema == nil { + return false + } + modified := false + for _, key := range []string{"uniqueItems", "minProperties"} { + if _, exists := schema[key]; exists { + delete(schema, key) + modified = true + } + } + if rawType, exists := schema["type"]; !exists || rawType == nil { + switch { + case schema["properties"] != nil: + schema["type"] = "object" + modified = true + case schema["items"] != nil: + schema["type"] = "array" + modified = true + } + } + if properties, ok := schema["properties"].(map[string]any); ok { + for _, raw := range properties { + if child, ok := raw.(map[string]any); ok && normalizeOpenAIResponseJSONSchema(child) { + modified = true + } + } + } + switch items := schema["items"].(type) { + case map[string]any: + if normalizeOpenAIResponseJSONSchema(items) { + modified = true + } + case []any: + for _, raw := range items { + if child, ok := raw.(map[string]any); ok && normalizeOpenAIResponseJSONSchema(child) { + modified = true + } + } + } + for _, key := range []string{ + "additionalProperties", + "additionalItems", + "contains", + "not", + "if", + "then", + "else", + "propertyNames", + "unevaluatedProperties", + "unevaluatedItems", + } { + if child, ok := schema[key].(map[string]any); ok && normalizeOpenAIResponseJSONSchema(child) { + modified = true + } + } + for _, key := range []string{"anyOf", "oneOf", "allOf", "prefixItems"} { + children, _ := schema[key].([]any) + for _, raw := range children { + if child, ok := raw.(map[string]any); ok && normalizeOpenAIResponseJSONSchema(child) { + modified = true + } + } + } + for _, key := range []string{"$defs", "definitions", "patternProperties", "dependentSchemas"} { + children, _ := schema[key].(map[string]any) + for _, raw := range children { + if child, ok := raw.(map[string]any); ok && normalizeOpenAIResponseJSONSchema(child) { + modified = true + } + } + } + if dependencies, ok := schema["dependencies"].(map[string]any); ok { + for _, raw := range dependencies { + if child, ok := raw.(map[string]any); ok && normalizeOpenAIResponseJSONSchema(child) { + modified = true + } + } } return modified } @@ -1357,8 +1511,62 @@ func filterCodexInput(input []any, preserveReferences bool) []any { }) } +func normalizeCodexFilterCallID(itemType, id string, preserve bool) string { + if preserve && len(id) <= codexCallIDMaxLength { + return id + } + return normalizeCodexCallIDForItemType(itemType, id) +} + +func codexItemReferenceIDMappings(input []any, preserveCallIDs bool) map[string]string { + mappings := make(map[string]string) + ambiguous := make(map[string]struct{}) + for _, rawItem := range input { + item, ok := rawItem.(map[string]any) + if !ok { + continue + } + itemType := strings.TrimSpace(firstNonEmptyString(item["type"])) + if !isCodexToolCallItemType(itemType) { + continue + } + rawCallID := strings.TrimSpace(firstNonEmptyString(item["call_id"])) + if rawCallID == "" { + continue + } + normalized := normalizeCodexFilterCallID(itemType, rawCallID, preserveCallIDs) + if existing, exists := mappings[rawCallID]; exists && existing != normalized { + delete(mappings, rawCallID) + ambiguous[rawCallID] = struct{}{} + continue + } + if _, conflict := ambiguous[rawCallID]; !conflict { + mappings[rawCallID] = normalized + } + } + return mappings +} + +func codexInputItemIDs(input []any) map[string]struct{} { + itemIDs := make(map[string]struct{}) + for _, rawItem := range input { + item, ok := rawItem.(map[string]any) + if !ok || strings.TrimSpace(firstNonEmptyString(item["type"])) == "item_reference" { + continue + } + itemType := strings.TrimSpace(firstNonEmptyString(item["type"])) + id := strings.TrimSpace(firstNonEmptyString(item["id"])) + if id != "" && !shouldStripOpenAIResponsesInputItemID(itemType, id) { + itemIDs[id] = struct{}{} + } + } + return itemIDs +} + func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []any { filtered := make([]any, 0, len(input)) + referenceIDMappings := codexItemReferenceIDMappings(input, opts.PreserveCallIDs) + inputItemIDs := codexInputItemIDs(input) for _, item := range input { m, ok := item.(map[string]any) if !ok { @@ -1408,18 +1616,7 @@ func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []an // 仅修正真正的 tool/function call 标识,避免误改普通 message/reasoning id; // 若 item_reference 指向 legacy call_* 标识,则仅修正该引用本身。 fixCallIDPrefix := func(id string) string { - if opts.PreserveCallIDs { - // preserve 模式尽量原样透传客户端 id 以维持 tool_use/tool_result - // 配对,但上游对 call_id 有 64 字符硬上限,超长原样透传必然被 - // 400 拒绝("Invalid 'input[N].call_id': string too long")。 - // 超长时退回确定性压缩:同一逻辑 id 在 function_call 与 - // function_call_output 两侧结果一致,配对不受影响。 - if len(id) <= codexCallIDMaxLength { - return id - } - return compactCodexCallID(id) - } - return normalizeCodexCallID(id) + return normalizeCodexFilterCallID(typ, id, opts.PreserveCallIDs) } if typ == "item_reference" { @@ -1430,8 +1627,12 @@ func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []an for key, value := range m { newItem[key] = value } - if id, ok := newItem["id"].(string); ok && strings.HasPrefix(id, "call_") { - newItem["id"] = fixCallIDPrefix(id) + if id, ok := newItem["id"].(string); ok && strings.HasPrefix(strings.TrimSpace(id), "call_") { + trimmedID := strings.TrimSpace(id) + _, referencesExistingItem := inputItemIDs[trimmedID] + if normalizedID, mapped := referenceIDMappings[trimmedID]; mapped && !referencesExistingItem { + newItem["id"] = normalizedID + } } filtered = append(filtered, newItem) continue diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go index 9489ab60df..adeffadbfa 100644 --- a/backend/internal/service/openai_codex_transform_test.go +++ b/backend/internal/service/openai_codex_transform_test.go @@ -233,7 +233,7 @@ func TestApplyCodexOAuthTransform_CompactsOverlongCallIDsWhenPreserveRequested(t require.Len(t, compacted, codexCallIDMaxLength) require.True(t, strings.HasPrefix(compacted, codexCallIDPrefix)) require.Equal(t, compacted, output["call_id"], "两侧压缩结果必须一致以保持配对") - require.Equal(t, compactCodexCallID(callID), compacted, "压缩必须是确定性的") + require.Equal(t, normalizeCodexCallID(callID), compacted, "压缩必须是确定性的") } func TestApplyCodexOAuthTransform_ToolSearchOutputPreservesCallID(t *testing.T) { @@ -253,7 +253,7 @@ func TestApplyCodexOAuthTransform_ToolSearchOutputPreservesCallID(t *testing.T) first, ok := input[0].(map[string]any) require.True(t, ok) require.Equal(t, "tool_search_output", first["type"]) - require.Equal(t, "fc_1", first["call_id"]) + require.Equal(t, "tsc_1", first["call_id"]) } func TestApplyCodexOAuthTransform_CustomAndMCPToolOutputsPreserveCallID(t *testing.T) { @@ -273,13 +273,81 @@ func TestApplyCodexOAuthTransform_CustomAndMCPToolOutputsPreserveCallID(t *testi first, ok := input[0].(map[string]any) require.True(t, ok) - require.Equal(t, "fc_custom", first["call_id"]) + require.Equal(t, "ctc_custom", first["call_id"]) second, ok := input[1].(map[string]any) require.True(t, ok) require.Equal(t, "fc_mcp", second["call_id"]) } +func TestApplyCodexOAuthTransform_NormalizesNativeToolCallPairsByType(t *testing.T) { + reqBody := map[string]any{ + "model": "gpt-5.6-sol", + "input": []any{ + map[string]any{"type": "custom_tool_call", "id": "fc_custom", "call_id": "call_custom", "name": "apply_patch"}, + map[string]any{"type": "custom_tool_call_output", "call_id": "fc_custom", "output": "done"}, + map[string]any{"type": "tool_search_call", "id": "fc_search", "call_id": "call_search"}, + map[string]any{"type": "tool_search_output", "call_id": "fc_search", "output": "result"}, + }, + } + + applyCodexOAuthTransform(reqBody, false, false) + + input := reqBody["input"].([]any) + custom := input[0].(map[string]any) + customOutput := input[1].(map[string]any) + search := input[2].(map[string]any) + searchOutput := input[3].(map[string]any) + require.NotContains(t, custom, "id", "the invalid replay item id must be removed, not fabricated") + require.Equal(t, "ctc_custom", custom["call_id"]) + require.Equal(t, custom["call_id"], customOutput["call_id"]) + require.NotContains(t, search, "id", "the invalid replay item id must be removed, not fabricated") + require.Equal(t, "tsc_search", search["call_id"]) + require.Equal(t, search["call_id"], searchOutput["call_id"]) +} + +func TestApplyCodexOAuthTransform_PreservesNativeCallIDsWhenRequested(t *testing.T) { + reqBody := map[string]any{ + "model": "gpt-5.6-sol", + "input": []any{ + map[string]any{"type": "custom_tool_call", "id": "ctc_custom", "call_id": "call_custom", "name": "apply_patch"}, + map[string]any{"type": "custom_tool_call_output", "call_id": "call_custom", "output": "done"}, + map[string]any{"type": "tool_search_call", "id": "tsc_search", "call_id": "call_search"}, + map[string]any{"type": "tool_search_output", "call_id": "call_search", "output": "result"}, + }, + } + + applyCodexOAuthTransformWithOptions(reqBody, codexOAuthTransformOptions{PreserveToolCallIDs: true}) + + input := reqBody["input"].([]any) + require.Equal(t, "ctc_custom", input[0].(map[string]any)["id"]) + require.Equal(t, "call_custom", input[0].(map[string]any)["call_id"]) + require.Equal(t, "call_custom", input[1].(map[string]any)["call_id"]) + require.Equal(t, "tsc_search", input[2].(map[string]any)["id"]) + require.Equal(t, "call_search", input[2].(map[string]any)["call_id"]) + require.Equal(t, "call_search", input[3].(map[string]any)["call_id"]) +} + +func TestApplyCodexOAuthTransform_BoundsEquivalentNativeToolCallIDsWithPairing(t *testing.T) { + suffix := strings.Repeat("x", 70) + reqBody := map[string]any{ + "model": "gpt-5.6-sol", + "input": []any{ + map[string]any{"type": "custom_tool_call", "call_id": "call_" + suffix, "name": "apply_patch"}, + map[string]any{"type": "custom_tool_call_output", "call_id": "fc_" + suffix, "output": "done"}, + }, + } + + applyCodexOAuthTransformWithOptions(reqBody, codexOAuthTransformOptions{PreserveToolCallIDs: true}) + + input := reqBody["input"].([]any) + callID := input[0].(map[string]any)["call_id"].(string) + outputCallID := input[1].(map[string]any)["call_id"].(string) + require.LessOrEqual(t, len(callID), codexCallIDMaxLength) + require.True(t, strings.HasPrefix(callID, "ctc_")) + require.Equal(t, callID, outputCallID) +} + func TestApplyCodexOAuthTransform_ImageAndWebSearchCallsDoNotGainCallID(t *testing.T) { reqBody := map[string]any{ "model": "gpt-5.2", @@ -468,7 +536,7 @@ func TestApplyCodexOAuthTransform_PreservesFunctionCallInputName(t *testing.T) { item, ok := input[0].(map[string]any) require.True(t, ok) require.Equal(t, "shell", item["name"]) - require.Equal(t, "fc_1", item["call_id"]) + require.Equal(t, "ctc_1", item["call_id"]) } func TestApplyCodexOAuthTransform_PreservesMCPToolCallIDAndName(t *testing.T) { @@ -1651,11 +1719,14 @@ func TestApplyCodexOAuthTransform_StripsPromptCacheRetention(t *testing.T) { func TestApplyCodexOAuthTransform_StripsChatGPTInternalUnsupportedFields(t *testing.T) { reqBody := map[string]any{ "model": "gpt-5.4", + "chat_template_kwargs": map[string]any{"enable_thinking": true}, "user": "user_123", "metadata": map[string]any{"trace_id": "abc"}, "prompt_cache_retention": "24h", "safety_identifier": "sid", "stream_options": map[string]any{"include_usage": true}, + "truncation": "auto", + "stop_sequences": []any{"END"}, "input": []any{ map[string]any{"role": "user", "content": "hi"}, }, @@ -1669,6 +1740,40 @@ func TestApplyCodexOAuthTransform_StripsChatGPTInternalUnsupportedFields(t *test } } +func TestApplyCodexOAuthTransform_NormalizesPromptAndCommands(t *testing.T) { + reqBody := map[string]any{ + "model": "gpt-5.5", + "prompt": "hello", + "commands": []any{"unsupported"}, + } + + result := applyCodexOAuthTransform(reqBody, true, false) + require.True(t, result.Modified) + require.Equal(t, []any{ + map[string]any{"type": "message", "role": "user", "content": "hello"}, + }, reqBody["input"]) + require.NotContains(t, reqBody, "prompt") + require.NotContains(t, reqBody, "commands") +} + +func TestNormalizeOpenAIResponsesImageGenerationTools_StripsGPTImage2InputFidelity(t *testing.T) { + reqBody := map[string]any{"tools": []any{ + map[string]any{"type": "image_generation", "model": "gpt-image-2-codex", "input_fidelity": "high"}, + map[string]any{"type": "image_generation", "model": "gpt-image-1.5", "input_fidelity": "high"}, + }} + + require.True(t, normalizeOpenAIResponsesImageGenerationTools(reqBody)) + tools := reqBody["tools"].([]any) + require.NotContains(t, tools[0].(map[string]any), "input_fidelity") + require.Equal(t, "high", tools[1].(map[string]any)["input_fidelity"]) +} + +func TestOpenAIRequestBodyImageGenerationToolNeedsNormalization_GPTImage2InputFidelity(t *testing.T) { + body := []byte(`{"tools":[{"type":"image_generation","model":"gpt-image-2-codex","input_fidelity":"high"}]}`) + + require.True(t, openAIRequestBodyImageGenerationToolNeedsNormalization(body)) +} + func TestApplyCodexOAuthTransform_ExtractsSystemMessages(t *testing.T) { reqBody := map[string]any{ "model": "gpt-5.1", diff --git a/backend/internal/service/openai_codex_turn_state_test.go b/backend/internal/service/openai_codex_turn_state_test.go index a58e166299..141f089f52 100644 --- a/backend/internal/service/openai_codex_turn_state_test.go +++ b/backend/internal/service/openai_codex_turn_state_test.go @@ -7,6 +7,8 @@ import ( "testing" "time" + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/util/responseheaders" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" ) @@ -249,6 +251,19 @@ func TestWriteOpenAIPassthroughResponseHeaders_RelaysAndClearsTurnState(t *testi require.Empty(t, dst.Get("X-Codex-Turn-State")) } +func TestWriteOpenAIPassthroughResponseHeaders_RelaysReasoningIncluded(t *testing.T) { + dst := http.Header{} + src := http.Header{} + src.Set("X-Reasoning-Included", "1") + + writeOpenAIPassthroughResponseHeaders( + dst, + src, + responseheaders.CompileHeaderFilter(config.ResponseHeaderConfig{}), + ) + require.Equal(t, "1", dst.Get("X-Reasoning-Included")) +} + func TestEnsureOpenAIRemoteCompactionV2BetaFeature(t *testing.T) { t.Run("absent_sets_feature", func(t *testing.T) { h := http.Header{} diff --git a/backend/internal/service/openai_compact_body_signal.go b/backend/internal/service/openai_compact_body_signal.go index 83266e6c00..09d962529d 100644 --- a/backend/internal/service/openai_compact_body_signal.go +++ b/backend/internal/service/openai_compact_body_signal.go @@ -1,6 +1,7 @@ package service import ( + "encoding/json" "net/http" "strings" @@ -22,6 +23,47 @@ func MarkOpenAINativeCompactionV2(c *gin.Context) { } } +// NormalizeCompactionTriggerInputOrder keeps a single compaction trigger as +// the final Responses input item, as required by the upstream v2 wire format. +func NormalizeCompactionTriggerInputOrder(body []byte) ([]byte, bool, error) { + if len(body) == 0 { + return body, false, nil + } + var payload map[string]any + if err := json.Unmarshal(body, &payload); err != nil { + return body, false, err + } + input, ok := payload["input"].([]any) + if !ok || len(input) == 0 { + return body, false, nil + } + triggerCount := 0 + normalized := make([]any, 0, len(input)) + for _, raw := range input { + item, itemOK := raw.(map[string]any) + if itemOK && item["type"] == "compaction_trigger" { + triggerCount++ + continue + } + normalized = append(normalized, raw) + } + if triggerCount == 0 { + return body, false, nil + } + if triggerCount == 1 { + if last, ok := input[len(input)-1].(map[string]any); ok && last["type"] == "compaction_trigger" { + return body, false, nil + } + } + normalized = append(normalized, map[string]any{"type": "compaction_trigger"}) + payload["input"] = normalized + encoded, err := json.Marshal(payload) + if err != nil { + return body, false, err + } + return encoded, true, nil +} + func isOpenAINativeCompactionV2(c *gin.Context) bool { if c == nil { return false diff --git a/backend/internal/service/openai_compact_body_signal_test.go b/backend/internal/service/openai_compact_body_signal_test.go index 8e05f4f98b..3b0a21295e 100644 --- a/backend/internal/service/openai_compact_body_signal_test.go +++ b/backend/internal/service/openai_compact_body_signal_test.go @@ -6,6 +6,7 @@ import ( "testing" "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" ) func TestHasCompactionTriggerInInput_DetectsCompactSignal(t *testing.T) { @@ -54,3 +55,22 @@ func TestHasCompactionTriggerInInput_CompactTriggerOnly(t *testing.T) { body := []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"}]}`) require.True(t, HasCompactionTriggerInInput(body)) } + +func TestNormalizeCompactionTriggerInputOrder_MovesAndCollapsesTriggers(t *testing.T) { + body := []byte(`{"model":"gpt-5.5","input":[{"type":"compaction_trigger"},{"type":"message","role":"user","content":"tail"},{"type":"compaction_trigger"}]}`) + normalized, changed, err := NormalizeCompactionTriggerInputOrder(body) + require.NoError(t, err) + require.True(t, changed) + items := gjson.GetBytes(normalized, "input").Array() + require.Len(t, items, 2) + require.Equal(t, "message", items[0].Get("type").String()) + require.Equal(t, "compaction_trigger", items[1].Get("type").String()) +} + +func TestNormalizeCompactionTriggerInputOrder_AlreadyFinalPreservesBytes(t *testing.T) { + body := []byte(`{"input":[{"type":"message"},{"type":"compaction_trigger"}]}`) + normalized, changed, err := NormalizeCompactionTriggerInputOrder(body) + require.NoError(t, err) + require.False(t, changed) + require.Equal(t, string(body), string(normalized)) +} diff --git a/backend/internal/service/openai_gateway_apikey_item_id_test.go b/backend/internal/service/openai_gateway_apikey_item_id_test.go index 80865e2891..9c4990f545 100644 --- a/backend/internal/service/openai_gateway_apikey_item_id_test.go +++ b/backend/internal/service/openai_gateway_apikey_item_id_test.go @@ -39,6 +39,10 @@ func TestOpenAIGatewayService_APIKeyPassthrough_StripsInvalidInputItemIDs(t *tes {"type":"function_call","id":"item_bad_call","call_id":"call_123","name":"exec_command","arguments":"{}"}, {"type":"message","id":"msg_valid","role":"user","content":[{"type":"input_text","text":"continue"}]}, {"type":"function_call","id":"fc_valid","call_id":"call_456","name":"apply_patch","arguments":"{}"}, + {"type":"custom_tool_call","id":"fc_wrong_custom","call_id":"call_custom_1","name":"apply_patch","input":"patch"}, + {"type":"custom_tool_call","id":"ctc_valid","call_id":"call_custom_2","name":"apply_patch","input":"patch"}, + {"type":"tool_search_call","id":"fc_wrong_search","call_id":"call_search_1","arguments":{"query":"docs"}}, + {"type":"tool_search_call","id":"tsc_valid","call_id":"call_search_2","arguments":{"query":"docs"}}, {"type":"function_call_output","id":"item_output","call_id":"call_123","output":"done"}, {"type":"web_search_call","id":"item_unconstrained"} ] @@ -58,9 +62,54 @@ func TestOpenAIGatewayService_APIKeyPassthrough_StripsInvalidInputItemIDs(t *tes require.Equal(t, "{}", gjson.GetBytes(forwarded, "input.1.arguments").String()) require.Equal(t, "msg_valid", gjson.GetBytes(forwarded, "input.2.id").String()) require.Equal(t, "fc_valid", gjson.GetBytes(forwarded, "input.3.id").String()) - require.Equal(t, "item_output", gjson.GetBytes(forwarded, "input.4.id").String()) - require.Equal(t, "call_123", gjson.GetBytes(forwarded, "input.4.call_id").String()) - require.Equal(t, "item_unconstrained", gjson.GetBytes(forwarded, "input.5.id").String()) + require.False(t, gjson.GetBytes(forwarded, "input.4.id").Exists()) + require.Equal(t, "ctc_valid", gjson.GetBytes(forwarded, "input.5.id").String()) + require.False(t, gjson.GetBytes(forwarded, "input.6.id").Exists()) + require.Equal(t, "tsc_valid", gjson.GetBytes(forwarded, "input.7.id").String()) + require.Equal(t, "item_output", gjson.GetBytes(forwarded, "input.8.id").String()) + require.Equal(t, "call_123", gjson.GetBytes(forwarded, "input.8.call_id").String()) + require.Equal(t, "item_unconstrained", gjson.GetBytes(forwarded, "input.9.id").String()) +} + +func TestOpenAIGatewayService_OAuthPassthrough_SanitizesNativeToolItemIDs(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstreamSSE := "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_test\",\"model\":\"gpt-5.6-sol\",\"output\":[],\"usage\":{\"input_tokens\":1,\"output_tokens\":1,\"total_tokens\":2}}}\n\ndata: [DONE]\n\n" + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader(upstreamSSE)), + }} + svc := newOpenAIImageGenerationControlTestService(upstream) + c, _ := newOpenAIImageGenerationControlTestContext(true, "codex_cli_rs/0.144.1") + account := newOpenAIImageGenerationControlTestAccount() + account.Type = AccountTypeOAuth + account.Credentials = map[string]any{ + "access_token": "oauth-token", + "chatgpt_account_id": "chatgpt-account", + } + account.Extra = map[string]any{"openai_passthrough": true} + + body := []byte(`{ + "model":"gpt-5.6-sol", + "stream":true, + "instructions":"test", + "input":[ + {"type":"custom_tool_call","id":"fc_wrong_custom","call_id":"call_custom_1","name":"apply_patch","input":"patch"}, + {"type":"custom_tool_call","id":"ctc_valid","call_id":"call_custom_2","name":"apply_patch","input":"patch"}, + {"type":"tool_search_call","id":"fc_wrong_search","call_id":"call_search_1","arguments":{"query":"docs"}}, + {"type":"tool_search_call","id":"tsc_valid","call_id":"call_search_2","arguments":{"query":"docs"}} + ] + }`) + + result, err := svc.Forward(context.Background(), c, account, body) + require.NoError(t, err) + require.NotNil(t, result) + require.NotNil(t, upstream.lastReq) + require.False(t, gjson.GetBytes(upstream.lastBody, "input.0.id").Exists()) + require.Equal(t, "ctc_valid", gjson.GetBytes(upstream.lastBody, "input.1.id").String()) + require.False(t, gjson.GetBytes(upstream.lastBody, "input.2.id").Exists()) + require.Equal(t, "tsc_valid", gjson.GetBytes(upstream.lastBody, "input.3.id").String()) } // TestOpenAIGatewayService_APIKeyPassthrough_StripsInvalidReasoningItemIDs @@ -119,7 +168,12 @@ func TestShouldStripOpenAIResponsesInputItemID_Reasoning(t *testing.T) { {"message msg id", "message", "msg_abc", false}, {"message item id", "message", "item_x", true}, {"function_call fc id", "function_call", "fc_abc", false}, + {"function_call ctc id", "function_call", "ctc_abc", true}, {"function_call item id", "function_call", "item_x", true}, + {"custom tool ctc id", "custom_tool_call", "ctc_abc", false}, + {"custom tool fc id", "custom_tool_call", "fc_abc", true}, + {"tool search tsc id", "tool_search_call", "tsc_abc", false}, + {"tool search fc id", "tool_search_call", "fc_abc", true}, {"unconstrained type", "web_search_call", "ws_001", false}, } for _, tc := range cases { diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 8308c4c0c3..3cd4a99bd9 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -60,6 +60,7 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( defaultMappedModel string, ) (*OpenAIForwardResult, error) { beginUpstreamResponseModelObservation(c) + setCodexToolNameReverse(c, nil) restrictionResult := s.detectCodexClientRestriction(c, account, body) logCodexCLIOnlyDetection(ctx, c, account, getAPIKeyIDFromContext(c), restrictionResult, body) @@ -212,6 +213,11 @@ func (s *OpenAIGatewayService) ForwardAsChatCompletions( SkipDefaultInstructions: !isResponsesShape, OmitPromotedSystemMessagesFromInput: !isResponsesShape && !isJSONObjectFormat, }) + if codexResult.Error != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": codexResult.Error.Error()}}) + return nil, codexResult.Error + } + setCodexToolNameReverse(c, codexResult.ToolNameReverse) if !isResponsesShape { ensureCodexOAuthInstructionsField(reqBody) } @@ -421,7 +427,7 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse( ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") - finalResponse, usage, acc, err := s.readOpenAICompatBufferedTerminal(resp, "openai chat_completions buffered", requestID) + finalResponse, usage, acc, err := s.readOpenAICompatBufferedTerminal(resp, c, "openai chat_completions buffered", requestID) if err != nil { return nil, s.newOpenAICompatBufferedReadFailoverError(c, account, resp, requestID, err) } @@ -635,6 +641,7 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse( } processDataLine := func(payload string) bool { + payload = string(restoreCodexToolNamesFromContext(c, []byte(payload))) if firstChunk { firstChunk = false ms := int(time.Since(startTime).Milliseconds()) diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 47c02d866e..0562ab92a7 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -23,6 +23,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco clearGrokResponsesClientToolMapping(c) clearOpenAIResponsesClientToolMapping(c) clearOpenAIResponsesNamespaceNames(c) + setCodexToolNameReverse(c, nil) startTime := time.Now() // 固定渠道映射后的请求级 canonical body;账号 normalize/strip 不得改写跨 failover hint。 canonicalImageIntentBody := body @@ -58,6 +59,22 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if toolSchemaSanitized { body = sanitizedToolBody } + patternSanitizedBody, patternSanitized, patternErr := sanitizeOpenAIResponsesToolSchemaPatterns(body) + if patternErr != nil { + return nil, fmt.Errorf("sanitize OpenAI Responses tool schema patterns: %w", patternErr) + } + if patternSanitized { + body = patternSanitizedBody + } + if account.IsOpenAI() && account.IsOAuth() { + reasoningBody, reasoningChanged, reasoningErr := normalizeOpenAIResponsesReasoningMode(body) + if reasoningErr != nil { + return nil, fmt.Errorf("normalize OpenAI Responses reasoning.mode: %w", reasoningErr) + } + if reasoningChanged { + body = reasoningBody + } + } if account.IsOpenAIOAuth() && isOpenAIResponsesLiteHeader(c.GetHeader(responsesLiteHeader)) { liteBody, changed, liteErr := normalizeOpenAIResponsesLiteToolsPayload(body) if liteErr != nil { @@ -119,7 +136,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if shouldForwardOpenAIResponsesViaRawChatCompletions(account) { return s.forwardResponsesViaRawChatCompletions(ctx, c, account, body) } - if account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey { + if account.Platform == PlatformOpenAI && (account.Type == AccountTypeAPIKey || account.Type == AccountTypeOAuth) { sanitizedBody, changed, sanitizeErr := sanitizeOpenAIResponsesInputItemIDs(body) if sanitizeErr != nil { return nil, fmt.Errorf("sanitize OpenAI Responses input item IDs: %w", sanitizeErr) @@ -325,6 +342,17 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco markPatchSet("reasoning.effort", "none") logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized reasoning.effort: minimal -> none (account: %s)", account.Name) } + if strings.TrimSpace(gjson.GetBytes(body, "text.format.type").String()) == "json_schema" || + strings.TrimSpace(gjson.GetBytes(body, "response_format.type").String()) == "json_schema" { + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } + if normalizeOpenAIResponseFormatSchemas(decoded) { + markDecodedModified() + logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Normalized Responses JSON schema compatibility") + } + } imageIntent = imageIntent || IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, nil) || isOpenAIImageGenerationModel(upstreamModel) if imageIntent && !imageGenerationAllowed { @@ -415,6 +443,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } else { codexResult = applyCodexOAuthTransform(decoded, isCodexCLI, isCompactRequest) } + if codexResult.Error != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": gin.H{"type": "invalid_request_error", "message": codexResult.Error.Error()}}) + return nil, codexResult.Error + } + setCodexToolNameReverse(c, codexResult.ToolNameReverse) if codexResult.Modified { markDecodedModified() } @@ -538,6 +571,29 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } } + if account.Type == AccountTypeOAuth { + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } + if input, ok := decoded["input"].([]any); ok && sanitizeOpenAIResponsesOrphanToolOutputs( + decoded, + input, + strings.TrimSpace(firstNonEmptyString(decoded["previous_response_id"])) != "", + ) { + markDecodedModified() + } + } + if reqBody != nil || openAIResponsesInputMayNeedTruncation(body) { + decoded, decodeErr := ensureReqBody() + if decodeErr != nil { + return nil, decodeErr + } + if truncateOpenAIResponsesInputText(decoded) { + markDecodedModified() + } + } + if bodyModified { if requestView.HasPatches() { if patchedBody, patchErr := requestView.ApplyPatches(); patchErr == nil { diff --git a/backend/internal/service/openai_gateway_messages.go b/backend/internal/service/openai_gateway_messages.go index 4b0aa58edf..914db873f0 100644 --- a/backend/internal/service/openai_gateway_messages.go +++ b/backend/internal/service/openai_gateway_messages.go @@ -35,6 +35,7 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( defaultMappedModel string, ) (*OpenAIForwardResult, error) { beginUpstreamResponseModelObservation(c) + setCodexToolNameReverse(c, nil) // 入口分流(国产供应商 Anthropic 协议):上游为供应商原生 Anthropic 端点时, // /v1/messages 请求零转换直通(仅模型名映射 + 少量 body 清洗),完整保留 @@ -193,6 +194,11 @@ func (s *OpenAIGatewayService) ForwardAsAnthropic( SkipDefaultInstructions: true, PreserveToolCallIDs: true, }) + if codexResult.Error != nil { + writeAnthropicError(c, http.StatusBadRequest, "invalid_request_error", codexResult.Error.Error()) + return nil, codexResult.Error + } + setCodexToolNameReverse(c, codexResult.ToolNameReverse) forcedTemplateText := "" if s.cfg != nil { forcedTemplateText = s.cfg.Gateway.ForcedCodexInstructionsTemplate @@ -551,7 +557,7 @@ func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse( ) (*OpenAIForwardResult, error) { requestID := resp.Header.Get("x-request-id") - finalResponse, usage, acc, err := s.readOpenAICompatBufferedTerminal(resp, "openai messages buffered", requestID) + finalResponse, usage, acc, err := s.readOpenAICompatBufferedTerminal(resp, c, "openai messages buffered", requestID) if err != nil { var readErr *openAICompatBufferedReadError if errors.As(err, &readErr) && readErr != nil { @@ -690,6 +696,7 @@ func (e *openAICompatBufferedReadError) Unwrap() error { return e.cause } func (s *OpenAIGatewayService) readOpenAICompatBufferedTerminal( resp *http.Response, + c *gin.Context, logPrefix string, requestID string, ) (*apicompat.ResponsesResponse, OpenAIUsage, *apicompat.BufferedResponseAccumulator, error) { @@ -769,6 +776,7 @@ func (s *OpenAIGatewayService) readOpenAICompatBufferedTerminal( if !ok { if frame, ok := parser.Finish(); ok { payload := openAICompatPayloadWithEventType(frame.Data, frame.EventType) + payload = string(restoreCodexToolNamesFromContext(c, []byte(payload))) var event apicompat.ResponsesStreamEvent if err := json.Unmarshal([]byte(payload), &event); err == nil { acc.ProcessEvent(&event) @@ -807,6 +815,7 @@ func (s *OpenAIGatewayService) readOpenAICompatBufferedTerminal( continue } payload := openAICompatPayloadWithEventType(frame.Data, frame.EventType) + payload = string(restoreCodexToolNamesFromContext(c, []byte(payload))) var event apicompat.ResponsesStreamEvent if err := json.Unmarshal([]byte(payload), &event); err != nil { @@ -918,6 +927,7 @@ func (s *OpenAIGatewayService) handleAnthropicStreamingResponse( // processDataLine handles a single "data: ..." SSE line from upstream. processDataLine := func(payload string) bool { + payload = string(restoreCodexToolNamesFromContext(c, []byte(payload))) if firstChunk { firstChunk = false ms := int(time.Since(startTime).Milliseconds()) diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index c15a78ef48..98aadf1bfd 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -1572,6 +1572,7 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( if restoreErr != nil { return resultWithUsage(), fmt.Errorf("restore OpenAI passthrough namespace response: %w", restoreErr) } + restoredData = restoreCodexToolNamesFromContext(c, restoredData) if !bytes.Equal(restoredData, dataBytes) { dataBytes = restoredData trimmedData = strings.TrimSpace(string(restoredData)) @@ -1818,6 +1819,7 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough( if err != nil { return nil, fmt.Errorf("restore OpenAI passthrough namespace response: %w", err) } + body = restoreCodexToolNamesFromContext(c, body) if mapping, ok := openAIResponsesClientToolMapping(c); ok && json.Valid(body) { body, _, err = apicompat.RestoreResponsesClientToolPayload(body, mapping) if err != nil { @@ -1869,6 +1871,7 @@ func (s *OpenAIGatewayService) handlePassthroughSSEToJSON(resp *http.Response, c if restoreErr != nil { return nil, fmt.Errorf("restore OpenAI passthrough namespace response: %w", restoreErr) } + restoredBody = restoreCodexToolNamesFromContext(c, restoredBody) body = restoredBody } else { terminalType, terminalPayload, terminalOK := extractOpenAISSETerminalEvent(bodyText) diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index f2efddcf77..980788284a 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -729,6 +729,188 @@ func extractOpenAIRequestMetaFromBody(body []byte) (model string, stream bool, p return view.Model, view.Stream, view.PromptCacheKey } +func normalizeOpenAIOAuthResponsesCompatibilityFields(reqBody map[string]any) bool { + if reqBody == nil { + return false + } + changed := false + if prompt, exists := reqBody["prompt"]; exists { + if input, hasInput := reqBody["input"]; !hasInput || input == nil { + if prompt != nil { + reqBody["input"] = prompt + } + } + delete(reqBody, "prompt") + changed = true + } + if _, exists := reqBody["commands"]; exists { + delete(reqBody, "commands") + changed = true + } + return changed +} + +func normalizeOpenAIOAuthResponsesCompatibilityBody(body []byte) ([]byte, bool, error) { + if len(body) == 0 { + return body, false, nil + } + normalized := body + changed := false + prompt := gjson.GetBytes(normalized, "prompt") + if prompt.Exists() { + input := gjson.GetBytes(normalized, "input") + if prompt.Type != gjson.Null && (!input.Exists() || input.Type == gjson.Null) { + next, err := sjson.SetRawBytes(normalized, "input", []byte(prompt.Raw)) + if err != nil { + return body, false, fmt.Errorf("normalize oauth responses prompt: %w", err) + } + normalized = next + } + next, err := sjson.DeleteBytes(normalized, "prompt") + if err != nil { + return body, false, fmt.Errorf("normalize oauth responses delete prompt: %w", err) + } + normalized = next + changed = true + } + if gjson.GetBytes(normalized, "commands").Exists() { + next, err := sjson.DeleteBytes(normalized, "commands") + if err != nil { + return body, false, fmt.Errorf("normalize oauth responses delete commands: %w", err) + } + normalized = next + changed = true + } + return normalized, changed, nil +} + +func normalizeOpenAIResponsesReasoningMode(body []byte) ([]byte, bool, error) { + if len(body) == 0 { + return body, false, nil + } + mode := gjson.GetBytes(body, "reasoning.mode") + if !mode.Exists() || mode.Type != gjson.String { + return body, false, nil + } + updated := body + effort := gjson.GetBytes(body, "reasoning.effort") + if (!effort.Exists() || effort.Type == gjson.Null || strings.TrimSpace(effort.String()) == "") && + strings.EqualFold(strings.TrimSpace(mode.String()), "pro") { + var err error + updated, err = sjson.SetBytes(updated, "reasoning.effort", "max") + if err != nil { + return body, false, fmt.Errorf("set reasoning effort for mode=pro: %w", err) + } + } + updated, err := sjson.DeleteBytes(updated, "reasoning.mode") + if err != nil { + return body, false, fmt.Errorf("delete unsupported reasoning.mode: %w", err) + } + if reasoning := gjson.GetBytes(updated, "reasoning"); reasoning.Exists() && reasoning.IsObject() && len(reasoning.Map()) == 0 { + updated, err = sjson.DeleteBytes(updated, "reasoning") + if err != nil { + return body, false, fmt.Errorf("delete empty reasoning object: %w", err) + } + } + return updated, true, nil +} + +func normalizeOpenAIResponseFormatSchemasBody(body []byte) ([]byte, bool, error) { + if len(body) == 0 { + return body, false, nil + } + textFormat := strings.TrimSpace(gjson.GetBytes(body, "text.format.type").String()) + responseFormat := strings.TrimSpace(gjson.GetBytes(body, "response_format.type").String()) + if textFormat != "json_schema" && responseFormat != "json_schema" { + return body, false, nil + } + var reqBody map[string]any + decoder := json.NewDecoder(bytes.NewReader(body)) + decoder.UseNumber() + if err := decoder.Decode(&reqBody); err != nil { + return body, false, fmt.Errorf("normalize responses schema body: %w", err) + } + if !normalizeOpenAIResponseFormatSchemas(reqBody) { + return body, false, nil + } + normalized, err := json.Marshal(reqBody) + if err != nil { + return body, false, fmt.Errorf("serialize normalized responses schema body: %w", err) + } + return normalized, true, nil +} + +func normalizeOpenAIResponsesWebSocketCompatibilityBody(body []byte, account *Account) ([]byte, bool, error) { + normalized := body + changed := false + if sanitized, idsChanged, err := sanitizeOpenAIResponsesInputItemIDs(normalized); err != nil { + return body, false, fmt.Errorf("sanitize websocket Responses input item IDs: %w", err) + } else if idsChanged { + normalized = sanitized + changed = true + } + if account != nil && account.IsOpenAI() && account.IsOAuth() { + if reasoningBody, reasoningChanged, err := normalizeOpenAIResponsesReasoningMode(normalized); err != nil { + return body, false, err + } else if reasoningChanged { + normalized = reasoningBody + changed = true + } + } + if account != nil && account.IsOpenAIOAuth() { + oauthBody, oauthChanged, err := normalizeOpenAIOAuthResponsesCompatibilityBody(normalized) + if err != nil { + return body, false, err + } + normalized = oauthBody + changed = changed || oauthChanged + for _, field := range openAIChatGPTInternalUnsupportedFields { + if !gjson.GetBytes(normalized, field).Exists() { + continue + } + next, deleteErr := sjson.DeleteBytes(normalized, field) + if deleteErr != nil { + return body, false, fmt.Errorf("normalize websocket body delete %s: %w", field, deleteErr) + } + normalized = next + changed = true + } + } + if schemaBody, schemaChanged, err := normalizeOpenAIResponseFormatSchemasBody(normalized); err != nil { + return body, false, err + } else if schemaChanged { + normalized = schemaBody + changed = true + } + if openAIRequestBodyImageGenerationToolNeedsNormalization(normalized) { + var reqBody map[string]any + if err := json.Unmarshal(normalized, &reqBody); err != nil { + return body, false, fmt.Errorf("normalize websocket image tool body: %w", err) + } + if normalizeOpenAIResponsesImageGenerationTools(reqBody) { + next, err := json.Marshal(reqBody) + if err != nil { + return body, false, fmt.Errorf("serialize normalized websocket image tool body: %w", err) + } + normalized = next + changed = true + } + } + if toolBody, toolChanged, err := sanitizeOpenAIResponsesToolParameterTypes(normalized); err != nil { + return body, false, fmt.Errorf("normalize websocket tool parameter types: %w", err) + } else if toolChanged { + normalized = toolBody + changed = true + } + if patternBody, patternChanged, err := sanitizeOpenAIResponsesToolSchemaPatterns(normalized); err != nil { + return body, false, fmt.Errorf("normalize websocket tool schema patterns: %w", err) + } else if patternChanged { + normalized = patternBody + changed = true + } + return normalized, changed, nil +} + // normalizeOpenAIPassthroughOAuthBody 将透传 OAuth 请求体收敛为旧链路关键行为: // 1) 删除 ChatGPT internal API 不支持的顶层 Responses 参数 // 2) store=false 3) 非 compact 保持 stream=true;compact 强制 stream=false @@ -737,8 +919,16 @@ func normalizeOpenAIPassthroughOAuthBody(body []byte, compact bool) ([]byte, boo return body, false, nil } - normalized := body - changed := false + normalized, changed, err := normalizeOpenAIOAuthResponsesCompatibilityBody(body) + if err != nil { + return body, false, err + } + if reasoningBody, reasoningChanged, reasoningErr := normalizeOpenAIResponsesReasoningMode(normalized); reasoningErr != nil { + return body, false, reasoningErr + } else if reasoningChanged { + normalized = reasoningBody + changed = true + } for _, field := range openAIChatGPTInternalUnsupportedFields { if value := gjson.GetBytes(normalized, field); !value.Exists() { @@ -751,6 +941,12 @@ func normalizeOpenAIPassthroughOAuthBody(body []byte, compact bool) ([]byte, boo normalized = next changed = true } + if schemaBody, schemaChanged, schemaErr := normalizeOpenAIResponseFormatSchemasBody(normalized); schemaErr != nil { + return body, false, schemaErr + } else if schemaChanged { + normalized = schemaBody + changed = true + } if inputResult := gjson.GetBytes(normalized, "input"); inputResult.Exists() { switch { diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index 5577e90e4b..bc6043fdc9 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -566,6 +566,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. streamEarlyErr = fmt.Errorf("restore OpenAI namespace response: %w", restoreErr) return } + restoredData = restoreCodexToolNamesFromContext(c, restoredData) if !bytes.Equal(restoredData, dataBytes) { dataBytes = restoredData data = string(restoredData) @@ -1060,21 +1061,43 @@ func (s *OpenAIGatewayService) parseSSEUsage(data string, usage *OpenAIUsage) { } func (s *OpenAIGatewayService) parseSSEUsageBytes(data []byte, usage *OpenAIUsage) { - if usage == nil || len(data) == 0 || bytes.Equal(data, []byte("[DONE]")) { + if usage == nil || len(data) == 0 || bytes.Equal(bytes.TrimSpace(data), []byte("[DONE]")) { return } - // 选择性解析:仅在数据中包含终止事件标识时才进入字段提取。 - if len(data) < 72 { + parsedUsage, ok := extractOpenAIUsageFromJSONBytes(data) + if !ok { return } - eventType := gjson.GetBytes(data, "type").String() - if eventType != "response.completed" && eventType != "response.done" && eventType != "response.failed" && - eventType != "response.incomplete" && eventType != "response.cancelled" && eventType != "response.canceled" { - return - } - - if parsedUsage, ok := extractOpenAIUsageFromJSONBytes(data); ok { + if openAIStreamEventTypeIsTerminal(strings.TrimSpace(gjson.GetBytes(data, "type").String())) { *usage = parsedUsage + return + } + mergeOpenAIUsageNonZero(usage, parsedUsage) +} + +// Compatible Responses upstreams may report usage before the terminal event. +// Retain those non-zero fields as a fallback; terminal usage remains authoritative. +func mergeOpenAIUsageNonZero(dst *OpenAIUsage, src OpenAIUsage) { + if dst == nil { + return + } + if src.InputTokens > 0 { + dst.InputTokens = src.InputTokens + } + if src.ImageInputTokens > 0 { + dst.ImageInputTokens = src.ImageInputTokens + } + if src.OutputTokens > 0 { + dst.OutputTokens = src.OutputTokens + } + if src.CacheCreationInputTokens > 0 { + dst.CacheCreationInputTokens = src.CacheCreationInputTokens + } + if src.CacheReadInputTokens > 0 { + dst.CacheReadInputTokens = src.CacheReadInputTokens + } + if src.ImageOutputTokens > 0 { + dst.ImageOutputTokens = src.ImageOutputTokens } } @@ -1310,6 +1333,7 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r if err != nil { return nil, fmt.Errorf("restore OpenAI namespace response: %w", err) } + body = restoreCodexToolNamesFromContext(c, body) responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter) // Codex 协议要求 /responses/compact JSON 响应携带 x-codex-turn-state // (codex-api/src/endpoint/compact.rs 从响应头捕获),显式回传。 @@ -1391,6 +1415,7 @@ func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Conte if restoreErr != nil { return nil, fmt.Errorf("restore OpenAI namespace response: %w", restoreErr) } + restoredBody = restoreCodexToolNamesFromContext(c, restoredBody) body = restoredBody } else { terminalType, terminalPayload, terminalOK := extractOpenAISSETerminalEvent(bodyText) diff --git a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go index c8acab0358..850c5d0ce8 100644 --- a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go +++ b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go @@ -327,10 +327,29 @@ func TestIsOpenAIContextWindowError(t *testing.T) { "maximum context length exceeded", nil, )) + require.True(t, isOpenAIContextWindowError( + "", + []byte(`maximum context length exceeded`), + )) require.False(t, isOpenAIContextWindowError( "context canceled", nil, )) + require.False(t, isOpenAIContextWindowError( + "upstream unavailable", + []byte(`{"error":{"message":"upstream unavailable","code":"upstream_error"},"echo":"context_length_exceeded maximum context length"}`), + )) +} + +func TestOpenAITransientAndCapacityClassificationIgnoresEchoedJSON(t *testing.T) { + body := []byte(`{"error":{"message":"upstream unavailable","code":"upstream_error"},"echo":"server is overloaded; selected model is at capacity"}`) + + require.False(t, isOpenAITransientProcessingError(http.StatusBadRequest, "upstream unavailable", body)) + require.False(t, isOpenAIRequestScopedCapacityShed("upstream unavailable", body)) + + plainText := []byte(`server is overloaded; please retry later`) + require.True(t, isOpenAITransientProcessingError(http.StatusServiceUnavailable, "", plainText)) + require.True(t, isOpenAIRequestScopedCapacityShed("", plainText)) } func TestShouldFailoverOpenAIUpstreamResponseContextWindow502(t *testing.T) { @@ -339,6 +358,11 @@ func TestShouldFailoverOpenAIUpstreamResponseContextWindow502(t *testing.T) { require.False(t, svc.shouldFailoverOpenAIUpstreamResponse(http.StatusBadGateway, "", body)) require.True(t, svc.shouldFailoverOpenAIUpstreamResponse(http.StatusBadGateway, "temporary upstream outage", []byte(`{"error":{"message":"temporary upstream outage"}}`))) + require.True(t, svc.shouldFailoverOpenAIUpstreamResponse( + http.StatusBadGateway, + "temporary upstream outage", + []byte(`{"error":{"message":"temporary upstream outage"},"echo":"context_length_exceeded"}`), + )) } func TestOpenAIGatewayService_Forward_LogsInstructionsRequiredDetails(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 1d18d3d70c..ddaed2a8df 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -3391,10 +3391,10 @@ func TestParseSSEUsage_SelectiveParsing(t *testing.T) { svc := &OpenAIGatewayService{} usage := &OpenAIUsage{InputTokens: 9, OutputTokens: 8, CacheReadInputTokens: 7} - // 非 completed 事件,不应覆盖 usage + // 非终态事件中的显式 usage 作为兼容 fallback,非零字段会被合并。 svc.parseSSEUsage(`{"type":"response.in_progress","response":{"usage":{"input_tokens":1,"output_tokens":2}}}`, usage) - require.Equal(t, 9, usage.InputTokens) - require.Equal(t, 8, usage.OutputTokens) + require.Equal(t, 1, usage.InputTokens) + require.Equal(t, 2, usage.OutputTokens) require.Equal(t, 7, usage.CacheReadInputTokens) // completed 事件,应提取 usage @@ -3421,6 +3421,43 @@ func TestParseSSEUsage_SelectiveParsing(t *testing.T) { require.Equal(t, 6, usage.CacheReadInputTokens) } +func TestParseSSEUsage_NonTerminalUsageMergesNonZeroFields(t *testing.T) { + svc := &OpenAIGatewayService{} + usage := &OpenAIUsage{} + + svc.parseSSEUsage(`{"type":"response.in_progress","usage":{"input_tokens":17,"output_tokens":1,"input_tokens_details":{"cached_tokens":4}}}`, usage) + svc.parseSSEUsage(`{"type":"response.output_text.done","usage":{"input_tokens":0,"output_tokens":5,"input_tokens_details":{"cached_tokens":0,"cache_write_tokens":3}}}`, usage) + + require.Equal(t, 17, usage.InputTokens) + require.Equal(t, 5, usage.OutputTokens) + require.Equal(t, 4, usage.CacheReadInputTokens) + require.Equal(t, 3, usage.CacheCreationInputTokens) +} + +func TestParseSSEUsage_TerminalUsageReplacesFallback(t *testing.T) { + svc := &OpenAIGatewayService{} + usage := &OpenAIUsage{} + + svc.parseSSEUsage(`{"type":"response.output_text.done","usage":{"input_tokens":17,"output_tokens":5,"input_tokens_details":{"cached_tokens":4}}}`, usage) + svc.parseSSEUsage(`{"type":"response.completed","response":{"usage":{"input_tokens":19,"output_tokens":7}}}`, usage) + + require.Equal(t, 19, usage.InputTokens) + require.Equal(t, 7, usage.OutputTokens) + require.Zero(t, usage.CacheReadInputTokens) +} + +func TestParseSSEUsage_TerminalWithoutUsageKeepsFallback(t *testing.T) { + svc := &OpenAIGatewayService{} + usage := &OpenAIUsage{} + + svc.parseSSEUsage(`{"type":"response.in_progress","usage":{"input_tokens":17,"output_tokens":5}}`, usage) + svc.parseSSEUsage(`{"type":"response.completed","response":{"id":"resp_1"}}`, usage) + svc.parseSSEUsage(" [DONE]\n", usage) + + require.Equal(t, 17, usage.InputTokens) + require.Equal(t, 5, usage.OutputTokens) +} + func TestExtractOpenAIUsageFromJSONBytes_AcceptsResponseAndChatUsageShapes(t *testing.T) { usage, ok := extractOpenAIUsageFromJSONBytes([]byte(`{"id":"resp_1","usage":{"input_tokens":9,"output_tokens":5,"input_tokens_details":{"cached_tokens":2,"cache_write_tokens":4}}}`)) require.True(t, ok) diff --git a/backend/internal/service/openai_gateway_upstream_errors.go b/backend/internal/service/openai_gateway_upstream_errors.go index 32f04b2c6a..3fa2c3db75 100644 --- a/backend/internal/service/openai_gateway_upstream_errors.go +++ b/backend/internal/service/openai_gateway_upstream_errors.go @@ -135,7 +135,7 @@ func isOpenAITransientProcessingError(upstreamStatusCode int, upstreamMsg string if isOpenAICapacityShedMessage(upstreamMsg) || isOpenAICapacityShedMessage(gjson.GetBytes(upstreamBody, "error.message").String()) || isOpenAICapacityShedMessage(gjson.GetBytes(upstreamBody, "response.error.message").String()) || - isOpenAICapacityShedMessage(string(upstreamBody)) { + (!gjson.ValidBytes(upstreamBody) && isOpenAICapacityShedMessage(string(upstreamBody))) { return true } if upstreamStatusCode != http.StatusBadRequest && upstreamStatusCode != http.StatusServiceUnavailable { @@ -170,7 +170,14 @@ func isOpenAITransientProcessingError(upstreamStatusCode int, upstreamMsg string if match(gjson.GetBytes(upstreamBody, "error.message").String()) { return true } - return match(string(upstreamBody)) + if match(gjson.GetBytes(upstreamBody, "response.error.message").String()) || + match(gjson.GetBytes(upstreamBody, "message").String()) { + return true + } + // A valid JSON error may echo arbitrary request content. Only its explicit + // error fields are authoritative; scan the whole body only for non-JSON + // providers that return a plain-text error response. + return !gjson.ValidBytes(upstreamBody) && match(string(upstreamBody)) } func isOpenAICapacityShedMessage(text string) bool { @@ -183,7 +190,7 @@ func isOpenAICapacityShedMessage(text string) bool { func isOpenAIRequestScopedCapacityShed(upstreamMsg string, upstreamBody []byte) bool { return isOpenAIUpstreamCapacityShedEvent(upstreamBody) || isOpenAICapacityShedMessage(upstreamMsg) || - isOpenAICapacityShedMessage(string(upstreamBody)) + (!gjson.ValidBytes(upstreamBody) && isOpenAICapacityShedMessage(string(upstreamBody))) } func isOpenAIContextWindowError(upstreamMsg string, upstreamBody []byte) bool { @@ -228,7 +235,10 @@ func isOpenAIContextWindowError(upstreamMsg string, upstreamBody []byte) bool { return true } } - return match(string(upstreamBody)) + // Do not let echoed request content in a structured JSON error change the + // retry/client-status classification. Plain-text upstream errors remain + // supported by scanning the whole body only when it is not valid JSON. + return !gjson.ValidBytes(upstreamBody) && match(string(upstreamBody)) } func (s *OpenAIGatewayService) shouldFailoverUpstreamError(statusCode int) bool { diff --git a/backend/internal/service/openai_passthrough_normalization_test.go b/backend/internal/service/openai_passthrough_normalization_test.go index 15b803a7de..e7d06e4b8f 100644 --- a/backend/internal/service/openai_passthrough_normalization_test.go +++ b/backend/internal/service/openai_passthrough_normalization_test.go @@ -20,6 +20,194 @@ func TestNormalizeOpenAIPassthroughOAuthBody_RemovesUnsupportedUser(t *testing.T require.False(t, gjson.GetBytes(normalized, "store").Bool()) } +func TestNormalizeOpenAIPassthroughOAuthBody_NormalizesCompatibilityFields(t *testing.T) { + body := []byte(`{"model":"gpt-5.5","prompt":"hello","commands":["unsupported"],"truncation":"auto","stop_sequences":["END"],"chat_template_kwargs":{"enable_thinking":true}}`) + + normalized, changed, err := normalizeOpenAIPassthroughOAuthBody(body, false) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "hello", gjson.GetBytes(normalized, "input.0.content").String()) + for _, field := range []string{"prompt", "commands", "truncation", "stop_sequences", "chat_template_kwargs"} { + require.False(t, gjson.GetBytes(normalized, field).Exists(), field) + } +} + +func TestNormalizeOpenAIPassthroughOAuthBody_NormalizesReasoningMode(t *testing.T) { + body := []byte(`{"model":"gpt-5.6-sol","input":"hello","reasoning":{"mode":"pro"}}`) + + normalized, changed, err := normalizeOpenAIPassthroughOAuthBody(body, false) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "max", gjson.GetBytes(normalized, "reasoning.effort").String()) + require.False(t, gjson.GetBytes(normalized, "reasoning.mode").Exists()) +} + +func TestNormalizeOpenAIOAuthResponsesCompatibilityBody_PreservesExplicitInput(t *testing.T) { + body := []byte(`{"model":"gpt-5.5","input":"explicit","prompt":"legacy"}`) + + normalized, changed, err := normalizeOpenAIOAuthResponsesCompatibilityBody(body) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "explicit", gjson.GetBytes(normalized, "input").String()) + require.False(t, gjson.GetBytes(normalized, "prompt").Exists()) +} + +func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_OnlyStripsOAuthFields(t *testing.T) { + body := []byte(`{"type":"response.create","prompt":"hello","commands":{},"truncation":"auto","stop_sequences":["END"],"chat_template_kwargs":{"enable_thinking":true}}`) + + oauthBody, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "hello", gjson.GetBytes(oauthBody, "input").String()) + for _, field := range []string{"prompt", "commands", "truncation", "stop_sequences", "chat_template_kwargs"} { + require.False(t, gjson.GetBytes(oauthBody, field).Exists(), field) + } + + apiKeyBody, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}) + require.NoError(t, err) + require.False(t, changed) + require.JSONEq(t, string(body), string(apiKeyBody)) +} + +func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_SanitizesNativeItemIDs(t *testing.T) { + body := []byte(`{"type":"response.create","model":"gpt-5.6-sol","input":[` + + `{"type":"custom_tool_call","id":"fc_wrong_custom","call_id":"call_custom_1","name":"apply_patch","input":"patch"},` + + `{"type":"custom_tool_call","id":"ctc_valid","call_id":"call_custom_2","name":"apply_patch","input":"patch"},` + + `{"type":"tool_search_call","id":"fc_wrong_search","call_id":"call_search_1","arguments":{"query":"docs"}},` + + `{"type":"tool_search_call","id":"tsc_valid","call_id":"call_search_2","arguments":{"query":"docs"}}]}`) + + for _, oauth := range []bool{false, true} { + accountType := AccountTypeAPIKey + if oauth { + accountType = AccountTypeOAuth + } + normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: accountType}) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "response.create", gjson.GetBytes(normalized, "type").String()) + require.False(t, gjson.GetBytes(normalized, "input.0.id").Exists()) + require.Equal(t, "ctc_valid", gjson.GetBytes(normalized, "input.1.id").String()) + require.False(t, gjson.GetBytes(normalized, "input.2.id").Exists()) + require.Equal(t, "tsc_valid", gjson.GetBytes(normalized, "input.3.id").String()) + // Native Responses call_id values are correlation keys, not item IDs. + require.Equal(t, "call_custom_1", gjson.GetBytes(normalized, "input.0.call_id").String()) + require.Equal(t, "call_search_1", gjson.GetBytes(normalized, "input.2.call_id").String()) + } +} + +func TestNormalizeOpenAIResponsesReasoningMode(t *testing.T) { + tests := []struct { + name string + body string + wantEffort string + }{ + {name: "pro maps to max", body: `{"reasoning":{"mode":"pro"}}`, wantEffort: "max"}, + {name: "explicit effort wins", body: `{"reasoning":{"mode":"pro","effort":"high"}}`, wantEffort: "high"}, + {name: "other mode only removed", body: `{"reasoning":{"mode":"standard"}}`}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + normalized, changed, err := normalizeOpenAIResponsesReasoningMode([]byte(tt.body)) + require.NoError(t, err) + require.True(t, changed) + require.False(t, gjson.GetBytes(normalized, "reasoning.mode").Exists()) + require.Equal(t, tt.wantEffort, gjson.GetBytes(normalized, "reasoning.effort").String()) + }) + } +} + +func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_ReasoningModeAccountScope(t *testing.T) { + body := []byte(`{"type":"response.create","reasoning":{"mode":"pro"}}`) + for _, accountType := range []string{AccountTypeOAuth, AccountTypeSetupToken} { + normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: accountType}) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "max", gjson.GetBytes(normalized, "reasoning.effort").String()) + require.False(t, gjson.GetBytes(normalized, "reasoning.mode").Exists()) + } + apiKeyBody, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}) + require.NoError(t, err) + require.False(t, changed) + require.JSONEq(t, string(body), string(apiKeyBody)) +} + +func TestNormalizeOpenAIResponsesWebSocketCompatibilityBody_SanitizesToolSchemas(t *testing.T) { + body := []byte(`{"type":"response.create","tools":[{"type":"function","name":"search","parameters":{"type":null,"properties":{"q":{"type":"string","pattern":"^(?=.*foo).+$"}}}}]}`) + for _, accountType := range []string{AccountTypeAPIKey, AccountTypeOAuth} { + normalized, changed, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(body, &Account{Platform: PlatformOpenAI, Type: accountType}) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(normalized, "tools.0.parameters.type").String()) + require.False(t, gjson.GetBytes(normalized, "tools.0.parameters.properties.q.pattern").Exists()) + } +} + +func TestNormalizeOpenAIResponseFormatSchemasBody_PreservesNonStrictOptionalFields(t *testing.T) { + body := []byte(`{"text":{"format":{"type":"json_schema","strict":false,"schema":{"properties":{"tags":{"items":{"type":"string"},"uniqueItems":true}},"minProperties":1,"maxProperties":4}}}}`) + + normalized, changed, err := normalizeOpenAIResponseFormatSchemasBody(body) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(normalized, "text.format.schema.type").String()) + require.Equal(t, "array", gjson.GetBytes(normalized, "text.format.schema.properties.tags.type").String()) + require.False(t, gjson.GetBytes(normalized, "text.format.schema.minProperties").Exists()) + require.False(t, gjson.GetBytes(normalized, "text.format.schema.properties.tags.uniqueItems").Exists()) + require.Equal(t, int64(4), gjson.GetBytes(normalized, "text.format.schema.maxProperties").Int()) + require.False(t, gjson.GetBytes(normalized, "text.format.schema.required").Exists()) +} + +func TestNormalizeOpenAIResponseFormatSchemasBody_DoesNotExpandStrictSchema(t *testing.T) { + body := []byte(`{"response_format":{"type":"json_schema","json_schema":{"strict":true,"schema":{"properties":{"name":{"type":"string"}}}}}}`) + + normalized, changed, err := normalizeOpenAIResponseFormatSchemasBody(body) + require.NoError(t, err) + require.True(t, changed) // Safe type inference still applies. + require.Equal(t, "object", gjson.GetBytes(normalized, "response_format.json_schema.schema.type").String()) + require.False(t, gjson.GetBytes(normalized, "response_format.json_schema.schema.required").Exists()) + require.False(t, gjson.GetBytes(normalized, "response_format.json_schema.schema.additionalProperties").Exists()) +} + +func TestNormalizeOpenAIResponseFormatSchemasBody_TraversesNestedSchemaContainers(t *testing.T) { + body := []byte(`{ + "text":{"format":{"type":"json_schema","schema":{ + "$defs":{"entry":{"properties":{"name":{"type":"string"}},"minProperties":1,"maxProperties":3}}, + "additionalProperties":{"items":{"type":"string"},"uniqueItems":true}, + "prefixItems":[{"properties":{"id":{"type":"string"}},"minProperties":1}], + "dependentSchemas":{"kind":{"properties":{"value":{"type":"string"}},"uniqueItems":true}}, + "not":{"items":{"type":"string"},"minProperties":1} + }}} + }`) + + normalized, changed, err := normalizeOpenAIResponseFormatSchemasBody(body) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(normalized, "text.format.schema.$defs.entry.type").String()) + require.False(t, gjson.GetBytes(normalized, "text.format.schema.$defs.entry.minProperties").Exists()) + require.Equal(t, int64(3), gjson.GetBytes(normalized, "text.format.schema.$defs.entry.maxProperties").Int()) + require.Equal(t, "array", gjson.GetBytes(normalized, "text.format.schema.additionalProperties.type").String()) + require.False(t, gjson.GetBytes(normalized, "text.format.schema.additionalProperties.uniqueItems").Exists()) + require.Equal(t, "object", gjson.GetBytes(normalized, "text.format.schema.prefixItems.0.type").String()) + require.False(t, gjson.GetBytes(normalized, "text.format.schema.prefixItems.0.minProperties").Exists()) + require.Equal(t, "object", gjson.GetBytes(normalized, "text.format.schema.dependentSchemas.kind.type").String()) + require.False(t, gjson.GetBytes(normalized, "text.format.schema.dependentSchemas.kind.uniqueItems").Exists()) + require.Equal(t, "array", gjson.GetBytes(normalized, "text.format.schema.not.type").String()) + require.False(t, gjson.GetBytes(normalized, "text.format.schema.not.minProperties").Exists()) + require.False(t, gjson.GetBytes(normalized, "text.format.schema.required").Exists()) +} + +func TestNormalizeOpenAIResponseFormatSchemasBody_PreservesExistingTypeValues(t *testing.T) { + body := []byte(`{"text":{"format":{"type":"json_schema","schema":{"properties":{"union":{"type":["object","null"],"properties":{"name":{"type":"string"}}},"custom":{"type":{"vendor":"shape"},"properties":{"id":{"type":"string"}}},"inferred":{"type":null,"items":{"type":"string"}}}}}}}`) + + normalized, changed, err := normalizeOpenAIResponseFormatSchemasBody(body) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(normalized, "text.format.schema.type").String()) + require.Equal(t, "object", gjson.GetBytes(normalized, "text.format.schema.properties.union.type.0").String()) + require.Equal(t, "null", gjson.GetBytes(normalized, "text.format.schema.properties.union.type.1").String()) + require.True(t, gjson.GetBytes(normalized, "text.format.schema.properties.custom.type").IsObject()) + require.Equal(t, "array", gjson.GetBytes(normalized, "text.format.schema.properties.inferred.type").String()) +} + func TestNormalizeOpenAIPassthroughOAuthBody_CompactRemovesUnsupportedUser(t *testing.T) { body := []byte(`{"model":"gpt-5.4","input":"hello","user":"user_123","metadata":{"user_id":"user_123"},"stream":true,"store":true}`) diff --git a/backend/internal/service/openai_responses_input_compat.go b/backend/internal/service/openai_responses_input_compat.go new file mode 100644 index 0000000000..6e01e28d85 --- /dev/null +++ b/backend/internal/service/openai_responses_input_compat.go @@ -0,0 +1,152 @@ +package service + +import ( + "bytes" + "strings" +) + +const openAIResponsesInputTextMaxChars = 10000000 + +// sanitizeOpenAIResponsesOrphanToolOutputs removes tool-output items that have +// no matching call or item reference anywhere in the current input. +func sanitizeOpenAIResponsesOrphanToolOutputs(reqBody map[string]any, input []any, hasPreviousResponseID bool) bool { + if len(input) == 0 || hasPreviousResponseID { + return false + } + + toolCallIDs := make(map[string]struct{}, len(input)) + referenceIDs := make(map[string]struct{}, len(input)) + for _, rawItem := range input { + item, ok := rawItem.(map[string]any) + if !ok { + continue + } + itemType := strings.TrimSpace(firstNonEmptyString(item["type"])) + if itemType == "item_reference" { + if id := strings.TrimSpace(firstNonEmptyString(item["id"])); id != "" { + referenceIDs[id] = struct{}{} + } + continue + } + if !isCodexToolCallContextItemType(itemType) { + continue + } + if id := strings.TrimSpace(firstNonEmptyString(item["call_id"], item["id"])); id != "" { + toolCallIDs[id] = struct{}{} + } + } + + modified := false + normalized := make([]any, 0, len(input)) + for _, rawItem := range input { + item, ok := rawItem.(map[string]any) + if !ok || !isCodexToolCallOutputItemType(strings.TrimSpace(firstNonEmptyString(item["type"]))) { + normalized = append(normalized, rawItem) + continue + } + + callID := strings.TrimSpace(firstNonEmptyString(item["call_id"])) + _, hasToolCall := toolCallIDs[callID] + _, hasReference := referenceIDs[callID] + if callID != "" && (hasToolCall || hasReference) { + normalized = append(normalized, rawItem) + continue + } + + modified = true + } + if !modified { + return false + } + reqBody["input"] = normalized + return true +} + +func truncateOpenAIResponsesInputText(reqBody map[string]any) bool { + input, ok := reqBody["input"].([]any) + if !ok || len(input) == 0 { + return false + } + modified := false + for _, rawItem := range input { + item, ok := rawItem.(map[string]any) + if !ok { + continue + } + itemType := strings.TrimSpace(firstNonEmptyString(item["type"])) + if isCodexToolCallOutputItemType(itemType) { + if output, ok := item["output"].(string); ok { + if truncated, changed := truncateOpenAIResponsesInputString(output); changed { + item["output"] = truncated + modified = true + } + } + } + if truncateOpenAIResponsesMessageText(item) { + modified = true + } + } + return modified +} + +func openAIResponsesInputMayNeedTruncation(body []byte) bool { + if len(body) <= openAIResponsesInputTextMaxChars { + return false + } + if bytes.Contains(body, []byte(`"text"`)) && bytes.Contains(body, []byte(`"content"`)) { + return true + } + for _, itemType := range []string{ + "function_call_output", + "tool_search_output", + "custom_tool_call_output", + "mcp_tool_call_output", + } { + if bytes.Contains(body, []byte(itemType)) { + return true + } + } + return false +} + +func truncateOpenAIResponsesMessageText(item map[string]any) bool { + itemType := strings.TrimSpace(firstNonEmptyString(item["type"])) + role := strings.TrimSpace(firstNonEmptyString(item["role"])) + if itemType != "message" && role == "" { + return false + } + parts, ok := item["content"].([]any) + if !ok { + return false + } + modified := false + for _, rawPart := range parts { + part, ok := rawPart.(map[string]any) + if !ok { + continue + } + text, ok := part["text"].(string) + if !ok { + continue + } + if truncated, changed := truncateOpenAIResponsesInputString(text); changed { + part["text"] = truncated + modified = true + } + } + return modified +} + +func truncateOpenAIResponsesInputString(value string) (string, bool) { + if len(value) <= openAIResponsesInputTextMaxChars { + return value, false + } + chars := 0 + for index := range value { + if chars == openAIResponsesInputTextMaxChars { + return value[:index], true + } + chars++ + } + return value, false +} diff --git a/backend/internal/service/openai_responses_input_compat_test.go b/backend/internal/service/openai_responses_input_compat_test.go new file mode 100644 index 0000000000..86a8841521 --- /dev/null +++ b/backend/internal/service/openai_responses_input_compat_test.go @@ -0,0 +1,159 @@ +package service + +import ( + "context" + "net/http" + "strings" + "testing" + "unicode/utf8" + + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestSanitizeOpenAIResponsesOrphanToolOutputs(t *testing.T) { + t.Run("preserves matches regardless of item order", func(t *testing.T) { + input := []any{ + map[string]any{"type": "tool_search_output", "call_id": "search_1", "output": "first"}, + map[string]any{"type": "tool_search_call", "id": "search_1", "query": "docs"}, + map[string]any{"type": "custom_tool_call_output", "call_id": "custom_1", "output": "second"}, + map[string]any{"type": "item_reference", "id": "custom_1"}, + } + reqBody := map[string]any{"input": input} + + require.False(t, sanitizeOpenAIResponsesOrphanToolOutputs(reqBody, input, false)) + require.Equal(t, input, reqBody["input"]) + }) + + t.Run("outputs do not legitimize each other", func(t *testing.T) { + input := []any{ + map[string]any{"type": "function_call_output", "call_id": "missing", "output": "one"}, + map[string]any{"type": "tool_search_output", "call_id": "missing", "output": "two"}, + map[string]any{"type": "custom_tool_call_output", "call_id": "missing", "output": "three"}, + map[string]any{"type": "mcp_tool_call_output", "call_id": "missing", "output": "four"}, + } + reqBody := map[string]any{"input": input} + + require.True(t, sanitizeOpenAIResponsesOrphanToolOutputs(reqBody, input, false)) + got := reqBody["input"].([]any) + require.Empty(t, got) + }) + + t.Run("preserves all output variants with matching calls", func(t *testing.T) { + pairs := []struct { + callType string + outputType string + }{ + {callType: "function_call", outputType: "function_call_output"}, + {callType: "tool_search_call", outputType: "tool_search_output"}, + {callType: "custom_tool_call", outputType: "custom_tool_call_output"}, + {callType: "mcp_tool_call", outputType: "mcp_tool_call_output"}, + } + input := make([]any, 0, len(pairs)*2) + for index, pair := range pairs { + callID := string(rune('a' + index)) + input = append(input, + map[string]any{"type": pair.callType, "call_id": callID}, + map[string]any{"type": pair.outputType, "call_id": callID, "output": "ok"}, + ) + } + reqBody := map[string]any{"input": input} + + require.False(t, sanitizeOpenAIResponsesOrphanToolOutputs(reqBody, input, false)) + }) + + t.Run("previous response may contain the missing call", func(t *testing.T) { + input := []any{map[string]any{"type": "function_call_output", "call_id": "remote", "output": "ok"}} + reqBody := map[string]any{"input": input, "previous_response_id": "resp_1"} + + require.False(t, sanitizeOpenAIResponsesOrphanToolOutputs(reqBody, input, true)) + require.Equal(t, input, reqBody["input"]) + }) +} + +func TestTruncateOpenAIResponsesInputText(t *testing.T) { + oversized := strings.Repeat("a", openAIResponsesInputTextMaxChars+1) + input := []any{ + map[string]any{"type": "function_call_output", "call_id": "a", "output": oversized}, + map[string]any{"type": "tool_search_output", "call_id": "b", "output": oversized}, + map[string]any{"type": "custom_tool_call_output", "call_id": "c", "output": oversized}, + map[string]any{"type": "mcp_tool_call_output", "call_id": "d", "output": oversized}, + map[string]any{ + "role": "user", + "content": []any{ + map[string]any{"type": "input_text", "text": "short"}, + map[string]any{"type": "input_text", "text": oversized}, + }, + }, + } + reqBody := map[string]any{"input": input} + + require.True(t, truncateOpenAIResponsesInputText(reqBody)) + for _, rawItem := range input[:4] { + item := rawItem.(map[string]any) + require.Len(t, item["output"].(string), openAIResponsesInputTextMaxChars) + } + content := input[4].(map[string]any)["content"].([]any) + require.Equal(t, "short", content[0].(map[string]any)["text"]) + require.Len(t, content[1].(map[string]any)["text"].(string), openAIResponsesInputTextMaxChars) +} + +func TestTruncateOpenAIResponsesInputStringPreservesUTF8Boundary(t *testing.T) { + value := strings.Repeat("a", openAIResponsesInputTextMaxChars-1) + "中中" + + got, changed := truncateOpenAIResponsesInputString(value) + + require.True(t, changed) + require.True(t, utf8.ValidString(got)) + require.Equal(t, openAIResponsesInputTextMaxChars, utf8.RuneCountInString(got)) + require.Equal(t, "中", got[len(got)-len("中"):]) +} + +func TestOpenAIResponsesInputMayNeedTruncation(t *testing.T) { + short := []byte(`{"input":[{"type":"function_call_output","output":"ok"}]}`) + largeUnrelated := []byte(`{"input":"` + strings.Repeat("x", openAIResponsesInputTextMaxChars+1) + `"}`) + largeOutput := []byte(`{"input":[{"type":"function_call_output","output":"` + strings.Repeat("x", openAIResponsesInputTextMaxChars+1) + `"}]}`) + + require.False(t, openAIResponsesInputMayNeedTruncation(short)) + require.False(t, openAIResponsesInputMayNeedTruncation(largeUnrelated)) + require.True(t, openAIResponsesInputMayNeedTruncation(largeOutput)) +} + +func TestOpenAIGatewayService_OAuthDropsOrphanAfterDroppingPreviousResponse(t *testing.T) { + body := []byte(`{"model":"gpt-5.5","stream":false,"previous_response_id":"resp_missing","input":[{"type":"function_call_output","call_id":"call_missing","output":"keep this result"}]}`) + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + newOpenAIRejectedFieldTestResponse(http.StatusOK, `{"id":"resp_ok","output":[],"usage":{"input_tokens":1,"output_tokens":1,"input_tokens_details":{"cached_tokens":0}}}`), + }} + + result, err := newOpenAIRejectedFieldTestService(upstream).Forward( + context.Background(), + newOpenAIRejectedFieldTestContext(body), + newOpenAIOAuthNamespaceTestAccount(), + body, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.bodies, 1) + require.False(t, gjson.GetBytes(upstream.bodies[0], "previous_response_id").Exists()) + require.Empty(t, gjson.GetBytes(upstream.bodies[0], "input").Array()) +} + +func TestOpenAIGatewayService_TruncatesOversizedToolOutputBeforeForward(t *testing.T) { + body := []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"function_call","call_id":"call_1","name":"lookup","arguments":"{}"},{"type":"function_call_output","call_id":"call_1","output":"` + strings.Repeat("x", openAIResponsesInputTextMaxChars+1) + `"}]}`) + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + newOpenAIRejectedFieldTestResponse(http.StatusOK, `{"id":"resp_ok","output":[],"usage":{"input_tokens":1,"output_tokens":1,"input_tokens_details":{"cached_tokens":0}}}`), + }} + + result, err := newOpenAIRejectedFieldTestService(upstream).Forward( + context.Background(), + newOpenAIRejectedFieldTestContext(body), + newOpenAIRejectedFieldTestAccount(), + body, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.bodies, 1) + require.Equal(t, openAIResponsesInputTextMaxChars, len(gjson.GetBytes(upstream.bodies[0], "input.1.output").String())) +} diff --git a/backend/internal/service/openai_responses_item_id.go b/backend/internal/service/openai_responses_item_id.go index 34443944b2..62a990afbb 100644 --- a/backend/internal/service/openai_responses_item_id.go +++ b/backend/internal/service/openai_responses_item_id.go @@ -8,22 +8,46 @@ import ( "github.com/tidwall/sjson" ) +func openAIResponsesInputItemIDPrefix(itemType string) (string, bool) { + switch strings.TrimSpace(itemType) { + case "message": + return "msg", true + case "reasoning": + return "rs", true + case "custom_tool_call": + return openAIResponsesToolCallIDPrefix(itemType), true + case "tool_search_call": + return openAIResponsesToolCallIDPrefix(itemType), true + default: + if isCodexToolCallInputType(itemType) { + return openAIResponsesToolCallIDPrefix(itemType), true + } + return "", false + } +} + +func openAIResponsesToolCallIDPrefix(itemType string) string { + switch strings.TrimSpace(itemType) { + case "custom_tool_call", "custom_tool_call_output": + return "ctc" + case "tool_search_call", "tool_search_output": + return "tsc" + default: + return "fc" + } +} + // Invalid replayed IDs are removed rather than rewritten because a fabricated -// msg/fc ID may point at a different upstream object. +// ID may point at a different upstream object. func shouldStripOpenAIResponsesInputItemID(itemType, id string) bool { if id == "" { return false } - if itemType == "message" { - return !strings.HasPrefix(id, "msg") + prefix, constrained := openAIResponsesInputItemIDPrefix(itemType) + if !constrained { + return false } - if itemType == "reasoning" { - return !strings.HasPrefix(id, "rs") - } - if isCodexToolCallInputType(itemType) { - return !strings.HasPrefix(id, "fc") - } - return false + return !strings.HasPrefix(id, prefix) } func sanitizeOpenAIResponsesInputItemIDs(body []byte) ([]byte, bool, error) { diff --git a/backend/internal/service/openai_responses_rejected_field_retry.go b/backend/internal/service/openai_responses_rejected_field_retry.go index d8d2bb02bf..592d1ac31a 100644 --- a/backend/internal/service/openai_responses_rejected_field_retry.go +++ b/backend/internal/service/openai_responses_rejected_field_retry.go @@ -15,8 +15,13 @@ import ( const maxOpenAIResponsesRejectedFieldRetries = 6 var ( - openAIResponsesRejectedNamespaceParamPattern = regexp.MustCompile(`(?i)^input\[(\d+)\]\.namespace$`) - openAIResponsesRejectedMessageParamPattern = regexp.MustCompile(`(?i)(?:unknown|unsupported)[ _-]+parameter\s*(?::|=|is)?\s*["']?(max_output_tokens|input\[\d+\]\.namespace)(?:["']|\b)`) + openAIResponsesRejectedNamespaceParamPattern = regexp.MustCompile(`(?i)^input\[(\d+)\]\.namespace$`) + openAIResponsesRejectedStatusParamPattern = regexp.MustCompile(`(?i)^input\[(\d+)\]\.status$`) + openAIResponsesRejectedContentParamPattern = regexp.MustCompile(`(?i)^input\[(\d+)\]\.content$`) + openAIResponsesRejectedCacheParamPattern = regexp.MustCompile(`(?i)^input\[(\d+)\]\.prompt_cache_breakpoint$`) + openAIResponsesRejectedMessageParamPattern = regexp.MustCompile(`(?i)(?:unknown|unsupported)[ _-]+parameter\s*(?::|=|is)?\s*["']?(max_output_tokens|input\[\d+\]\.(?:namespace|status))(?:["']|\b)`) + openAIResponsesInvalidTypeMessageParamPattern = regexp.MustCompile(`(?i)invalid[ _-]+type\s+for\s+["']?(input\[\d+\]\.content)(?:["']|\b)[^\n]*\b(?:got|received)\s+null\b`) + openAIResponsesCacheModelRejectionPattern = regexp.MustCompile(`(?i)["']?(prompt_cache_breakpoint|input\[\d+\]\.prompt_cache_breakpoint)["']?\s+is\s+not\s+supported\s+on\s+this\s+model\b`) ) type openAIResponsesRejectedFieldRetryState struct { @@ -62,23 +67,53 @@ func normalizeOpenAIResponsesRejectedFieldRetryBody(statusCode int, body, respon code := strings.ToLower(strings.TrimSpace(extractUpstreamErrorCode(responseBody))) message := strings.ToLower(strings.TrimSpace(extractUpstreamErrorMessage(responseBody))) - if !isExplicitOpenAIResponsesFieldRejection(code, message) { - return nil, "", false, nil + param := strings.ToLower(strings.TrimSpace(gjson.GetBytes(responseBody, "error.param").String())) + cacheMessageParam := openAIResponsesCacheModelRejectionParamFromMessage(message) + cacheParam := param + if cacheParam == "" { + cacheParam = cacheMessageParam + } + cacheParamMatchesMessage := cacheMessageParam == "" || cacheParam == cacheMessageParam + cacheModelRejection := code == "invalid_parameter" || cacheMessageParam != "" + if cacheParam != "" && cacheParamMatchesMessage && cacheModelRejection { + if cacheParam == "prompt_cache_breakpoint" && gjson.GetBytes(body, cacheParam).Exists() { + retryBody, err := sjson.DeleteBytes(body, cacheParam) + if err != nil { + return nil, "", false, fmt.Errorf("delete rejected prompt_cache_breakpoint: %w", err) + } + return retryBody, "prompt_cache_breakpoint parameter rejection", true, nil + } + if index, ok := openAIResponsesRejectedCacheIndex(cacheParam); ok { + return removeOpenAIResponsesRejectedCacheAtIndex(body, index) + } + } + if isExplicitOpenAIResponsesFieldRejection(code, message) { + if param == "" { + param = openAIResponsesRejectedParamFromMessage(message) + } + if index, ok := openAIResponsesRejectedNamespaceIndex(param); ok { + return removeOpenAIResponsesRejectedNamespaceAtIndex(body, index) + } + if index, ok := openAIResponsesRejectedStatusIndex(param); ok { + return removeOpenAIResponsesRejectedStatusAtIndex(body, index) + } + if param == "max_output_tokens" && gjson.GetBytes(body, "max_output_tokens").Exists() { + retryBody, err := sjson.DeleteBytes(body, "max_output_tokens") + if err != nil { + return nil, "", false, fmt.Errorf("delete rejected max_output_tokens: %w", err) + } + return retryBody, "max_output_tokens parameter rejection", true, nil + } } - param := strings.ToLower(strings.TrimSpace(gjson.GetBytes(responseBody, "error.param").String())) - if param == "" { - param = openAIResponsesRejectedParamFromMessage(message) + messageContentParam := openAIResponsesInvalidTypeParamFromMessage(message) + contentParam := param + if contentParam == "" { + contentParam = messageContentParam } - if index, ok := openAIResponsesRejectedNamespaceIndex(param); ok { - return removeOpenAIResponsesRejectedNamespaceAtIndex(body, index) - } - if param == "max_output_tokens" && gjson.GetBytes(body, "max_output_tokens").Exists() { - retryBody, err := sjson.DeleteBytes(body, "max_output_tokens") - if err != nil { - return nil, "", false, fmt.Errorf("delete rejected max_output_tokens: %w", err) - } - return retryBody, "max_output_tokens parameter rejection", true, nil + if index, ok := openAIResponsesRejectedContentIndex(contentParam); ok && + contentParam == messageContentParam && isExplicitOpenAIResponsesNullContentRejection(code, message) { + return normalizeOpenAIResponsesRejectedNullContentAtIndex(body, index) } return nil, "", false, nil } @@ -100,8 +135,46 @@ func openAIResponsesRejectedParamFromMessage(message string) string { return strings.ToLower(strings.TrimSpace(match[1])) } +func openAIResponsesInvalidTypeParamFromMessage(message string) string { + match := openAIResponsesInvalidTypeMessageParamPattern.FindStringSubmatch(strings.TrimSpace(message)) + if len(match) != 2 { + return "" + } + return strings.ToLower(strings.TrimSpace(match[1])) +} + +func openAIResponsesCacheModelRejectionParamFromMessage(message string) string { + match := openAIResponsesCacheModelRejectionPattern.FindStringSubmatch(strings.TrimSpace(message)) + if len(match) != 2 { + return "" + } + return strings.ToLower(strings.TrimSpace(match[1])) +} + +func isExplicitOpenAIResponsesNullContentRejection(code, message string) bool { + code = strings.TrimSpace(code) + return (code == "invalid_type" || code == "invalid_request_error" || code == "") && + openAIResponsesInvalidTypeMessageParamPattern.MatchString(strings.TrimSpace(message)) +} + func openAIResponsesRejectedNamespaceIndex(param string) (int, bool) { - match := openAIResponsesRejectedNamespaceParamPattern.FindStringSubmatch(strings.TrimSpace(param)) + return openAIResponsesRejectedInputIndex(openAIResponsesRejectedNamespaceParamPattern, param) +} + +func openAIResponsesRejectedStatusIndex(param string) (int, bool) { + return openAIResponsesRejectedInputIndex(openAIResponsesRejectedStatusParamPattern, param) +} + +func openAIResponsesRejectedContentIndex(param string) (int, bool) { + return openAIResponsesRejectedInputIndex(openAIResponsesRejectedContentParamPattern, param) +} + +func openAIResponsesRejectedCacheIndex(param string) (int, bool) { + return openAIResponsesRejectedInputIndex(openAIResponsesRejectedCacheParamPattern, param) +} + +func openAIResponsesRejectedInputIndex(pattern *regexp.Regexp, param string) (int, bool) { + match := pattern.FindStringSubmatch(strings.TrimSpace(param)) if len(match) != 2 { return 0, false } @@ -112,6 +185,67 @@ func openAIResponsesRejectedNamespaceIndex(param string) (int, bool) { return 0, false } +func removeOpenAIResponsesRejectedStatusAtIndex(body []byte, index int) ([]byte, string, bool, error) { + itemPath := fmt.Sprintf("input.%d", index) + if !gjson.GetBytes(body, itemPath).IsObject() { + return nil, "", false, nil + } + statusPath := itemPath + ".status" + if !gjson.GetBytes(body, statusPath).Exists() { + return nil, "", false, nil + } + retryBody, err := sjson.DeleteBytes(body, statusPath) + if err != nil { + return nil, "", false, fmt.Errorf("delete rejected status at input[%d]: %w", index, err) + } + return retryBody, "indexed status parameter rejection", true, nil +} + +func removeOpenAIResponsesRejectedCacheAtIndex(body []byte, index int) ([]byte, string, bool, error) { + itemPath := fmt.Sprintf("input.%d", index) + if !gjson.GetBytes(body, itemPath).IsObject() { + return nil, "", false, nil + } + cachePath := itemPath + ".prompt_cache_breakpoint" + if !gjson.GetBytes(body, cachePath).Exists() { + return nil, "", false, nil + } + retryBody, err := sjson.DeleteBytes(body, cachePath) + if err != nil { + return nil, "", false, fmt.Errorf("delete rejected prompt_cache_breakpoint at input[%d]: %w", index, err) + } + return retryBody, "indexed prompt_cache_breakpoint parameter rejection", true, nil +} + +func normalizeOpenAIResponsesRejectedNullContentAtIndex(body []byte, index int) ([]byte, string, bool, error) { + itemPath := fmt.Sprintf("input.%d", index) + item := gjson.GetBytes(body, itemPath) + content := gjson.GetBytes(body, itemPath+".content") + if !item.IsObject() || !content.Exists() || content.Type != gjson.Null { + return nil, "", false, nil + } + + itemType := strings.ToLower(strings.TrimSpace(item.Get("type").String())) + role := strings.TrimSpace(item.Get("role").String()) + contentPath := itemPath + ".content" + switch { + case itemType == "reasoning": + retryBody, err := sjson.DeleteBytes(body, contentPath) + if err != nil { + return nil, "", false, fmt.Errorf("delete rejected null content at input[%d]: %w", index, err) + } + return retryBody, "indexed reasoning null content rejection", true, nil + case itemType == "message" || role != "": + retryBody, err := sjson.SetBytes(body, contentPath, "") + if err != nil { + return nil, "", false, fmt.Errorf("normalize rejected null content at input[%d]: %w", index, err) + } + return retryBody, "indexed message null content rejection", true, nil + default: + return nil, "", false, nil + } +} + func removeOpenAIResponsesRejectedNamespaceAtIndex(body []byte, index int) ([]byte, string, bool, error) { itemPath := fmt.Sprintf("input.%d", index) itemType := strings.ToLower(strings.TrimSpace(gjson.GetBytes(body, itemPath+".type").String())) diff --git a/backend/internal/service/openai_responses_rejected_field_retry_test.go b/backend/internal/service/openai_responses_rejected_field_retry_test.go index 724fd90c37..4929ac3a0c 100644 --- a/backend/internal/service/openai_responses_rejected_field_retry_test.go +++ b/backend/internal/service/openai_responses_rejected_field_retry_test.go @@ -114,6 +114,270 @@ func TestNormalizeOpenAIResponsesRejectedFieldRetryBodyBindsMaxOutputTokensToRej require.False(t, gjson.GetBytes(retryBody, "max_output_tokens").Exists()) } +func TestNormalizeOpenAIResponsesRejectedFieldRetryBodyRemovesExactIndexedStatus(t *testing.T) { + body := []byte(`{"input":[{"type":"message","status":"keep","content":"one"},{"type":"reasoning","status":"remove","summary":[]}]}`) + responses := []struct { + name string + body []byte + }{ + { + name: "structured param", + body: []byte(`{"error":{"code":"unknown_parameter","message":"Unknown parameter: 'input[1].status'.","param":"input[1].status"}}`), + }, + { + name: "message param", + body: []byte(`{"error":{"code":"unsupported_parameter","message":"Unsupported parameter: input[1].status."}}`), + }, + } + for _, tt := range responses { + t.Run(tt.name, func(t *testing.T) { + retryBody, _, changed, err := normalizeOpenAIResponsesRejectedFieldRetryBody(http.StatusBadRequest, body, tt.body) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "keep", gjson.GetBytes(retryBody, "input.0.status").String()) + require.False(t, gjson.GetBytes(retryBody, "input.1.status").Exists()) + }) + } +} + +func TestNormalizeOpenAIResponsesRejectedFieldRetryBodyNormalizesExactNullContent(t *testing.T) { + tests := []struct { + name string + body []byte + wantChange bool + wantValue string + wantExists bool + }{ + { + name: "message becomes empty string", + body: []byte(`{"input":[{"type":"message","role":"assistant","content":null}]}`), + wantChange: true, + wantValue: "", + wantExists: true, + }, + { + name: "reasoning content is removed", + body: []byte(`{"input":[{"type":"reasoning","content":null,"summary":[]}]}`), + wantChange: true, + wantExists: false, + }, + { + name: "unknown item is unchanged", + body: []byte(`{"input":[{"type":"future_item","content":null}]}`), + wantChange: false, + }, + { + name: "non null content is unchanged", + body: []byte(`{"input":[{"type":"message","content":"keep"}]}`), + wantChange: false, + }, + } + responseBody := []byte(`{"error":{"code":"invalid_type","message":"Invalid type for 'input[0].content': expected one of a string or a list of input items, but got null instead.","param":"input[0].content"}}`) + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + retryBody, _, changed, err := normalizeOpenAIResponsesRejectedFieldRetryBody(http.StatusBadRequest, tt.body, responseBody) + require.NoError(t, err) + require.Equal(t, tt.wantChange, changed) + if !tt.wantChange { + require.Nil(t, retryBody) + return + } + content := gjson.GetBytes(retryBody, "input.0.content") + require.Equal(t, tt.wantExists, content.Exists()) + if tt.wantExists { + require.Equal(t, tt.wantValue, content.String()) + require.Equal(t, gjson.String, content.Type) + } + }) + } +} + +func TestNormalizeOpenAIResponsesRejectedFieldRetryBodyRejectsUnsafeIndexedMutations(t *testing.T) { + tests := []struct { + name string + body []byte + responseBody []byte + }{ + { + name: "nested status path", + body: []byte(`{"input":[{"type":"message","content":{"status":"keep"}}]}`), + responseBody: []byte(`{"error":{"code":"unknown_parameter","message":"Unknown parameter: input[0].content.status.","param":"input[0].content.status"}}`), + }, + { + name: "status index out of bounds", + body: []byte(`{"input":[{"type":"message","status":"keep"}]}`), + responseBody: []byte(`{"error":{"code":"unknown_parameter","message":"Unknown parameter: input[4].status.","param":"input[4].status"}}`), + }, + { + name: "status path only mentioned", + body: []byte(`{"input":[{"type":"message","status":"keep"}]}`), + responseBody: []byte(`{"error":{"code":"invalid_request_error","message":"input[0].status must be completed","param":"input[0].status"}}`), + }, + { + name: "content param and message disagree", + body: []byte(`{"input":[{"type":"message","content":null},{"type":"message","content":null}]}`), + responseBody: []byte(`{"error":{"code":"invalid_type","message":"Invalid type for input[1].content: got null instead.","param":"input[0].content"}}`), + }, + { + name: "content error only mentions null", + body: []byte(`{"input":[{"type":"message","content":null}]}`), + responseBody: []byte(`{"error":{"code":"invalid_type","message":"content cannot be null","param":"input[0].content"}}`), + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + retryBody, _, changed, err := normalizeOpenAIResponsesRejectedFieldRetryBody(http.StatusBadRequest, tt.body, tt.responseBody) + require.NoError(t, err) + require.False(t, changed) + require.Nil(t, retryBody) + }) + } +} + +func TestOpenAIGatewayService_OAuthRetriesExactRejectedStatus(t *testing.T) { + body := []byte(`{"model":"gpt-5.5","stream":true,"instructions":"test","input":[{"type":"message","role":"user","status":"completed","content":"hello"}]}`) + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + newOpenAIRejectedFieldTestResponse(http.StatusBadRequest, `{"error":{"code":"unknown_parameter","message":"Unknown parameter: 'input[0].status'.","param":"input[0].status"}}`), + newOpenAIRejectedFieldTestResponse(http.StatusOK, "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_ok\",\"output\":[],\"usage\":{\"input_tokens\":1,\"output_tokens\":1,\"total_tokens\":2}}}\n\ndata: [DONE]\n\n"), + }} + upstream.responses[1].Header.Set("Content-Type", "text/event-stream") + + result, err := newOpenAIRejectedFieldTestService(upstream).Forward( + context.Background(), newOpenAIRejectedFieldTestContext(body), newOpenAIOAuthNamespaceTestAccount(), body, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.bodies, 2) + require.Equal(t, "completed", gjson.GetBytes(upstream.bodies[0], "input.0.status").String()) + require.False(t, gjson.GetBytes(upstream.bodies[1], "input.0.status").Exists()) +} + +func TestOpenAIGatewayService_APIKeyRetriesExactRejectedNullMessageContent(t *testing.T) { + body := []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"message","role":"assistant","content":null},{"type":"message","role":"user","content":"continue"}]}`) + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + newOpenAIRejectedFieldTestResponse(http.StatusBadRequest, `{"error":{"code":"invalid_type","message":"Invalid type for 'input[0].content': expected one of a string or a list of input items, but got null instead.","param":"input[0].content"}}`), + newOpenAIRejectedFieldTestResponse(http.StatusOK, `{"output":[],"usage":{"input_tokens":1,"output_tokens":1,"input_tokens_details":{"cached_tokens":0}}}`), + }} + + result, err := newOpenAIRejectedFieldTestService(upstream).Forward( + context.Background(), newOpenAIRejectedFieldTestContext(body), newOpenAIRejectedFieldTestAccount(), body, + ) + + require.NoError(t, err) + require.NotNil(t, result) + require.Len(t, upstream.bodies, 2) + require.Equal(t, gjson.Null, gjson.GetBytes(upstream.bodies[0], "input.0.content").Type) + require.Equal(t, gjson.String, gjson.GetBytes(upstream.bodies[1], "input.0.content").Type) + require.Equal(t, "continue", gjson.GetBytes(upstream.bodies[1], "input.1.content").String()) +} + +func TestNormalizeOpenAIResponsesRejectedFieldRetryBodyRemovesModelRejectedPromptCacheBreakpoint(t *testing.T) { + tests := []struct { + name string + body []byte + responseBody []byte + removedPath string + preserved string + reason string + }{ + { + name: "top level", + body: []byte(`{"model":"gpt-5.6-sol","prompt_cache_breakpoint":{"type":"message_start"},"input":"hello"}`), + responseBody: []byte(`{"error":{"code":"invalid_parameter","message":"prompt_cache_breakpoint is not supported on this model","param":"prompt_cache_breakpoint"}}`), + removedPath: "prompt_cache_breakpoint", + preserved: "input", + reason: "prompt_cache_breakpoint parameter rejection", + }, + { + name: "indexed path from message", + body: []byte(`{"input":[{"type":"message","prompt_cache_breakpoint":{"type":"message_start"}},{"type":"message","prompt_cache_breakpoint":{"type":"message_end"}}]}`), + responseBody: []byte(`{"error":{"code":"invalid_parameter","message":"input[1].prompt_cache_breakpoint is not supported on this model"}}`), + removedPath: "input.1.prompt_cache_breakpoint", + preserved: "input.0.prompt_cache_breakpoint", + reason: "indexed prompt_cache_breakpoint parameter rejection", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + retryBody, reason, changed, err := normalizeOpenAIResponsesRejectedFieldRetryBody(http.StatusBadRequest, tt.body, tt.responseBody) + + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, tt.reason, reason) + require.False(t, gjson.GetBytes(retryBody, tt.removedPath).Exists()) + require.True(t, gjson.GetBytes(retryBody, tt.preserved).Exists()) + }) + } +} + +func TestNormalizeOpenAIResponsesRejectedFieldRetryBodyRejectsAmbiguousPromptCacheBreakpointErrors(t *testing.T) { + body := []byte(`{"prompt_cache_breakpoint":{"type":"message_start"},"input":[{"type":"message","prompt_cache_breakpoint":{"type":"message_end"}}]}`) + tests := []struct { + name string + responseBody []byte + }{ + { + name: "structured param disagrees", + responseBody: []byte(`{"error":{"code":"invalid_parameter","message":"input[0].prompt_cache_breakpoint is not supported on this model","param":"prompt_cache_breakpoint"}}`), + }, + { + name: "index out of bounds", + responseBody: []byte(`{"error":{"code":"invalid_parameter","message":"input[4].prompt_cache_breakpoint is not supported on this model","param":"input[4].prompt_cache_breakpoint"}}`), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + retryBody, _, changed, err := normalizeOpenAIResponsesRejectedFieldRetryBody(http.StatusBadRequest, body, tt.responseBody) + + require.NoError(t, err) + require.False(t, changed) + require.Nil(t, retryBody) + }) + } +} + +func TestNormalizeOpenAIResponsesRejectedFieldRetryBodyAcceptsEitherCacheModelRejectionSignal(t *testing.T) { + tests := []struct { + name string + responseBody []byte + }{ + { + name: "invalid parameter code", + responseBody: []byte(`{"error":{"code":"invalid_parameter","message":"This optional cache hint cannot be used here","param":"prompt_cache_breakpoint"}}`), + }, + { + name: "model rejection message", + responseBody: []byte(`{"error":{"code":"invalid_request_error","message":"prompt_cache_breakpoint is not supported on this model","param":"prompt_cache_breakpoint"}}`), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + retryBody, _, changed, err := normalizeOpenAIResponsesRejectedFieldRetryBody(http.StatusBadRequest, []byte(`{"prompt_cache_breakpoint":true,"input":"keep"}`), tt.responseBody) + + require.NoError(t, err) + require.True(t, changed) + require.False(t, gjson.GetBytes(retryBody, "prompt_cache_breakpoint").Exists()) + require.Equal(t, "keep", gjson.GetBytes(retryBody, "input").String()) + }) + } +} + +func TestOpenAIResponsesRejectedFieldRetryStateAllowsPromptCacheBreakpointVariantOnce(t *testing.T) { + body := []byte(`{"input":[{"prompt_cache_breakpoint":{"type":"message_start"}}]}`) + responseBody := []byte(`{"error":{"code":"invalid_parameter","message":"input[0].prompt_cache_breakpoint is not supported on this model","param":"input[0].prompt_cache_breakpoint"}}`) + state := newOpenAIResponsesRejectedFieldRetryState(body) + + retryBody, _, changed, err := normalizeOpenAIResponsesRejectedFieldRetryBody(http.StatusBadRequest, body, responseBody) + require.NoError(t, err) + require.True(t, changed) + require.True(t, state.Allow(retryBody)) + require.False(t, state.Allow(retryBody)) +} + func TestOpenAIGatewayService_APIKeyStripsAllIndexedNamespacesBeforeFirstForward(t *testing.T) { body := []byte(`{"model":"gpt-5.5","stream":false,"input":[{"type":"function_call","name":"first","namespace":"remove-first","arguments":"{}"},{"type":"custom_tool_call","name":"second","namespace":"remove-second","input":"{}"}]}`) upstream := &httpUpstreamRecorder{responses: []*http.Response{ diff --git a/backend/internal/service/openai_responses_tool_schema.go b/backend/internal/service/openai_responses_tool_schema.go index 8c7e7b8d18..5063785066 100644 --- a/backend/internal/service/openai_responses_tool_schema.go +++ b/backend/internal/service/openai_responses_tool_schema.go @@ -2,11 +2,103 @@ package service import ( "bytes" + "encoding/json" "sort" + "strings" "github.com/tidwall/gjson" ) +func sanitizeOpenAIResponsesToolSchemaPatterns(body []byte) ([]byte, bool, error) { + if len(body) == 0 || !bytes.Contains(body, []byte(`"pattern"`)) { + return body, false, nil + } + var root any + decoder := json.NewDecoder(bytes.NewReader(body)) + decoder.UseNumber() + if err := decoder.Decode(&root); err != nil { + return nil, false, err + } + changed := false + var visitSchema func(any) + visitSchema = func(value any) { + node, ok := value.(map[string]any) + if !ok { + return + } + if pattern, ok := node["pattern"].(string); ok && hasRegexLookaround(pattern) { + delete(node, "pattern") + changed = true + } + for _, key := range []string{ + "additionalProperties", "additionalItems", "contains", "not", "if", "then", "else", + "propertyNames", "unevaluatedProperties", "unevaluatedItems", + } { + visitSchema(node[key]) + } + if items, ok := node["items"].(map[string]any); ok { + visitSchema(items) + } else if items, ok := node["items"].([]any); ok { + for _, child := range items { + visitSchema(child) + } + } + for _, key := range []string{"anyOf", "oneOf", "allOf", "prefixItems"} { + children, _ := node[key].([]any) + for _, child := range children { + visitSchema(child) + } + } + for _, key := range []string{"properties", "patternProperties", "$defs", "definitions", "dependentSchemas"} { + children, _ := node[key].(map[string]any) + for _, child := range children { + visitSchema(child) + } + } + if dependencies, ok := node["dependencies"].(map[string]any); ok { + for _, child := range dependencies { + visitSchema(child) + } + } + } + var visitTools func(any) + visitTools = func(value any) { + switch node := value.(type) { + case map[string]any: + if parameters, ok := node["parameters"]; ok { + visitSchema(parameters) + } + if function, ok := node["function"]; ok { + visitTools(function) + } + if tools, ok := node["tools"]; ok { + visitTools(tools) + } + case []any: + for _, child := range node { + visitTools(child) + } + } + } + if document, ok := root.(map[string]any); ok { + visitTools(document["tools"]) + visitTools(document["input"]) + } + if !changed { + return body, false, nil + } + sanitized, err := json.Marshal(root) + if err != nil { + return nil, false, err + } + return sanitized, true, nil +} + +func hasRegexLookaround(pattern string) bool { + return strings.Contains(pattern, "(?=") || strings.Contains(pattern, "(?!") || + strings.Contains(pattern, "(?<=") || strings.Contains(pattern, "(?