diff --git a/backend/internal/handler/openai_gateway_credential_failover_test.go b/backend/internal/handler/openai_gateway_credential_failover_test.go index 91a3881dfa..d94eae7fe9 100644 --- a/backend/internal/handler/openai_gateway_credential_failover_test.go +++ b/backend/internal/handler/openai_gateway_credential_failover_test.go @@ -326,12 +326,15 @@ func TestOpsRecoveredCredentialFailoverDoesNotCreateRequestError(t *testing.T) { 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: - } + require.Equal(t, int64(1), OpsErrorLogQueueLength()) + job := <-opsErrorLogQueue + require.Equal(t, http.StatusOK, job.entry.StatusCode) + require.Equal(t, string(service.GatewayFailureStageAccountAuth), job.entry.ErrorPhase) + 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, string(service.GatewayFailureStageAccountAuth), events[1].Stage) } func TestOpsWebSocketCredentialFailoverSuccessDoesNotCreateRequestError(t *testing.T) { @@ -354,12 +357,15 @@ func TestOpsWebSocketCredentialFailoverSuccessDoesNotCreateRequestError(t *testi 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: - } + require.Equal(t, int64(1), OpsErrorLogQueueLength()) + job := <-opsErrorLogQueue + require.Equal(t, http.StatusOK, job.entry.StatusCode) + require.Equal(t, string(service.GatewayFailureStageAccountAuth), job.entry.ErrorPhase) + require.NotNil(t, job.entry.UpstreamErrorsJSON) + events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON) + require.NoError(t, err) + require.Len(t, events, 1) + require.Equal(t, string(service.GatewayFailureStageAccountAuth), events[0].Stage) } func TestOpsWebSocketCredentialFailoverExhaustedIsRecorded(t *testing.T) { diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 89e00c8954..45144479c7 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -409,7 +409,27 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) { h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "previous_response_id must be a response.id (resp_*), not a message id") return } + groupID := int64(0) + if apiKey.GroupID != nil { + groupID = *apiKey.GroupID + } + owned, ownershipErr := h.gatewayService.ValidateOpenAIHTTPResponseOwner( + c.Request.Context(), + groupID, + previousResponseID, + subject.UserID, + apiKey.ID, + ) + if ownershipErr != nil { + reqLog.Warn("openai.previous_response_owner_lookup_failed", zap.Error(ownershipErr)) + } + if !owned { + reqLog.Warn("openai.request_validation_failed", zap.String("reason", "previous_response_owner_mismatch")) + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", "previous_response_id is not available for this user") + return + } } + service.SetOpenAIHTTPResponseOwner(c, subject.UserID, apiKey.ID) setOpsRequestContext(c, reqModel, reqStream) setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false))) diff --git a/backend/internal/handler/openai_gateway_handler_test.go b/backend/internal/handler/openai_gateway_handler_test.go index 6cc0eaffe7..86c56d4901 100644 --- a/backend/internal/handler/openai_gateway_handler_test.go +++ b/backend/internal/handler/openai_gateway_handler_test.go @@ -877,12 +877,61 @@ func TestOpenAIResponses_AcceptsHTTPContinuationPreviousResponseIDBeforeRouting( }) h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil) + require.NoError(t, h.gatewayService.BindOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_123456", 1, 101)) h.Responses(c) require.NotEqual(t, http.StatusBadRequest, w.Code) require.NotContains(t, w.Body.String(), "Responses WebSocket v2") } +func TestOpenAIResponses_RejectsHTTPContinuationOwnedByAnotherUser(t *testing.T) { + gin.SetMode(gin.TestMode) + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader( + `{"model":"gpt-5.1","stream":false,"previous_response_id":"resp_other_tenant","input":"hello"}`, + )) + c.Request.Header.Set("Content-Type", "application/json") + + groupID := int64(2) + c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{ + ID: 202, + UserID: 2, + GroupID: &groupID, + User: &service.User{ID: 2}, + }) + c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: 2, Concurrency: 1}) + + h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil) + require.NoError(t, h.gatewayService.BindOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_other_tenant", 1, 101)) + h.Responses(c) + + require.Equal(t, http.StatusBadRequest, w.Code) + require.Contains(t, w.Body.String(), "previous_response_id is not available for this user") +} + +func TestOpenAIResponses_RejectsUnownedHTTPContinuation(t *testing.T) { + gin.SetMode(gin.TestMode) + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader( + `{"model":"gpt-5.1","stream":false,"previous_response_id":"resp_unknown","input":"hello"}`, + )) + c.Request.Header.Set("Content-Type", "application/json") + + groupID := int64(2) + c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{ID: 101, UserID: 1, GroupID: &groupID, User: &service.User{ID: 1}}) + c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: 1, Concurrency: 1}) + + h := newOpenAIHandlerForPreviousResponseIDValidation(t, nil) + h.Responses(c) + + require.Equal(t, http.StatusBadRequest, w.Code) + require.Contains(t, w.Body.String(), "previous_response_id is not available for this user") +} + func TestOpenAIResponses_FunctionCallOutputHTTPGuidanceDoesNotSuggestPreviousResponseReuse(t *testing.T) { gin.SetMode(gin.TestMode) diff --git a/backend/internal/handler/ops_capture_writer_nil_test.go b/backend/internal/handler/ops_capture_writer_nil_test.go index fce7680a16..38d3f1433b 100644 --- a/backend/internal/handler/ops_capture_writer_nil_test.go +++ b/backend/internal/handler/ops_capture_writer_nil_test.go @@ -12,6 +12,18 @@ import ( "github.com/stretchr/testify/require" ) +type blockingOpsResponseWriter struct { + gin.ResponseWriter + writeStarted chan struct{} + writeRelease chan struct{} +} + +func (w *blockingOpsResponseWriter) WriteString(s string) (int, error) { + close(w.writeStarted) + <-w.writeRelease + return w.ResponseWriter.WriteString(s) +} + type deterministicOpsCaptureWriterStatePool struct { states []*opsCaptureWriterState } @@ -144,3 +156,56 @@ func TestOpsCaptureWriter_StaleLeaseCannotReachReacquiredState(t *testing.T) { defer releaseOpsCaptureWriter(other) require.NotSame(t, current.state, other.state) } + +func TestOpsCaptureWriter_ReleaseWaitsForDelegatedWriteWithoutHoldingStateMutex(t *testing.T) { + gin.SetMode(gin.TestMode) + pool := &deterministicOpsCaptureWriterStatePool{} + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + inner := &blockingOpsResponseWriter{ + ResponseWriter: ctx.Writer, + writeStarted: make(chan struct{}), + writeRelease: make(chan struct{}), + } + w := acquireOpsCaptureWriterFromPool(pool, inner) + + writeDone := make(chan struct{}) + go func() { + defer close(writeDone) + _, _ = w.WriteString("body") + }() + <-inner.writeStarted + + mutexAvailable := make(chan struct{}) + go func() { + w.state.mu.Lock() + w.state.mu.Unlock() + close(mutexAvailable) + }() + select { + case <-mutexAvailable: + case <-time.After(time.Second): + t.Fatal("state mutex remained held across the delegated network write") + } + + releaseDone := make(chan struct{}) + go func() { + releaseOpsCaptureWriter(w) + close(releaseDone) + }() + select { + case <-releaseDone: + t.Fatal("release returned while a delegated write was still active") + case <-time.After(20 * time.Millisecond): + } + require.Empty(t, pool.states) + + close(inner.writeRelease) + <-writeDone + select { + case <-releaseDone: + case <-time.After(time.Second): + t.Fatal("release did not finish after the delegated write returned") + } + require.Len(t, pool.states, 1) +} diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go index 5498f59706..41cd2119b9 100644 --- a/backend/internal/handler/ops_error_logger.go +++ b/backend/internal/handler/ops_error_logger.go @@ -528,6 +528,7 @@ type opsCaptureWriter struct { type opsCaptureWriterState struct { mu sync.RWMutex + inFlight sync.WaitGroup generation uint64 responseWriter gin.ResponseWriter limit int @@ -605,9 +606,15 @@ func releaseOpsCaptureWriter(w *opsCaptureWriter) { state.mu.Unlock() return } + // Invalidate the lease before waiting. No new delegated calls can start for + // this handle, while calls that already copied the writer keep it alive via + // inFlight until their network operation returns. state.generation++ state.responseWriter = nil state.ctx = nil + state.mu.Unlock() + state.inFlight.Wait() + state.mu.Lock() state.limit = opsCaptureWriterLimit state.probe = state.probe[:0] state.lineProbe = state.lineProbe[:0] @@ -657,6 +664,27 @@ func (w *opsCaptureWriter) lockActiveWrite() (*opsCaptureWriterState, gin.Respon return state, state.responseWriter } +func (w *opsCaptureWriter) beginDelegatedCall() (*opsCaptureWriterState, gin.ResponseWriter) { + if w == nil || w.state == nil { + return nil, nil + } + state := w.state + state.mu.Lock() + if state.generation != w.generation || state.responseWriter == nil { + state.mu.Unlock() + return nil, nil + } + rw := state.responseWriter + state.inFlight.Add(1) + return state, rw +} + +func finishDelegatedCall(state *opsCaptureWriterState) { + if state != nil { + state.inFlight.Done() + } +} + func (w *opsCaptureWriter) setContext(ctx *gin.Context) { state, _ := w.lockActiveWrite() if state == nil { @@ -702,19 +730,21 @@ func (w *opsCaptureWriter) Header() http.Header { return rw.Header() } func (w *opsCaptureWriter) WriteHeader(code int) { - state, rw := w.lockActive() + state, rw := w.beginDelegatedCall() if state == nil { return } - defer state.mu.RUnlock() + state.mu.Unlock() + defer finishDelegatedCall(state) rw.WriteHeader(code) } func (w *opsCaptureWriter) WriteHeaderNow() { - state, rw := w.lockActive() + state, rw := w.beginDelegatedCall() if state == nil { return } - defer state.mu.RUnlock() + state.mu.Unlock() + defer finishDelegatedCall(state) rw.WriteHeaderNow() } func (w *opsCaptureWriter) Status() int { @@ -742,19 +772,21 @@ func (w *opsCaptureWriter) Written() bool { return rw.Written() } func (w *opsCaptureWriter) Flush() { - state, rw := w.lockActive() + state, rw := w.beginDelegatedCall() if state == nil { return } - defer state.mu.RUnlock() + state.mu.Unlock() + defer finishDelegatedCall(state) rw.Flush() } func (w *opsCaptureWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { - state, rw := w.lockActive() + state, rw := w.beginDelegatedCall() if state == nil { return nil, nil, errors.New("response writer released") } - defer state.mu.RUnlock() + state.mu.Unlock() + defer finishDelegatedCall(state) return rw.Hijack() } func (w *opsCaptureWriter) CloseNotify() <-chan bool { @@ -777,26 +809,28 @@ func (w *opsCaptureWriter) Pusher() http.Pusher { } func (w *opsCaptureWriter) Write(b []byte) (int, error) { - state, rw := w.lockActiveWrite() + state, rw := w.beginDelegatedCall() if state == nil { return 0, nil } - defer state.mu.Unlock() if state.shouldCapture() { state.captureResponseChunk(b, rw.Status()) } + state.mu.Unlock() + defer finishDelegatedCall(state) return rw.Write(b) } func (w *opsCaptureWriter) WriteString(s string) (int, error) { - state, rw := w.lockActiveWrite() + state, rw := w.beginDelegatedCall() if state == nil { return 0, nil } - defer state.mu.Unlock() if state.shouldCapture() { state.captureResponseChunk([]byte(s), rw.Status()) } + state.mu.Unlock() + defer finishDelegatedCall(state) return rw.WriteString(s) } @@ -887,6 +921,13 @@ func (state *opsCaptureWriterState) captureResponseChunk(chunk []byte, status in state.appendCapturedResponse(chunk) return } + // Most stream writes contain one or more complete successful SSE frames. + // Skip the byte-wise frame parser when the chunk cannot contain a terminal + // event and leaves no split frame to carry into the next write. + if len(state.probe) == 0 && len(state.lineProbe) == 0 && endsAtOpsSSEFrameBoundary(chunk) && + !mayContainOpsTerminalSSE(chunk) { + return + } for i, b := range chunk { if state.skipLF { state.skipLF = false @@ -952,6 +993,20 @@ func (state *opsCaptureWriterState) captureResponseChunk(chunk []byte, status in } } +func endsAtOpsSSEFrameBoundary(chunk []byte) bool { + return bytes.HasSuffix(chunk, []byte("\n\n")) || + bytes.HasSuffix(chunk, []byte("\r\n\r\n")) || + bytes.HasSuffix(chunk, []byte("\r\r")) +} + +func mayContainOpsTerminalSSE(chunk []byte) bool { + if bytes.Contains(chunk, []byte("response.failed")) { + return true + } + return bytes.Contains(chunk, []byte("error")) && + (bytes.Contains(chunk, []byte("event")) || bytes.Contains(chunk, []byte(`"type"`))) +} + func isOpsTerminalSSEEventLine(line []byte) bool { line = bytes.TrimSpace(line) field, value, found := bytes.Cut(line, []byte{':'}) @@ -1066,10 +1121,14 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { 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) + // A marked in-band error is a visible request failure even though its + // wire status is already 200. Otherwise retain recovered attempts as a + // provider-health row whose 2xx status keeps it outside request SLA. + if len(service.GetOpsStreamErrors(c)) > 0 { + logOpsStreamError(c, ops, status) + } else { + logOpsRecoveredUpstream(c, ops, status) + } return } } @@ -1214,6 +1273,126 @@ func OpsErrorLoggerMiddleware(ops *service.OpsService) gin.HandlerFunc { } } +func logOpsRecoveredUpstream(c *gin.Context, ops *service.OpsService, finalStatus int) { + if c == nil || ops == nil || finalStatus >= 400 { + return + } + + entry := &service.OpsInsertErrorLogInput{StatusCode: finalStatus} + applyOpsUpstreamFieldsFromContext(c, entry) + if entry.UpstreamStatusCode == nil && entry.UpstreamErrorMessage == nil && + entry.UpstreamErrorDetail == nil && len(entry.UpstreamErrors) == 0 { + return + } + + lastStatus := 0 + if entry.UpstreamStatusCode != nil { + lastStatus = *entry.UpstreamStatusCode + } + lastStage := "" + for i := len(entry.UpstreamErrors) - 1; i >= 0; i-- { + if event := entry.UpstreamErrors[i]; event != nil { + lastStage = event.Stage + if event.AccountID > 0 { + accountID := event.AccountID + entry.AccountID = &accountID + } + break + } + } + if entry.AccountID == nil { + if accountID, ok := c.Get(opsAccountIDKey); ok { + if value, ok := accountID.(int64); ok && value > 0 { + entry.AccountID = &value + } + } + } + + entry.ErrorPhase = "upstream" + entry.ErrorType = "upstream_error" + entry.ErrorSource = "upstream_http" + entry.ErrorOwner = "provider" + entry.Severity = classifyOpsSeverity(entry.ErrorType, lastStatus) + entry.IsCountTokens = isCountTokensRequest(c) + entry.CreatedAt = time.Now() + entry.ErrorMessage = "Recovered upstream error" + if lastStage == string(service.GatewayFailureStageAccountAuth) { + entry.ErrorPhase = string(service.GatewayFailureStageAccountAuth) + entry.ErrorMessage = "Recovered account authentication failure" + } else if lastStatus > 0 { + entry.ErrorMessage += " " + strconv.Itoa(lastStatus) + } + if entry.UpstreamErrorMessage != nil && strings.TrimSpace(*entry.UpstreamErrorMessage) != "" { + entry.ErrorMessage += ": " + strings.TrimSpace(*entry.UpstreamErrorMessage) + } + entry.ErrorMessage = truncateString(entry.ErrorMessage, 2048) + + if c.Request != nil { + entry.UserAgent = c.GetHeader("User-Agent") + if c.Request.URL != nil { + entry.RequestPath = c.Request.URL.Path + } + if c.Request.Context() != nil { + entry.ClientRequestID, _ = c.Request.Context().Value(ctxkey.ClientRequestID).(string) + entry.RequestID, _ = c.Request.Context().Value(ctxkey.RequestID).(string) + } + } + entry.RequestID = strings.TrimSpace(entry.RequestID) + if entry.RequestID == "" { + entry.RequestID = c.Writer.Header().Get("X-Request-Id") + } + entry.Model = c.GetString(opsModelKey) + entry.RequestedModel = entry.Model + entry.Stream = c.GetBool(opsStreamKey) + entry.InboundEndpoint = GetInboundEndpoint(c) + entry.UpstreamModel = c.GetString(opsUpstreamModelKey) + entry.RequestType = opsRequestTypeFromContext(c) + + apiKey := getOpsAPIKey(c) + fallbackPlatform := guessPlatformFromPath(entry.RequestPath) + var requestContext context.Context = context.Background() + if c.Request != nil { + requestContext = c.Request.Context() + } + entry.Platform = resolveOpsPlatform(requestContext, apiKey, fallbackPlatform) + entry.UpstreamEndpoint = GetUpstreamEndpoint(c, entry.Platform) + 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 + } + if apiKey.Group != nil && apiKey.Group.Platform != "" { + entry.Platform = apiKey.Group.Platform + } + } + if clientIP := strings.TrimSpace(ip.GetClientIP(c)); clientIP != "" { + entry.ClientIP = &clientIP + } + applyOpsLatencyFieldsFromContext(c, entry) + enqueueOpsErrorLog(ops, entry) +} + +func opsRequestTypeFromContext(c *gin.Context) *int16 { + if c == nil { + return nil + } + if value, ok := c.Get(opsRequestTypeKey); ok { + switch typed := value.(type) { + case int16: + result := typed + return &result + case int: + result := int16(typed) + return &result + } + } + return nil +} + // logOpsStreamError 记录一次挂在已固化 HTTP 200 SSE 流上的就地错误。 // 由于 wire 状态码停留在 200,常规的 status>=400 捕获路径永远不会触发; // handleStreamingAwareError 通过 service.MarkOpsStreamError 标记这类错误, diff --git a/backend/internal/handler/ops_error_logger_test.go b/backend/internal/handler/ops_error_logger_test.go index 1956ad2776..38176cb1b0 100644 --- a/backend/internal/handler/ops_error_logger_test.go +++ b/backend/internal/handler/ops_error_logger_test.go @@ -38,15 +38,18 @@ func (r *ingressRejectSettingRepo) Set(context.Context, string, string) error { type ingressRejectOpsRepo struct { service.OpsRepository insertCalls int + entries []*service.OpsInsertErrorLogInput } -func (r *ingressRejectOpsRepo) InsertErrorLog(context.Context, *service.OpsInsertErrorLogInput) (int64, error) { +func (r *ingressRejectOpsRepo) InsertErrorLog(_ context.Context, entry *service.OpsInsertErrorLogInput) (int64, error) { r.insertCalls++ + r.entries = append(r.entries, entry) return 0, nil } -func (r *ingressRejectOpsRepo) BatchInsertErrorLogs(context.Context, []*service.OpsInsertErrorLogInput) (int64, error) { +func (r *ingressRejectOpsRepo) BatchInsertErrorLogs(_ context.Context, entries []*service.OpsInsertErrorLogInput) (int64, error) { r.insertCalls++ + r.entries = append(r.entries, entries...) return 0, nil } @@ -328,11 +331,12 @@ func TestOpsErrorLoggerMiddleware_OrdinaryPermissionStillRecords(t *testing.T) { require.Equal(t, http.StatusForbidden, job.entry.StatusCode) } -func TestOpsErrorLoggerMiddleware_SkipsRecoveredUpstreamErrorOnSuccessfulRequest(t *testing.T) { +func TestOpsErrorLoggerMiddleware_RecordsRecoveredUpstreamTelemetryOutsideFailureSLA(t *testing.T) { setupOpsErrorLogTestQueue(t, 2) gin.SetMode(gin.TestMode) - ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + repo := &ingressRejectOpsRepo{} + ops := service.NewOpsService(repo, 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) { @@ -347,7 +351,50 @@ func TestOpsErrorLoggerMiddleware_SkipsRecoveredUpstreamErrorOnSuccessfulRequest router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil)) require.Equal(t, http.StatusOK, recorder.Code) - require.Zero(t, OpsErrorLogQueueLength()) + require.Equal(t, int64(1), OpsErrorLogQueueLength()) + job := <-opsErrorLogQueue + require.Nil(t, job.entry.UpstreamErrors, "raw attempts must be released before async queueing") + require.NotNil(t, job.entry.UpstreamErrorsJSON) + queuedEvents, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON) + require.NoError(t, err) + require.Len(t, queuedEvents, 1) + require.Equal(t, http.StatusTooManyRequests, queuedEvents[0].UpstreamStatusCode) + + flushOpsErrorLogBatch([]opsErrorLogJob{job}) + require.Equal(t, 1, repo.insertCalls) + require.Len(t, repo.entries, 1) + persisted := repo.entries[0] + require.Equal(t, http.StatusOK, persisted.StatusCode, "recovered telemetry must remain outside failed-request SLA") + require.Equal(t, "upstream", persisted.ErrorPhase) + require.Equal(t, "upstream_error", persisted.ErrorType) + require.Equal(t, "Recovered upstream error 429: earlier attempt was rate limited", persisted.ErrorMessage) + require.NotNil(t, persisted.UpstreamErrorsJSON) + persistedEvents, err := service.ParseOpsUpstreamErrors(*persisted.UpstreamErrorsJSON) + require.NoError(t, err) + require.Len(t, persistedEvents, 1) + require.Equal(t, http.StatusTooManyRequests, persistedEvents[0].UpstreamStatusCode) +} + +func TestOpsErrorLoggerMiddleware_IntermediateSkipMonitoringDoesNotHideFinalVisibleFailure(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.StatusBadGateway, Message: "hidden retry", SkipMonitoring: true}, + {UpstreamStatusCode: http.StatusServiceUnavailable, Message: "visible final"}, + }) + c.JSON(http.StatusServiceUnavailable, gin.H{"error": gin.H{"type": "upstream_error", "message": "visible final"}}) + }) + + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil)) + require.Equal(t, int64(1), OpsErrorLogQueueLength()) + job := <-opsErrorLogQueue + require.Equal(t, http.StatusServiceUnavailable, job.entry.StatusCode) + require.Equal(t, "visible final", job.entry.ErrorMessage) } func TestOpsErrorLoggerMiddleware_CapturesSplitResponsesFailedSSE(t *testing.T) { diff --git a/backend/internal/repository/gateway_cache.go b/backend/internal/repository/gateway_cache.go index f67e92b263..36935160da 100644 --- a/backend/internal/repository/gateway_cache.go +++ b/backend/internal/repository/gateway_cache.go @@ -270,27 +270,47 @@ func (c *gatewayCache) GetReasoningContent(ctx context.Context, itemID string) ( } const ( - cyberSessionBlockPrefix = "cyber_session_block:" - cyberSessionScopePrefix = "cyber_session_scope:" + cyberSessionBlockPrefix = "cyber_session_block:" + cyberSessionScopePrefix = "cyber_session_scope:" + cyberSessionRedisCommandMaxKeys = 128 ) -// SetCyberSessionBlocked atomically writes all exact blocks and their optional -// coarse source scope with the same TTL. +// SetCyberSessionBlocked writes exact blocks in bounded transactions. The +// coarse scope is activated only after all exact blocks have been stored. func (c *gatewayCache) SetCyberSessionBlocked(ctx context.Context, scopeKey string, keys []string, ttl time.Duration) error { if len(keys) == 0 { return nil } - pipe := c.rdb.TxPipeline() - for _, key := range keys { - if key != "" { + exactKeys := make([]string, 0, cyberSessionRedisCommandMaxKeys) + flush := func() error { + if len(exactKeys) == 0 { + return nil + } + pipe := c.rdb.TxPipeline() + for _, key := range exactKeys { pipe.Set(ctx, cyberSessionBlockPrefix+key, "1", ttl) } + _, err := pipe.Exec(ctx) + exactKeys = exactKeys[:0] + return err + } + for _, key := range keys { + if key != "" { + exactKeys = append(exactKeys, key) + if len(exactKeys) == cyberSessionRedisCommandMaxKeys { + if err := flush(); err != nil { + return err + } + } + } + } + if err := flush(); err != nil { + return err } if scopeKey != "" { - pipe.Set(ctx, cyberSessionScopePrefix+scopeKey, "1", ttl) + return c.rdb.Set(ctx, cyberSessionScopePrefix+scopeKey, "1", ttl).Err() } - _, err := pipe.Exec(ctx) - return err + return nil } func (c *gatewayCache) IsCyberSessionScopeActive(ctx context.Context, scopeKey string) (bool, error) { @@ -301,23 +321,29 @@ func (c *gatewayCache) IsCyberSessionScopeActive(ctx context.Context, scopeKey s return n > 0, nil } -// FindCyberSessionBlocked checks transcript-prefix candidates in one Redis -// round trip and returns the first blocked key in caller order. +// FindCyberSessionBlocked checks bounded batches in caller order and stops at +// the first blocked key, preserving the original earliest-match behavior. func (c *gatewayCache) FindCyberSessionBlocked(ctx context.Context, keys []string) (string, error) { if len(keys) == 0 { return "", nil } - redisKeys := make([]string, len(keys)) - for i, key := range keys { - redisKeys[i] = cyberSessionBlockPrefix + key - } - values, err := c.rdb.MGet(ctx, redisKeys...).Result() - if err != nil { - return "", err - } - for i, value := range values { - if value != nil { - return keys[i], nil + for start := 0; start < len(keys); start += cyberSessionRedisCommandMaxKeys { + end := start + cyberSessionRedisCommandMaxKeys + if end > len(keys) { + end = len(keys) + } + redisKeys := make([]string, end-start) + for i, key := range keys[start:end] { + redisKeys[i] = cyberSessionBlockPrefix + key + } + values, err := c.rdb.MGet(ctx, redisKeys...).Result() + if err != nil { + return "", err + } + for i, value := range values { + if value != nil { + return keys[start+i], nil + } } } return "", nil diff --git a/backend/internal/repository/gateway_cache_cyber_test.go b/backend/internal/repository/gateway_cache_cyber_test.go index 40396e7c26..03f29e82a1 100644 --- a/backend/internal/repository/gateway_cache_cyber_test.go +++ b/backend/internal/repository/gateway_cache_cyber_test.go @@ -2,6 +2,8 @@ package repository import ( "context" + "strconv" + "sync" "testing" "time" @@ -11,6 +13,42 @@ import ( "github.com/stretchr/testify/require" ) +type cyberRedisCommandHook struct { + mu sync.Mutex + mgetKeyCounts []int + setBatchSizes []int +} + +func (h *cyberRedisCommandHook) DialHook(next redis.DialHook) redis.DialHook { return next } + +func (h *cyberRedisCommandHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { + return func(ctx context.Context, cmd redis.Cmder) error { + if cmd.Name() == "mget" { + h.mu.Lock() + h.mgetKeyCounts = append(h.mgetKeyCounts, len(cmd.Args())-1) + h.mu.Unlock() + } + return next(ctx, cmd) + } +} + +func (h *cyberRedisCommandHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook { + return func(ctx context.Context, cmds []redis.Cmder) error { + setCount := 0 + for _, cmd := range cmds { + if cmd.Name() == "set" { + setCount++ + } + } + if setCount > 0 { + h.mu.Lock() + h.setBatchSizes = append(h.setBatchSizes, setCount) + h.mu.Unlock() + } + return next(ctx, cmds) + } +} + func TestGatewayCacheCyberBlockWritesScopeAndExactKeysTogether(t *testing.T) { server := miniredis.RunT(t) client := redis.NewClient(&redis.Options{Addr: server.Addr()}) @@ -29,3 +67,31 @@ func TestGatewayCacheCyberBlockWritesScopeAndExactKeysTogether(t *testing.T) { require.Greater(t, server.TTL(cyberSessionScopePrefix+"scope-1"), time.Duration(0)) require.Equal(t, server.TTL(cyberSessionBlockPrefix+"block-1"), server.TTL(cyberSessionBlockPrefix+"block-2")) } + +func TestGatewayCacheCyberBlockCommandsAreBoundedAndLookupShortCircuits(t *testing.T) { + server := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: server.Addr()}) + t.Cleanup(func() { _ = client.Close() }) + hook := &cyberRedisCommandHook{} + client.AddHook(hook) + store, ok := NewGatewayCache(client).(service.CyberSessionBlockStore) + require.True(t, ok) + + keys := make([]string, cyberSessionRedisCommandMaxKeys*2+44) + for i := range keys { + keys[i] = "block-" + strconv.Itoa(i) + } + ctx := context.Background() + require.NoError(t, store.SetCyberSessionBlocked(ctx, "large-scope", keys, time.Minute)) + require.Equal(t, []int{cyberSessionRedisCommandMaxKeys, cyberSessionRedisCommandMaxKeys, 44}, hook.setBatchSizes) + + lookup := make([]string, len(keys)) + for i := range lookup { + lookup[i] = "missing-" + strconv.Itoa(i) + } + lookup[cyberSessionRedisCommandMaxKeys+3] = keys[cyberSessionRedisCommandMaxKeys+3] + matched, err := store.FindCyberSessionBlocked(ctx, lookup) + require.NoError(t, err) + require.Equal(t, keys[cyberSessionRedisCommandMaxKeys+3], matched) + require.Equal(t, []int{cyberSessionRedisCommandMaxKeys, cyberSessionRedisCommandMaxKeys}, hook.mgetKeyCounts) +} diff --git a/backend/internal/repository/temp_unsched_cache.go b/backend/internal/repository/temp_unsched_cache.go index c145a30476..42c2fd648c 100644 --- a/backend/internal/repository/temp_unsched_cache.go +++ b/backend/internal/repository/temp_unsched_cache.go @@ -141,8 +141,3 @@ func (c *tempUnschedCache) RecordOpenAIAPIKeyHealthFailure(ctx context.Context, } return count, tripped == 1, nil } - -func (c *tempUnschedCache) ResetOpenAIAPIKeyHealthFailures(ctx context.Context, accountID int64) error { - key := c.openAIAPIKeyHealthKey(accountID) - return c.rdb.Del(ctx, key, key+":sequence").Err() -} diff --git a/backend/internal/repository/temp_unsched_cache_health_test.go b/backend/internal/repository/temp_unsched_cache_health_test.go index 092b85ac6b..1b8a1f10c6 100644 --- a/backend/internal/repository/temp_unsched_cache_health_test.go +++ b/backend/internal/repository/temp_unsched_cache_health_test.go @@ -3,6 +3,7 @@ package repository import ( "context" "testing" + "time" "github.com/Wei-Shaw/sub2api/internal/service" "github.com/alicebob/miniredis/v2" @@ -10,7 +11,7 @@ import ( "github.com/stretchr/testify/require" ) -func TestOpenAIAPIKeyHealthCacheTripsAndSuccessResetsWindow(t *testing.T) { +func TestOpenAIAPIKeyHealthCacheTripsWithinRollingWindow(t *testing.T) { server := miniredis.RunT(t) client := redis.NewClient(&redis.Options{Addr: server.Addr()}) t.Cleanup(func() { _ = client.Close() }) @@ -18,14 +19,6 @@ func TestOpenAIAPIKeyHealthCacheTripsAndSuccessResetsWindow(t *testing.T) { require.True(t, ok) ctx := context.Background() - for attempt := 1; attempt <= 2; attempt++ { - count, tripped, err := store.RecordOpenAIAPIKeyHealthFailure(ctx, 42, 1, 3) - require.NoError(t, err) - require.EqualValues(t, attempt, count) - require.False(t, tripped) - } - require.NoError(t, store.ResetOpenAIAPIKeyHealthFailures(ctx, 42)) - for attempt := 1; attempt <= 3; attempt++ { count, tripped, err := store.RecordOpenAIAPIKeyHealthFailure(ctx, 42, 1, 3) require.NoError(t, err) @@ -33,3 +26,25 @@ func TestOpenAIAPIKeyHealthCacheTripsAndSuccessResetsWindow(t *testing.T) { require.Equal(t, attempt == 3, tripped) } } + +func TestOpenAIAPIKeyHealthCacheDropsFailuresOutsideRollingWindow(t *testing.T) { + server := miniredis.RunT(t) + now := time.Unix(1_700_000_000, 0) + server.SetTime(now) + client := redis.NewClient(&redis.Options{Addr: server.Addr()}) + t.Cleanup(func() { _ = client.Close() }) + store, ok := NewTempUnschedCache(client).(service.OpenAIAPIKeyHealthCache) + require.True(t, ok) + + ctx := context.Background() + count, tripped, err := store.RecordOpenAIAPIKeyHealthFailure(ctx, 42, 1, 3) + require.NoError(t, err) + require.EqualValues(t, 1, count) + require.False(t, tripped) + + server.SetTime(now.Add(61 * time.Second)) + count, tripped, err = store.RecordOpenAIAPIKeyHealthFailure(ctx, 42, 1, 3) + require.NoError(t, err) + require.EqualValues(t, 1, count) + require.False(t, tripped) +} diff --git a/backend/internal/service/openai_access_state_failover_test.go b/backend/internal/service/openai_access_state_failover_test.go index 0acaa26182..a46f9ca571 100644 --- a/backend/internal/service/openai_access_state_failover_test.go +++ b/backend/internal/service/openai_access_state_failover_test.go @@ -24,22 +24,59 @@ func (r *openAIStream403AccountRepo) SetError(context.Context, int64, string) er return nil } +type openAIAuthPolicyAccountRepo struct { + AccountRepository + tempCalls int + setErrorCalls int +} + +func (r *openAIAuthPolicyAccountRepo) SetTempUnschedulable(context.Context, int64, time.Time, string) error { + r.tempCalls++ + return nil +} + +func (r *openAIAuthPolicyAccountRepo) SetError(context.Context, int64, string) error { + r.setErrorCalls++ + return nil +} + +type openAIAuthPolicy403Counter struct { + counts []int64 +} + +func (s *openAIAuthPolicy403Counter) IncrementOpenAI403Count(context.Context, int64, int) (int64, error) { + if len(s.counts) == 0 { + return 1, nil + } + count := s.counts[0] + s.counts = s.counts[1:] + return count, nil +} + +func (*openAIAuthPolicy403Counter) ResetOpenAI403Count(context.Context, int64) error { + return nil +} + func TestOpenAIUpstreamAccessStateClassification(t *testing.T) { tests := []struct { name string body string + want bool }{ - {"workspace_code", `{"detail":{"code":"deactivated_workspace"}}`}, - {"disabled_account", `{"error":{"message":"Your account is disabled"}}`}, - {"suspended_workspace", `{"response":{"error":{"message":"This workspace has been suspended"}}}`}, - {"deactivated_organization", `{"detail":{"message":"The organization is deactivated"}}`}, - {"scalar_detail", `{"detail":"This workspace has been disabled"}`}, - {"suspended_org_code", `{"error":{"code":"org_suspended"}}`}, + {"workspace_code", `{"detail":{"code":"deactivated_workspace"}}`, true}, + {"disabled_account_message", `{"error":{"message":"Your account is disabled"}}`, false}, + {"suspended_workspace_message", `{"response":{"error":{"message":"This workspace has been suspended"}}}`, false}, + {"deactivated_organization_message", `{"detail":{"message":"The organization is deactivated"}}`, false}, + {"scalar_detail", `{"detail":"This workspace has been disabled"}`, false}, + {"suspended_org_code", `{"error":{"code":"org_suspended"}}`, true}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { body := []byte(tt.body) - require.True(t, isOpenAIUpstreamAccessStateError("", body)) + require.Equal(t, tt.want, isOpenAIUpstreamAccessStateError("", body)) + if !tt.want { + return + } require.True(t, (&OpenAIGatewayService{}).shouldFailoverOpenAIUpstreamResponse(http.StatusForbidden, "", body)) require.True(t, shouldFailoverOpenAIPassthroughResponse(&Account{Type: AccountTypeOAuth}, http.StatusForbidden, body)) @@ -62,6 +99,93 @@ func TestOpenAIUpstreamAccessStateDoesNotScanEchoedJSON(t *testing.T) { require.False(t, shouldFailoverOpenAIPassthroughResponse(&Account{Type: AccountTypeOAuth}, http.StatusBadRequest, body)) } +func TestOpenAIHTTPAccessStateDoesNotTrustBadRequestMessage(t *testing.T) { + body := []byte(`{"error":{"type":"invalid_request_error","code":"unknown_parameter","message":"Unknown parameter: account disabled"}}`) + svc := &OpenAIGatewayService{} + + require.False(t, isOpenAIUpstreamAccessStateError("", body), "free-form stream messages are not durable account evidence") + require.False(t, isOpenAIHTTPUpstreamAccessStateError(http.StatusBadRequest, "", body)) + require.False(t, svc.shouldFailoverOpenAIUpstreamResponse(http.StatusBadRequest, "", body)) + require.False(t, shouldFailoverOpenAIPassthroughResponse(&Account{Type: AccountTypeOAuth}, http.StatusBadRequest, body)) + + err := newOpenAIUpstreamFailoverError(http.StatusBadRequest, nil, body, "", false) + require.False(t, err.IsCredentialFailure()) +} + +func TestOpenAIHTTPAccessStateBadRequestDoesNotDisableAccount(t *testing.T) { + repo := &openAIStream403AccountRepo{} + svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{accountRepo: repo}} + account := &Account{ID: 925, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true} + body := []byte(`{"error":{"code":"unknown_parameter","message":"Unknown parameter: account disabled"}}`) + + disabled := svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusBadRequest, nil, body) + + require.False(t, disabled) + require.Zero(t, repo.setErrorCalls) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestOpenAIStreamEchoedAccessStateMessageDoesNotDisableOrFailover(t *testing.T) { + repo := &openAIStream403AccountRepo{} + svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{accountRepo: repo}} + account := &Account{ID: 926, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true} + payload := []byte(`{"type":"response.failed","response":{"error":{"type":"invalid_request_error","code":"unknown_parameter","message":"Unknown parameter: account disabled"}}}`) + message := extractOpenAISSEErrorMessage(payload) + + require.False(t, isOpenAIUpstreamAccessStateError(message, payload)) + require.False(t, openAIStreamFailedEventShouldFailover(payload, message)) + status, disabled := svc.handleOpenAIStreamTerminalAccountSideEffects(nil, account, payload, message, nil) + require.Equal(t, http.StatusBadGateway, status) + require.False(t, disabled) + require.Zero(t, repo.setErrorCalls) + require.False(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestOpenAIHTTPAccessStateTrustsStructuredCode(t *testing.T) { + repo := &openAIStream403AccountRepo{} + svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{accountRepo: repo}} + account := &Account{ID: 930, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true} + body := []byte(`{"error":{"code":"organization_deactivated","message":"request rejected"}}`) + + require.True(t, isOpenAIHTTPUpstreamAccessStateError(http.StatusBadRequest, "", body)) + require.True(t, svc.shouldFailoverOpenAIUpstreamResponse(http.StatusBadRequest, "", body)) + require.True(t, svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusBadRequest, nil, body)) + require.Equal(t, 1, repo.setErrorCalls) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) +} + +func TestOpenAIHTTPAuthMessagesUseExistingStatusPolicies(t *testing.T) { + t.Run("oauth 401 remains recoverable", func(t *testing.T) { + repo := &openAIAuthPolicyAccountRepo{} + rateLimits := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + svc := &OpenAIGatewayService{rateLimitService: rateLimits} + account := &Account{ID: 931, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true, + Credentials: map[string]any{"refresh_token": "refreshable"}} + body := []byte(`{"error":{"message":"account is disabled"}}`) + + require.False(t, isOpenAIHTTPUpstreamAccessStateError(http.StatusUnauthorized, "", body)) + require.True(t, svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusUnauthorized, nil, body)) + require.Zero(t, repo.setErrorCalls) + require.Equal(t, 1, repo.tempCalls) + }) + + t.Run("403 uses counter cooldown", func(t *testing.T) { + repo := &openAIAuthPolicyAccountRepo{} + counter := &openAIAuthPolicy403Counter{counts: []int64{1}} + rateLimits := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + rateLimits.openAI403CounterCache = counter + svc := &OpenAIGatewayService{rateLimitService: rateLimits} + rateLimits.SetAccountRuntimeBlocker(svc) + account := &Account{ID: 932, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Status: StatusActive, Schedulable: true} + body := []byte(`{"error":{"message":"workspace has been suspended"}}`) + + require.False(t, isOpenAIHTTPUpstreamAccessStateError(http.StatusForbidden, "", body)) + require.True(t, svc.handleOpenAIAccountUpstreamError(context.Background(), account, http.StatusForbidden, nil, body)) + require.Zero(t, repo.setErrorCalls) + require.Equal(t, 1, repo.tempCalls) + }) +} + func TestOpenAICyberPolicyWrapped5xxNeverFailsOver(t *testing.T) { body := []byte(`{"error":{"code":"cyber_policy","message":"blocked"}}`) svc := &OpenAIGatewayService{} diff --git a/backend/internal/service/openai_account_runtime_block_fastpath.go b/backend/internal/service/openai_account_runtime_block_fastpath.go index a25d78c697..1d8da95913 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath.go @@ -65,7 +65,7 @@ func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Cont } stateCtx, cancel := openAIAccountStateContext(ctx) defer cancel() - if account != nil && account.Platform == PlatformOpenAI && isOpenAIUpstreamAccessStateError("", responseBody) { + if account != nil && account.Platform == PlatformOpenAI && isOpenAIHTTPUpstreamAccessStateError(statusCode, "", responseBody) { message := "OpenAI upstream account or workspace is unavailable" if upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(responseBody)); upstreamMsg != "" { message = upstreamMsg diff --git a/backend/internal/service/openai_account_runtime_block_fastpath_test.go b/backend/internal/service/openai_account_runtime_block_fastpath_test.go index 97bfae0eee..db620cc5b8 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath_test.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath_test.go @@ -15,11 +15,13 @@ import ( type oauth429RateLimitRepo struct { AccountRepository - setRateLimitedCalls int + setRateLimitedCalls int + lastRateLimitedUntil time.Time } -func (r *oauth429RateLimitRepo) SetRateLimited(context.Context, int64, time.Time) error { +func (r *oauth429RateLimitRepo) SetRateLimited(_ context.Context, _ int64, until time.Time) error { r.setRateLimitedCalls++ + r.lastRateLimitedUntil = until return nil } @@ -60,6 +62,52 @@ func TestOpenAI429FastPath_BlocksOAuthOnlyAfterRetryWindow(t *testing.T) { require.False(t, svc.shouldRetryOpenAIOAuth429OnSameAccount(account, http.StatusTooManyRequests, false)) } +func TestOpenAIStream429IgnoresSuccessfulQuotaSnapshotHeaders(t *testing.T) { + repo := &oauth429RateLimitRepo{} + rateLimits := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + svc := &OpenAIGatewayService{rateLimitService: rateLimits} + rateLimits.SetAccountRuntimeBlocker(svc) + account := &Account{ID: 421, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + svc.openaiOAuth429RetryStartedAt.Store(account.ID, time.Now().Add(-openAIOAuth429RetryWindow-time.Second)) + headers := http.Header{} + headers.Set("x-codex-primary-used-percent", "37") + headers.Set("x-codex-primary-reset-after-seconds", "604800") + headers.Set("x-codex-primary-window-minutes", "10080") + payload := []byte(`{"type":"error","error":{"type":"rate_limit_error","code":"rate_limit_exceeded","message":"slow down"}}`) + + status, disabled := svc.handleOpenAIStreamTerminalAccountSideEffects(nil, account, payload, "slow down", headers) + + require.Equal(t, http.StatusTooManyRequests, status) + require.False(t, disabled) + require.True(t, svc.isOpenAIAccountRuntimeBlocked(account)) + value, ok := svc.openaiAccountRuntimeBlockUntil.Load(account.ID) + require.True(t, ok) + blockedUntil, ok := value.(time.Time) + require.True(t, ok) + require.Less(t, time.Until(blockedUntil), time.Minute, "stream 429 must not inherit the normal seven-day quota snapshot") + if !repo.lastRateLimitedUntil.IsZero() { + require.Less(t, time.Until(repo.lastRateLimitedUntil), time.Minute) + } +} + +func TestOpenAIHTTP429StillUsesQuotaResetHeaders(t *testing.T) { + svc := &OpenAIGatewayService{rateLimitService: &RateLimitService{}} + account := &Account{ID: 422, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + svc.openaiOAuth429RetryStartedAt.Store(account.ID, time.Now().Add(-openAIOAuth429RetryWindow-time.Second)) + headers := http.Header{} + headers.Set("x-codex-primary-used-percent", "37") + headers.Set("x-codex-primary-reset-after-seconds", "604800") + headers.Set("x-codex-primary-window-minutes", "10080") + + svc.markOpenAIOAuth429RateLimited(context.Background(), account, headers, nil) + + value, ok := svc.openaiAccountRuntimeBlockUntil.Load(account.ID) + require.True(t, ok) + blockedUntil, ok := value.(time.Time) + require.True(t, ok) + require.Greater(t, time.Until(blockedUntil), 6*24*time.Hour, "real HTTP 429 must retain the upstream quota reset") +} + func TestOpenAI429RetryDelayHonorsBoundedRetryAfter(t *testing.T) { deadline := time.Now().Add(openAIOAuth429RetryWindow) require.Equal(t, openAIOAuth429RetryDelay, openAIOAuth429SameAccountRetryDelay(nil, deadline)) diff --git a/backend/internal/service/openai_alpha_search.go b/backend/internal/service/openai_alpha_search.go index 1954dc2670..c22545f5f1 100644 --- a/backend/internal/service/openai_alpha_search.go +++ b/backend/internal/service/openai_alpha_search.go @@ -104,7 +104,7 @@ func (s *OpenAIGatewayService) ForwardAlphaSearch(ctx context.Context, c *gin.Co if account.IsOpenAIOAuthLike() && resp.StatusCode == http.StatusTooManyRequests { return nil, s.newOpenAIAccountFailoverError(account, resp.StatusCode, resp.Header, respBody, upstreamMessage, shouldDisable, retryableOnSameAccount) } - if isOpenAIUpstreamAccessStateError(upstreamMessage, respBody) { + if isOpenAIHTTPUpstreamAccessStateError(resp.StatusCode, upstreamMessage, respBody) { return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMessage, retryableOnSameAccount) } return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBody, RetryableOnSameAccount: retryableOnSameAccount} @@ -182,7 +182,7 @@ func (s *OpenAIGatewayService) forwardAlphaSearchViaResponsesWebSearch( if account.IsOpenAIOAuthLike() && resp.StatusCode == http.StatusTooManyRequests { return nil, s.newOpenAIAccountFailoverError(account, resp.StatusCode, resp.Header, respBody, upstreamMessage, shouldDisable, retryableOnSameAccount) } - if isOpenAIUpstreamAccessStateError(upstreamMessage, respBody) { + if isOpenAIHTTPUpstreamAccessStateError(resp.StatusCode, upstreamMessage, respBody) { return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMessage, retryableOnSameAccount) } return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBody, RetryableOnSameAccount: retryableOnSameAccount} diff --git a/backend/internal/service/openai_apikey_health_breaker.go b/backend/internal/service/openai_apikey_health_breaker.go index c898cf1124..65581ae3d2 100644 --- a/backend/internal/service/openai_apikey_health_breaker.go +++ b/backend/internal/service/openai_apikey_health_breaker.go @@ -124,15 +124,7 @@ func (s *RateLimitService) ObserveOpenAIAPIKeyHealthFailure(ctx context.Context, return true } -func (s *RateLimitService) ObserveOpenAIAPIKeyHealthSuccess(ctx context.Context, account *Account) { - if s == nil || s.openAIAPIKeyHealth == nil || s.settingService == nil || !isOpenAIAPIKeyHealthBreakerAccount(account) { - return - } - settings, err := s.settingService.GetOpenAIAPIKeyHealthBreakerSettings(ctx) - if err != nil || settings == nil || !settings.Enabled { - return - } - if err := s.openAIAPIKeyHealth.ResetOpenAIAPIKeyHealthFailures(ctx, account.ID); err != nil { - logger.L().Warn("openai.apikey_health_breaker_reset_failed", zap.Int64("account_id", account.ID), zap.Error(err)) - } +func (s *RateLimitService) ObserveOpenAIAPIKeyHealthSuccess(context.Context, *Account) { + // Health failures are accumulated in a rolling time window. A success does + // not reset that window and must not add a Redis round trip to the hot path. } diff --git a/backend/internal/service/openai_apikey_health_breaker_test.go b/backend/internal/service/openai_apikey_health_breaker_test.go index a6e7aaad00..4c58ed1b4a 100644 --- a/backend/internal/service/openai_apikey_health_breaker_test.go +++ b/backend/internal/service/openai_apikey_health_breaker_test.go @@ -13,10 +13,12 @@ import ( type openAIAPIKeyHealthSettingRepo struct { SettingRepository - value string + value string + getCalls int } func (r *openAIAPIKeyHealthSettingRepo) GetValue(context.Context, string) (string, error) { + r.getCalls++ return r.value, nil } @@ -35,7 +37,6 @@ func (r *openAIAPIKeyHealthAccountRepo) SetTempUnschedulable(_ context.Context, type openAIAPIKeyHealthCacheStub struct { TempUnschedCache recordCalls int - resetCalls int setCalls int tripped bool } @@ -45,11 +46,6 @@ func (c *openAIAPIKeyHealthCacheStub) RecordOpenAIAPIKeyHealthFailure(context.Co return 3, c.tripped, nil } -func (c *openAIAPIKeyHealthCacheStub) ResetOpenAIAPIKeyHealthFailures(context.Context, int64) error { - c.resetCalls++ - return nil -} - func (c *openAIAPIKeyHealthCacheStub) SetTempUnsched(context.Context, int64, *TempUnschedState) error { c.setCalls++ return nil @@ -127,10 +123,11 @@ func TestOpenAIAPIKeyHealthBreakerTripsPersistedAndRuntimeState(t *testing.T) { require.Contains(t, repo.reason, openAIAPIKeyHealthBreakerReason) } -func TestOpenAIAPIKeyHealthSuccessResetsOnlyEligiblePoolAccount(t *testing.T) { +func TestOpenAIAPIKeyHealthSuccessDoesNotTouchSettingsOrCache(t *testing.T) { encoded, err := json.Marshal(OpenAIAPIKeyHealthBreakerSettings{Enabled: true, WindowMinutes: 1, FailureThreshold: 3, CooldownMinutes: 5}) require.NoError(t, err) - settings := NewSettingService(&openAIAPIKeyHealthSettingRepo{value: string(encoded)}, &config.Config{}) + settingRepo := &openAIAPIKeyHealthSettingRepo{value: string(encoded)} + settings := NewSettingService(settingRepo, &config.Config{}) cache := &openAIAPIKeyHealthCacheStub{} svc := NewRateLimitService(&openAIAPIKeyHealthAccountRepo{}, nil, &config.Config{}, nil, cache) svc.SetSettingService(settings) @@ -138,5 +135,6 @@ func TestOpenAIAPIKeyHealthSuccessResetsOnlyEligiblePoolAccount(t *testing.T) { svc.ObserveOpenAIAPIKeyHealthSuccess(context.Background(), openAIHealthPoolAccount()) svc.ObserveOpenAIAPIKeyHealthSuccess(context.Background(), &Account{ID: 43, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}) - require.Equal(t, 1, cache.resetCalls) + require.Zero(t, settingRepo.getCalls) + require.Zero(t, cache.recordCalls) } 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 1d798eec5e..18b2df3153 100644 --- a/backend/internal/service/openai_codex_function_call_id_test.go +++ b/backend/internal/service/openai_codex_function_call_id_test.go @@ -151,6 +151,28 @@ func TestFilterCodexInput_ExistingItemIDWinsOverLegacyCallIDMapping(t *testing.T require.Equal(t, "call_shared", filtered[2].(map[string]any)["id"]) } +func TestFilterCodexInput_NormalizesCrossTurnLegacyCallReference(t *testing.T) { + input := []any{ + map[string]any{"type": "item_reference", "id": "call_previous_turn"}, + } + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true}) + + require.Equal(t, "fc_previous_turn", filtered[0].(map[string]any)["id"]) +} + +func TestFilterCodexInput_PreservesNativeRemoteItemReferences(t *testing.T) { + for _, id := range []string{"fc_remote", "ctc_remote", "tsc_remote", "msg_remote", "rs_remote", "vendor_remote"} { + t.Run(id, func(t *testing.T) { + input := []any{map[string]any{"type": "item_reference", "id": id}} + + filtered := filterCodexInputWithOptions(input, codexInputFilterOptions{PreserveReferences: true}) + + require.Equal(t, id, filtered[0].(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_transform.go b/backend/internal/service/openai_codex_transform.go index bd24f29c66..fbd14ba967 100644 --- a/backend/internal/service/openai_codex_transform.go +++ b/backend/internal/service/openai_codex_transform.go @@ -1562,10 +1562,25 @@ func codexInputItemIDs(input []any) map[string]struct{} { return itemIDs } +func codexInputCallIDs(input []any) map[string]struct{} { + callIDs := make(map[string]struct{}) + for _, rawItem := range input { + item, ok := rawItem.(map[string]any) + if !ok || !isCodexToolCallItemType(strings.TrimSpace(firstNonEmptyString(item["type"]))) { + continue + } + if callID := strings.TrimSpace(firstNonEmptyString(item["call_id"])); callID != "" { + callIDs[callID] = struct{}{} + } + } + return callIDs +} + func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []any { filtered := make([]any, 0, len(input)) referenceIDMappings := codexItemReferenceIDMappings(input, opts.PreserveCallIDs) inputItemIDs := codexInputItemIDs(input) + inputCallIDs := codexInputCallIDs(input) for _, item := range input { m, ok := item.(map[string]any) if !ok { @@ -1629,8 +1644,14 @@ func filterCodexInputWithOptions(input []any, opts codexInputFilterOptions) []an 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 + if !referencesExistingItem { + if normalizedID, mapped := referenceIDMappings[trimmedID]; mapped { + newItem["id"] = normalizedID + } else if _, hasSameTurnCall := inputCallIDs[trimmedID]; !hasSameTurnCall { + // A bare call_* reference is a legacy function-call identifier. + // Normalize it even when its call item lives in an earlier turn. + newItem["id"] = normalizeCodexCallID(trimmedID) + } } } filtered = append(filtered, newItem) diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go index 56be9570b0..7895b03cc3 100644 --- a/backend/internal/service/openai_codex_transform_test.go +++ b/backend/internal/service/openai_codex_transform_test.go @@ -127,6 +127,26 @@ func TestApplyCodexOAuthTransform_ToolContinuationNormalizesToolReferenceIDsOnly require.Equal(t, "fc_1", second["call_id"]) } +func TestApplyCodexOAuthTransform_NormalizesIsolatedLegacyReferenceAcrossTurns(t *testing.T) { + reqBody := map[string]any{ + "model": "gpt-5.2", + "input": []any{ + map[string]any{"type": "item_reference", "id": "call_previous_turn"}, + map[string]any{"type": "item_reference", "id": "fc_remote_item"}, + map[string]any{"type": "item_reference", "id": "vendor_remote_item"}, + }, + } + + applyCodexOAuthTransform(reqBody, false, false) + + input, ok := reqBody["input"].([]any) + require.True(t, ok) + require.Len(t, input, 3) + require.Equal(t, "fc_previous_turn", input[0].(map[string]any)["id"]) + require.Equal(t, "fc_remote_item", input[1].(map[string]any)["id"]) + require.Equal(t, "vendor_remote_item", input[2].(map[string]any)["id"]) +} + func TestApplyCodexOAuthTransform_BoundsLongCallIDsAndPreservesPairing(t *testing.T) { suffix := strings.Repeat("z", 62) for _, tc := range []struct { diff --git a/backend/internal/service/openai_compact_fallback.go b/backend/internal/service/openai_compact_fallback.go index 517a7e2e7b..389a4e8beb 100644 --- a/backend/internal/service/openai_compact_fallback.go +++ b/backend/internal/service/openai_compact_fallback.go @@ -1,7 +1,10 @@ package service import ( + "bytes" + "encoding/json" "errors" + "io" "net/http" "strings" @@ -89,11 +92,7 @@ func isOpenAICompactModelFailure(statusCode int, upstreamMsg string, upstreamBod case "model_not_found", "model_not_available", "unsupported_model", "invalid_model": return true } - if strings.Contains(value, "model") && (strings.Contains(value, "not found") || - strings.Contains(value, "does not exist") || - strings.Contains(value, "unavailable") || - strings.Contains(value, "unsupported") || - strings.Contains(value, "not supported")) { + if isExplicitOpenAIModelAvailabilityMessage(value) { return true } } @@ -115,6 +114,121 @@ func isOpenAICompactModelFailure(statusCode int, upstreamMsg string, upstreamBod return false } +func isExplicitOpenAIModelAvailabilityMessage(value string) bool { + value = strings.ToLower(strings.TrimSpace(value)) + if value == "" { + return false + } + for _, phrase := range []string{ + "model not found", + "model does not exist", + "model is unavailable", + "model is not available", + "model is unsupported", + "model is not supported", + "unsupported model", + } { + if strings.Contains(value, phrase) { + return true + } + } + // OpenAI commonly identifies the missing model between the word "model" + // and the terminal availability phrase, for example: "The model `x` does + // not exist". Requiring the message to start with the model subject avoids + // treating unrelated feature errors such as "model output is not supported" + // as a signal to change models. + if strings.HasPrefix(value, "the model ") || strings.HasPrefix(value, "model ") { + return strings.Contains(value, " does not exist") || + strings.Contains(value, " was not found") || + strings.Contains(value, " is unavailable") || + strings.Contains(value, " is not available") + } + return false +} + +func openAICompactFallbackErrorResponse(resp *http.Response, signal *openAICompactFallbackSignal) (*http.Response, []byte) { + headers := make(http.Header) + if resp != nil { + headers = resp.Header.Clone() + } + if headers.Get("Content-Type") == "" { + headers.Set("Content-Type", "application/json") + } + payload := normalizeOpenAICompactFallbackHTTPErrorPayload(signal) + return &http.Response{ + StatusCode: http.StatusBadRequest, + Header: headers, + Body: io.NopCloser(bytes.NewReader(payload)), + }, payload +} + +func normalizeOpenAICompactFallbackHTTPErrorPayload(signal *openAICompactFallbackSignal) []byte { + if signal == nil { + return nil + } + payload := append([]byte(nil), signal.payload...) + var terminal struct { + Error json.RawMessage `json:"error"` + Response struct { + Error json.RawMessage `json:"error"` + } `json:"response"` + } + if json.Unmarshal(payload, &terminal) != nil || len(bytes.TrimSpace(terminal.Response.Error)) == 0 || + bytes.Equal(bytes.TrimSpace(terminal.Response.Error), []byte("null")) { + return payload + } + // Standard HTTP error handlers consume error.message/type/code. A streamed + // response.failed terminal nests the same object under response.error, so + // normalize only that envelope at the stream-to-HTTP boundary. + normalized, err := json.Marshal(struct { + Error json.RawMessage `json:"error"` + }{Error: terminal.Response.Error}) + if err != nil { + return payload + } + return normalized +} + +func (s *OpenAIGatewayService) appendOpenAICompactFallbackRetryOps( + c *gin.Context, + account *Account, + resp *http.Response, + payload []byte, + message string, + passthrough bool, +) { + if account == nil { + return + } + statusCode := http.StatusBadRequest + requestID := "" + if resp != nil { + statusCode = resp.StatusCode + requestID = resp.Header.Get("x-request-id") + } + detail := "" + if s != nil && s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody { + maxBytes := s.cfg.Gateway.LogUpstreamErrorBodyMaxBytes + if maxBytes <= 0 { + maxBytes = 2048 + } + detail = truncateString(string(payload), maxBytes) + } + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: statusCode, + UpstreamRequestID: requestID, + Passthrough: passthrough, + Kind: "retry", + Reason: "compact_model_fallback", + Message: sanitizeUpstreamErrorMessage(strings.TrimSpace(message)), + Detail: detail, + UpstreamResponseBody: detail, + }) +} + // prepareOpenAICompactFallbackRetry returns a body for one safe, same-account // retry. Callers invoke it only before any downstream response has been // written; it changes the model and deliberately leaves path, trigger, and @@ -164,6 +278,7 @@ func (s *OpenAIGatewayService) applyOpenAIPassthroughCompactFallbackFromSignal( if !retry { return body, "", false } + s.appendOpenAICompactFallbackRetryOps(c, account, resp, signal.payload, signal.message, true) if resp != nil && resp.Body != nil { _ = resp.Body.Close() } diff --git a/backend/internal/service/openai_compact_fallback_test.go b/backend/internal/service/openai_compact_fallback_test.go index 944f7e1cfd..2e4c7050eb 100644 --- a/backend/internal/service/openai_compact_fallback_test.go +++ b/backend/internal/service/openai_compact_fallback_test.go @@ -3,11 +3,14 @@ package service import ( "bytes" "context" + "errors" "io" "net/http" "net/http/httptest" + "strconv" "strings" "testing" + "time" "github.com/Wei-Shaw/sub2api/internal/config" "github.com/gin-gonic/gin" @@ -138,6 +141,29 @@ func TestPrepareOpenAICompactFallbackRetryDoesNotHideSpecificBusinessFailure(t * require.Equal(t, body, retryBody) } +func TestIsOpenAICompactModelFailureRequiresExplicitModelAvailabilityMessage(t *testing.T) { + tests := []struct { + name string + message string + want bool + }{ + {name: "explicit unsupported model", message: "The requested model is not supported", want: true}, + {name: "named missing model", message: "The model `gpt-5.5` does not exist", want: true}, + {name: "unsupported model code-like message", message: "unsupported model: gpt-5.5", want: true}, + {name: "unsupported model feature", message: "This model output format is not supported", want: false}, + {name: "unsupported parameter for model", message: "Parameter tools is not supported for this model", want: false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, isOpenAICompactModelFailure( + http.StatusBadRequest, + tt.message, + []byte(`{"error":{"message":`+strconv.Quote(tt.message)+`}}`), + )) + }) + } +} + func TestPrepareOpenAICompactFallbackRetrySkipsSameModel(t *testing.T) { gin.SetMode(gin.TestMode) svc := &OpenAIGatewayService{cfg: &config.Config{Gateway: config.GatewayConfig{OpenAICompactModel: "gpt-5.5"}}} @@ -191,6 +217,14 @@ func TestOpenAIGatewayForwardRetriesExplicitNativeCompactHTTPFailureOnce(t *test require.True(t, HasCompactionTriggerInInput(upstream.bodies[1])) require.Equal(t, upstream.requests[0].URL.Path, upstream.requests[1].URL.Path) require.NotContains(t, upstream.requests[1].URL.Path, "/compact") + rawEvents, ok := c.Get(OpsUpstreamErrorsKey) + require.True(t, ok) + events, ok := rawEvents.([]*OpsUpstreamErrorEvent) + require.True(t, ok) + require.Len(t, events, 1) + require.Equal(t, "retry", events[0].Kind) + require.Equal(t, "compact_model_fallback", events[0].Reason) + require.Equal(t, http.StatusBadRequest, events[0].UpstreamStatusCode) } func TestOpenAIGatewayForwardRetriesExplicitNativeCompactSSEFailureBeforeOutput(t *testing.T) { @@ -309,4 +343,63 @@ func TestOpenAIGatewayForwardDoesNotRecurseWhenCompactFallbackAlsoFails(t *testi require.Len(t, upstream.bodies, 2) require.Equal(t, "gpt-5.5", gjson.GetBytes(upstream.bodies[0], "model").String()) require.Equal(t, "gpt-5.4", gjson.GetBytes(upstream.bodies[1], "model").String()) + var compactSignal *openAICompactFallbackSignal + require.False(t, errors.As(err, &compactSignal)) + require.Equal(t, http.StatusBadRequest, recorder.Code) + require.Contains(t, recorder.Body.String(), "model not found") + rawEvents, ok := c.Get(OpsUpstreamErrorsKey) + require.True(t, ok) + events, ok := rawEvents.([]*OpsUpstreamErrorEvent) + require.True(t, ok) + require.Len(t, events, 2) + require.Equal(t, "retry", events[0].Kind) + require.Equal(t, "compact_model_fallback", events[0].Reason) + require.Equal(t, "http_error", events[1].Kind) +} + +func TestOpenAIPassthroughCompactFallbackSecondStreamFailureUsesStandardErrorPath(t *testing.T) { + gin.SetMode(gin.TestMode) + body := []byte(`{"model":"gpt-5.5","stream":true,"input":[{"type":"compaction_trigger"}]}`) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + c.Request.Header.Set("Content-Type", "application/json") + MarkOpenAINativeCompactionV2(c) + + failed := "event: response.failed\n" + + `data: {"type":"response.failed","response":{"status":"failed","error":{"code":"context_length_exceeded","message":"context window exceeded"}}}` + "\n\n" + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(failed))}, + {StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: io.NopCloser(strings.NewReader(failed))}, + }} + svc := &OpenAIGatewayService{ + cfg: &config.Config{Gateway: config.GatewayConfig{OpenAICompactModel: "gpt-5.4"}}, + httpUpstream: upstream, + } + account := &Account{ + ID: 1, Name: "openai-oauth", Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1, + Credentials: map[string]any{"access_token": "oauth-token", "chatgpt_account_id": "chatgpt-account"}, + Status: StatusActive, Schedulable: true, + } + + result, err := svc.forwardOpenAIPassthrough( + context.Background(), c, account, body, body, "gpt-5.5", false, nil, true, time.Now(), + ) + + require.Error(t, err) + require.Nil(t, result) + require.Len(t, upstream.bodies, 2) + var compactSignal *openAICompactFallbackSignal + require.False(t, errors.As(err, &compactSignal)) + require.Equal(t, http.StatusBadRequest, recorder.Code) + require.Contains(t, recorder.Body.String(), "context window exceeded") + rawEvents, ok := c.Get(OpsUpstreamErrorsKey) + require.True(t, ok) + events, ok := rawEvents.([]*OpsUpstreamErrorEvent) + require.True(t, ok) + require.Len(t, events, 2) + require.Equal(t, "retry", events[0].Kind) + require.Equal(t, "compact_model_fallback", events[0].Reason) + require.Equal(t, "http_error", events[1].Kind) + require.True(t, events[1].Passthrough) } diff --git a/backend/internal/service/openai_cyber_session_block.go b/backend/internal/service/openai_cyber_session_block.go index 797e746561..44c3ece444 100644 --- a/backend/internal/service/openai_cyber_session_block.go +++ b/backend/internal/service/openai_cyber_session_block.go @@ -21,6 +21,8 @@ type CyberSessionBlockStore interface { FindCyberSessionBlocked(ctx context.Context, keys []string) (string, error) } +const cyberSessionTranscriptLookupOverflowBlockKey = "transcript_lookup_limit_exceeded" + // CyberSessionExplicitBlockKey returns an inexpensive exact key when the // client supplies a stable session signal. func CyberSessionExplicitBlockKey(apiKeyID int64, c *gin.Context, body []byte) string { @@ -142,7 +144,13 @@ func (s *OpenAIGatewayService) FindCyberSessionBlockedForRequest(ctx context.Con if !active { return "" } - keys := CyberSessionTranscriptLookupKeys(apiKeyID, body) + transcript := deriveOpenAICyberTranscriptBlockKeys(apiKeyID, body) + if transcript.lookupKeysTruncated { + // Once the coarse scope is active, silently dropping old candidates would + // let a blocked client evade prefix matching by appending dummy items. + return cyberSessionTranscriptLookupOverflowBlockKey + } + keys := transcript.lookupKeys if len(keys) == 0 { return "" } diff --git a/backend/internal/service/openai_cyber_session_block_test.go b/backend/internal/service/openai_cyber_session_block_test.go index 7080450042..aff46df25d 100644 --- a/backend/internal/service/openai_cyber_session_block_test.go +++ b/backend/internal/service/openai_cyber_session_block_test.go @@ -2,8 +2,10 @@ package service import ( "context" + "encoding/json" "errors" "net/http/httptest" + "strconv" "strings" "testing" "time" @@ -75,11 +77,32 @@ func TestCyberTranscriptBlockKeysWebSocketResponseCreate(t *testing.T) { require.Len(t, CyberSessionTranscriptBlockKeys(88, body), 2) } +func TestCyberTranscriptLookupKeysAreBoundedAndKeepNewestOrder(t *testing.T) { + messages := make([]map[string]string, maxOpenAICyberTranscriptLookupKeys+44) + for i := range messages { + messages[i] = map[string]string{"role": "user", "content": "message-" + strconv.Itoa(i)} + } + body, err := json.Marshal(map[string]any{"messages": messages}) + require.NoError(t, err) + + keys := CyberSessionTranscriptLookupKeys(77, body) + require.Len(t, keys, maxOpenAICyberTranscriptLookupKeys) + + firstRetainedBody, err := json.Marshal(map[string]any{"messages": messages[:45]}) + require.NoError(t, err) + firstRetainedPrefix := CyberSessionTranscriptLookupKeys(77, firstRetainedBody) + require.Equal(t, firstRetainedPrefix[len(firstRetainedPrefix)-1], keys[0]) + + fullKey := CyberSessionTranscriptBlockKeys(77, body)[0] + require.Equal(t, fullKey, keys[len(keys)-1]) +} + // --- fakes --- type fakeCyberBlockStore struct { - blocked map[string]bool - scopes map[string]bool + blocked map[string]bool + scopes map[string]bool + findCalls int } var _ CyberSessionBlockStore = (*fakeCyberBlockStore)(nil) @@ -105,6 +128,7 @@ func (f *fakeCyberBlockStore) IsCyberSessionScopeActive(_ context.Context, scope } func (f *fakeCyberBlockStore) FindCyberSessionBlocked(_ context.Context, keys []string) (string, error) { + f.findCalls++ for _, key := range keys { if f.blocked[key] { return key, nil @@ -271,6 +295,34 @@ func TestFindCyberSessionBlockedForRequestUsesScopeForTranscript(t *testing.T) { require.Equal(t, blockKey, svc.FindCyberSessionBlockedForRequest(ctx, 9, nextCtx, nextBody, clientIP, "Codex CLI 1.2.4")) } +func TestFindCyberSessionBlockedForRequestFailsClosedOnScopedTranscriptOverflow(t *testing.T) { + settingSvc := &SettingService{settingRepo: &fakeSettingRepo{vals: map[string]string{ + SettingKeyCyberSessionBlockEnabled: "true", + SettingKeyCyberSessionBlockTTLSeconds: "60", + }}} + combo := &comboCacheAndStore{} + svc := &OpenAIGatewayService{cache: combo, settingService: settingSvc} + ctx := context.Background() + const apiKeyID = int64(9) + const clientIP = "203.0.113.20" + const userAgent = "Codex CLI 1.2.3" + + messages := make([]map[string]string, maxOpenAICyberTranscriptLookupKeys+1) + for i := range messages { + messages[i] = map[string]string{"role": "user", "content": "message-" + strconv.Itoa(i)} + } + body, err := json.Marshal(map[string]any{"messages": messages}) + require.NoError(t, err) + c, _ := newCyberBlockTestCtx(nil, string(body)) + require.Empty(t, svc.FindCyberSessionBlockedForRequest(ctx, apiKeyID, c, body, clientIP, userAgent), + "overflow alone must not bypass the scope gate") + combo.store.scopes = map[string]bool{CyberSessionScopeKey(apiKeyID, clientIP, userAgent): true} + + require.Equal(t, cyberSessionTranscriptLookupOverflowBlockKey, + svc.FindCyberSessionBlockedForRequest(ctx, apiKeyID, c, body, clientIP, userAgent)) + require.Zero(t, combo.store.findCalls, "overflow must not issue an unbounded Redis lookup") +} + func TestCyberSessionScopeKeyNormalizesUserAgentVersion(t *testing.T) { base := CyberSessionScopeKey(7, "203.0.113.10", "Codex CLI 1.2.3") require.NotEmpty(t, base) diff --git a/backend/internal/service/openai_cyber_transcript.go b/backend/internal/service/openai_cyber_transcript.go index 1314d54c51..b286c04769 100644 --- a/backend/internal/service/openai_cyber_transcript.go +++ b/backend/internal/service/openai_cyber_transcript.go @@ -11,10 +11,15 @@ import ( ) type openAICyberTranscriptBlockKeys struct { - lookupKeys []string - preLatestUserKey string + lookupKeys []string + preLatestUserKey string + lookupKeysTruncated bool } +// Bound the Redis lookup work for a single request while retaining the most +// recent transcript prefixes, where a continuation is most likely to match. +const maxOpenAICyberTranscriptLookupKeys = 256 + // deriveOpenAICyberTranscriptBlockKeys returns cumulative semantic-history // hashes plus the context key immediately before the latest user turn. The // context key requires model-generated history so shared first-turn templates @@ -53,8 +58,11 @@ func deriveOpenAICyberTranscriptBlockKeys(apiKeyID int64, body []byte) openAICyb return openAICyberTranscriptBlockKeys{} } result := openAICyberTranscriptBlockKeys{ - lookupKeys: make([]string, 0, int(sequence.Get("#").Int())), + lookupKeys: make([]string, 0, maxOpenAICyberTranscriptLookupKeys), } + nextLookupKey := 0 + lookupKeysRotated := false + lastLookupKey := "" // This is an entropy heuristic, not provenance proof: authenticated // server-side history would be required to distinguish fixed few-shot // assistant items perfectly. @@ -71,17 +79,31 @@ func deriveOpenAICyberTranscriptBlockKeys(apiKeyID int64, body []byte) openAICyb if strings.TrimSpace(canonical) == "" { return true } - if openAICyberTranscriptItemStartsUserTurn(item) && hasModelGeneratedItem && len(result.lookupKeys) > 0 { - result.preLatestUserKey = result.lookupKeys[len(result.lookupKeys)-1] + if openAICyberTranscriptItemStartsUserTurn(item) && hasModelGeneratedItem && lastLookupKey != "" { + result.preLatestUserKey = lastLookupKey } _, _ = h.Write([]byte("|item=")) _, _ = h.Write([]byte(canonical)) - result.lookupKeys = append(result.lookupKeys, hex.EncodeToString(h.Sum(nil))) + lastLookupKey = hex.EncodeToString(h.Sum(nil)) + if len(result.lookupKeys) < maxOpenAICyberTranscriptLookupKeys { + result.lookupKeys = append(result.lookupKeys, lastLookupKey) + } else { + result.lookupKeys[nextLookupKey] = lastLookupKey + nextLookupKey = (nextLookupKey + 1) % maxOpenAICyberTranscriptLookupKeys + lookupKeysRotated = true + result.lookupKeysTruncated = true + } if openAICyberTranscriptItemIsModelGenerated(item) { hasModelGeneratedItem = true } return true }) + if lookupKeysRotated { + ordered := make([]string, 0, len(result.lookupKeys)) + ordered = append(ordered, result.lookupKeys[nextLookupKey:]...) + ordered = append(ordered, result.lookupKeys[:nextLookupKey]...) + result.lookupKeys = ordered + } return result } diff --git a/backend/internal/service/openai_embeddings.go b/backend/internal/service/openai_embeddings.go index ef0c88b2d7..366e11a22f 100644 --- a/backend/internal/service/openai_embeddings.go +++ b/backend/internal/service/openai_embeddings.go @@ -140,7 +140,7 @@ func (s *OpenAIGatewayService) ForwardEmbeddings( if account.IsOpenAIOAuth() && resp.StatusCode == http.StatusTooManyRequests { return nil, s.newOpenAIAccountFailoverError(account, resp.StatusCode, resp.Header, respBody, upstreamMsg, shouldDisable, retryableOnSameAccount) } - if isOpenAIUpstreamAccessStateError(upstreamMsg, respBody) { + if isOpenAIHTTPUpstreamAccessStateError(resp.StatusCode, upstreamMsg, respBody) { return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, retryableOnSameAccount) } return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBody, RetryableOnSameAccount: retryableOnSameAccount} 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 9605a0acee..e3a93bebb3 100644 --- a/backend/internal/service/openai_gateway_apikey_item_id_test.go +++ b/backend/internal/service/openai_gateway_apikey_item_id_test.go @@ -258,7 +258,7 @@ func TestSanitizeOpenAIResponsesInputItemIDs_AllocationGrowthIsLinear(t *testing "10x more input items must not cause quadratic whole-body allocation growth") } -func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyClosesInvalidIDReferences(t *testing.T) { +func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyPreservesOpaqueReferences(t *testing.T) { body := []byte(`{"type":"response.create","input":[ {"type":"custom_tool_call","id":"ctc_call","call_id":"call_custom","name":"apply_patch","input":"patch"}, {"type":"custom_tool_call_output","id":"ctco_bad","call_id":"call_custom","output":"done"}, @@ -275,11 +275,12 @@ func TestNormalizeOpenAIResponsesWebSocketCompatibilityBodyClosesInvalidIDRefere require.NoError(t, err) require.True(t, changed) - require.Len(t, gjson.GetBytes(normalized, "input").Array(), 3) + require.Len(t, gjson.GetBytes(normalized, "input").Array(), 4) require.Equal(t, "ctc_call", gjson.GetBytes(normalized, "input.0.id").String()) require.Equal(t, "call_custom", gjson.GetBytes(normalized, "input.1.call_id").String()) require.False(t, gjson.GetBytes(normalized, "input.1.id").Exists()) - require.Equal(t, "item_future", gjson.GetBytes(normalized, "input.2.id").String()) + require.Equal(t, "ctco_bad", gjson.GetBytes(normalized, "input.2.id").String()) + require.Equal(t, "item_future", gjson.GetBytes(normalized, "input.3.id").String()) second, changedAgain, err := normalizeOpenAIResponsesWebSocketCompatibilityBody(normalized, &Account{ Platform: PlatformOpenAI, diff --git a/backend/internal/service/openai_gateway_forward.go b/backend/internal/service/openai_gateway_forward.go index 1d2c9f5f39..99d89a7fac 100644 --- a/backend/internal/service/openai_gateway_forward.go +++ b/backend/internal/service/openai_gateway_forward.go @@ -59,21 +59,10 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco // 在分流到 passthrough / Codex transform / 原生 ChatCompletions 之前统一修正 // 显式为 null 的工具 Schema type,否则 upstream 的 400 会被归一成可重试的 502, // 同一份坏定义在账号池里反复重放。 - if shouldSanitizeOpenAIResponsesToolSchemas(account.Platform) { - sanitizedToolBody, toolSchemaSanitized, toolSchemaErr := sanitizeOpenAIResponsesToolParameterTypes(body) - if toolSchemaErr != nil { - return nil, fmt.Errorf("sanitize OpenAI Responses tool parameters: %w", toolSchemaErr) - } - 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 sanitizedToolBody, toolSchemaSanitized, toolSchemaErr := sanitizeOpenAIResponsesToolSchemasForPlatform(body, account.Platform); toolSchemaErr != nil { + return nil, toolSchemaErr + } else if toolSchemaSanitized { + body = sanitizedToolBody } if account.IsOpenAIOAuthLike() { reasoningBody, reasoningChanged, reasoningErr := normalizeOpenAIResponsesReasoningMode(body) @@ -1015,6 +1004,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if retryBody, fallbackModel, retry := s.prepareOpenAICompactFallbackRetry( c, account, requestedModel, body, resp.StatusCode, upstreamMsg, respBody, compactModelFallbackRetried, ); retry { + s.appendOpenAICompactFallbackRetryOps(c, account, resp, respBody, upstreamMsg, false) fromModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) body = retryBody requestView = newOpenAIRequestView(body) @@ -1082,6 +1072,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco if retryBody, fallbackModel, retry := s.prepareOpenAICompactFallbackRetry( c, account, requestedModel, body, http.StatusBadRequest, signal.message, signal.payload, compactModelFallbackRetried, ); retry { + s.appendOpenAICompactFallbackRetryOps(c, account, resp, signal.payload, signal.message, false) body = retryBody requestView = newOpenAIRequestView(body) upstreamModel = fallbackModel @@ -1089,6 +1080,27 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco SetOpsUpstreamModel(c, fallbackModel) continue } + if resp.Body != nil { + _ = resp.Body.Close() + } + compactResp, compactBody := openAICompactFallbackErrorResponse(resp, signal) + if s.shouldFailoverOpenAIUpstreamResponse(compactResp.StatusCode, signal.message, compactBody) { + appendOpsUpstreamError(c, OpsUpstreamErrorEvent{ + Platform: account.Platform, + AccountID: account.ID, + AccountName: account.Name, + UpstreamStatusCode: compactResp.StatusCode, + UpstreamRequestID: compactResp.Header.Get("x-request-id"), + Kind: "failover", + Message: signal.message, + }) + shouldDisable := s.handleFailoverSideEffects(ctx, compactResp, account, compactBody, upstreamModel) + return nil, s.newOpenAIAccountFailoverError( + account, compactResp.StatusCode, compactResp.Header, compactBody, signal.message, shouldDisable, + !shouldDisable && account.IsPoolMode() && (account.IsPoolModeRetryableStatus(compactResp.StatusCode) || isOpenAITransientProcessingError(compactResp.StatusCode, signal.message, compactBody)), + ) + } + return s.handleErrorResponse(ctx, compactResp, c, account, body, resolveOpenAIErrorSchedulingModel(billingModel, upstreamModel)) } return nil, err } diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index 06653094ce..f5c8cfe106 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -377,6 +377,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( if retryBody, fallbackModel, retry := s.prepareOpenAICompactFallbackRetry( c, account, requestedModel, body, resp.StatusCode, upstreamMsg, probeBody, compactModelFallbackRetried, ); retry { + s.appendOpenAICompactFallbackRetryOps(c, account, resp, probeBody, upstreamMsg, true) fromModel := strings.TrimSpace(gjson.GetBytes(body, "model").String()) body = retryBody upstreamPassthroughModel = fallbackModel @@ -425,6 +426,14 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( compactModelFallbackRetried = true continue } + if signal, ok := asOpenAICompactFallbackSignal(handleErr); ok { + _ = resp.Body.Close() + compactResp, compactBody := openAICompactFallbackErrorResponse(resp, signal) + if shouldFailoverOpenAIPassthroughResponse(account, compactResp.StatusCode, compactBody) { + return nil, s.handleFailoverErrorResponsePassthrough(ctx, compactResp, c, account, body, compactBody) + } + return nil, s.handleErrorResponsePassthrough(ctx, compactResp, c, account, body, compactBody) + } _ = resp.Body.Close() return nil, handleErr } @@ -444,6 +453,14 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( compactModelFallbackRetried = true continue } + if signal, ok := asOpenAICompactFallbackSignal(handleErr); ok { + _ = resp.Body.Close() + compactResp, compactBody := openAICompactFallbackErrorResponse(resp, signal) + if shouldFailoverOpenAIPassthroughResponse(account, compactResp.StatusCode, compactBody) { + return nil, s.handleFailoverErrorResponsePassthrough(ctx, compactResp, c, account, body, compactBody) + } + return nil, s.handleErrorResponsePassthrough(ctx, compactResp, c, account, body, compactBody) + } _ = resp.Body.Close() return nil, handleErr } @@ -713,7 +730,7 @@ func shouldFailoverOpenAIPassthroughResponse(account *Account, statusCode int, r if isOpenAIContextWindowError("", responseBody) { return false } - if isOpenAIUpstreamAccessStateError("", responseBody) { + if isOpenAIHTTPUpstreamAccessStateError(statusCode, "", responseBody) { return true } if isOpenAIRequestBodyTooLargeError(statusCode, "", responseBody) { @@ -1520,7 +1537,14 @@ func (s *OpenAIGatewayService) handleOpenAIStreamTerminalAccountSideEffects( if c != nil && c.Request != nil { ctx = c.Request.Context() } - return statusCode, s.handleOpenAIAccountUpstreamError(ctx, account, statusCode, headers, payload) + accountHeaders := headers + if statusCode == http.StatusTooManyRequests { + // The enclosing HTTP response succeeded. Its quota snapshot describes + // normal account state and must not become the reset for a semantic 429 + // carried by a stream terminal event. + accountHeaders = nil + } + return statusCode, s.handleOpenAIAccountUpstreamError(ctx, account, statusCode, accountHeaders, payload) default: return statusCode, false } diff --git a/backend/internal/service/openai_gateway_request_body.go b/backend/internal/service/openai_gateway_request_body.go index e07e041ced..79ee4c4284 100644 --- a/backend/internal/service/openai_gateway_request_body.go +++ b/backend/internal/service/openai_gateway_request_body.go @@ -1033,17 +1033,11 @@ func normalizeOpenAIResponsesWebSocketCompatibilityBody(body []byte, account *Ac changed = true } } - if account != nil && shouldSanitizeOpenAIResponsesToolSchemas(account.Platform) { - 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 + if account != nil { + if schemaBody, schemaChanged, err := sanitizeOpenAIResponsesToolSchemasForPlatform(normalized, account.Platform); err != nil { + return body, false, fmt.Errorf("normalize websocket tool schemas: %w", err) + } else if schemaChanged { + normalized = schemaBody changed = true } } diff --git a/backend/internal/service/openai_gateway_response_flush_test.go b/backend/internal/service/openai_gateway_response_flush_test.go index 16483567ad..fe6e3ea6b2 100644 --- a/backend/internal/service/openai_gateway_response_flush_test.go +++ b/backend/internal/service/openai_gateway_response_flush_test.go @@ -441,6 +441,62 @@ func TestOpenAIResponseFlush_FailedAndErrorEventsFlushAtBoundaries(t *testing.T) }) } +func TestOpenAIResponseFlush_BareErrorFollowedByCompletedUsesCompletedTerminal(t *testing.T) { + body := "data: {\"type\":\"error\",\"error\":{\"code\":\"transient\",\"message\":\"retrying\"}}\n\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_recovered\",\"status\":\"completed\",\"output\":[],\"usage\":{\"input_tokens\":7,\"output_tokens\":3}}}\n\n" + recorder := newOpenAIResponseFlushRecorder() + + result, err := runOpenAIResponseFlushTest(recorder, io.NopCloser(strings.NewReader(body)), config.GatewayConfig{}) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, 7, result.usage.InputTokens) + require.Equal(t, 3, result.usage.OutputTokens) + gotBody, _ := recorder.snapshot() + require.NotContains(t, gotBody, `"type":"error"`) + require.NotContains(t, gotBody, `"type":"response.failed"`) + require.Contains(t, gotBody, `"type":"response.completed"`) +} + +func TestOpenAIResponseFlush_CompatibleAPIKeyDoesNotUseCodexBareErrorSynthesis(t *testing.T) { + body := "data: {\"type\":\"error\",\"error\":{\"code\":\"provider_error\",\"message\":\"provider failed\"}}\n\n" + recorder := newOpenAIResponseFlushRecorder() + account := &Account{ID: 2, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + + result, err := runOpenAIResponseFlushTestWithAccount(recorder, io.NopCloser(strings.NewReader(body)), config.GatewayConfig{}, account) + + require.Error(t, err) + require.NotNil(t, result) + gotBody, _ := recorder.snapshot() + require.Contains(t, gotBody, `"type":"error"`) + require.NotContains(t, gotBody, `"type":"response.failed"`) +} + +func TestOpenAIResponseFlush_RecentBareErrorAllowsCompletedBeforeIdleTimeout(t *testing.T) { + reader, writer := io.Pipe() + defer func() { _ = writer.Close() }() + recorder := newOpenAIResponseFlushRecorder() + resultCh, errCh := runOpenAIResponseFlushTestAsync(recorder, reader, config.GatewayConfig{StreamDataIntervalTimeout: 1}) + + // Place the bare error shortly before the first ticker firing. It is fresh + // data, so the ticker must leave the stream open for an authoritative event. + time.Sleep(700 * time.Millisecond) + _, err := io.WriteString(writer, "data: {\"type\":\"error\",\"error\":{\"code\":\"transient\",\"message\":\"retrying\"}}\n\n") + require.NoError(t, err) + time.Sleep(500 * time.Millisecond) + _, err = io.WriteString(writer, "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_late\",\"status\":\"completed\",\"output\":[],\"usage\":{\"input_tokens\":5,\"output_tokens\":1}}}\n\n") + require.NoError(t, err) + require.NoError(t, writer.Close()) + + require.NoError(t, <-errCh) + result := <-resultCh + require.NotNil(t, result) + require.Equal(t, 5, result.usage.InputTokens) + gotBody, _ := recorder.snapshot() + require.Contains(t, gotBody, `"type":"response.completed"`) + require.NotContains(t, gotBody, `"type":"response.failed"`) +} + func TestOpenAIResponseFlush_BareErrorTimeoutSynthesizesFailed(t *testing.T) { tests := []struct { name string @@ -534,6 +590,10 @@ func TestOpenAIResponseFlush_ClientDisconnectStillDrainsUsage(t *testing.T) { } func runOpenAIResponseFlushTest(recorder *openAIResponseFlushRecorder, body io.ReadCloser, gatewayCfg config.GatewayConfig) (*openaiStreamingResult, error) { + return runOpenAIResponseFlushTestWithAccount(recorder, body, gatewayCfg, &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth}) +} + +func runOpenAIResponseFlushTestWithAccount(recorder *openAIResponseFlushRecorder, body io.ReadCloser, gatewayCfg config.GatewayConfig, account *Account) (*openaiStreamingResult, error) { gin.SetMode(gin.TestMode) c, _ := gin.CreateTestContext(recorder) c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) @@ -546,7 +606,7 @@ func runOpenAIResponseFlushTest(recorder *openAIResponseFlushRecorder, body io.R Header: http.Header{"Content-Type": []string{"text/event-stream"}}, Body: body, } - return svc.handleStreamingResponse(context.Background(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI}, time.Now(), "gpt-5", "gpt-5") + return svc.handleStreamingResponse(context.Background(), resp, c, account, time.Now(), "gpt-5", "gpt-5") } func runOpenAIResponseFlushTestAsync(recorder *openAIResponseFlushRecorder, body io.ReadCloser, gatewayCfg config.GatewayConfig) (<-chan *openaiStreamingResult, <-chan error) { diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index 45d08e86f3..2675ebb995 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -251,7 +251,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. capacityFailoverSuppressedLogged := false failedMessage := "" clientOutputStarted := false - codexFailureTerminal := account != nil && account.Platform == PlatformOpenAI + codexFailureTerminal := account != nil && account.IsOpenAIOAuthLike() upstreamRequestID := strings.TrimSpace(resp.Header.Get("x-request-id")) var streamEarlyErr error terminalFailurePending := false @@ -478,6 +478,18 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. if data, ok := extractOpenAISSEDataLine(line); ok { dataBytes := []byte(data) eventType := effectiveOpenAISSEEventType(dataBytes, pendingSSEEventType) + if codexFailureTerminal && sawBareError && !sawResponseFailed && + (eventType == "response.completed" || eventType == "response.done") { + // A later successful terminal is authoritative over a pending bare + // error. Keep its usage and terminal visible to the client. + sawBareError = false + sawFailedEvent = false + terminalFailurePending = false + suppressCurrentEvent = false + bareErrorPayload = nil + bareErrorAccountSideEffectsPending = false + failedMessage = "" + } if codexFailureTerminal && sawBareError && !sawResponseFailed && eventType != "response.failed" { suppressCurrentEvent = true } @@ -862,10 +874,6 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. } case <-intervalCh: - if codexFailureTerminal && sawBareError && !sawResponseFailed { - _ = resp.Body.Close() - return finalizeStream() - } if failureDelivered { return resultWithUsage(), fmt.Errorf("upstream response failed: %s", failedMessage) } @@ -873,6 +881,10 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. if time.Since(lastRead) < streamInterval { continue } + if codexFailureTerminal && sawBareError && !sawResponseFailed { + _ = resp.Body.Close() + return finalizeStream() + } if clientDisconnected { return resultWithUsage(), fmt.Errorf("stream usage incomplete after timeout") } @@ -899,7 +911,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context. stopFirstOutputTimer() continue } - if codexFailureTerminal && sawBareError && !sawResponseFailed { + if codexFailureTerminal && sawBareError && !sawResponseFailed && len(events) == 0 { _ = resp.Body.Close() return finalizeStream() } @@ -1181,6 +1193,12 @@ func (s *OpenAIGatewayService) parseSSEUsageBytesWithType(data []byte, eventType if usage == nil || len(data) == 0 || bytes.Equal(bytes.TrimSpace(data), []byte("[DONE]")) { return } + // Usage is absent from nearly every delta event. Avoid full JSON validation + // and four gjson path scans on that hot path while retaining progressive + // usage from compatible upstreams on any event type. + if !bytes.Contains(data, []byte(`"usage"`)) { + return + } parsedUsage, ok := extractOpenAIUsageFromJSONBytes(data) if !ok { return @@ -1368,6 +1386,57 @@ func extractOpenAIResponseIDFromJSONBytes(body []byte) string { return strings.TrimSpace(gjson.GetBytes(body, "response.id").String()) } +const openAIHTTPResponseOwnerContextKey = "openai_http_response_owner" + +type openAIHTTPResponseOwner struct { + userID int64 + apiKeyID int64 +} + +// SetOpenAIHTTPResponseOwner marks the authenticated downstream owner whose +// successful Responses IDs may be used for later HTTP continuations. +func SetOpenAIHTTPResponseOwner(c *gin.Context, userID, apiKeyID int64) { + if c == nil || userID <= 0 || apiKeyID <= 0 { + return + } + c.Set(openAIHTTPResponseOwnerContextKey, openAIHTTPResponseOwner{userID: userID, apiKeyID: apiKeyID}) +} + +// ValidateOpenAIHTTPResponseOwner authorizes a continuation by downstream +// tenant. API key identity is retained in the binding, while keys owned by the +// same user remain interoperable. +func (s *OpenAIGatewayService) ValidateOpenAIHTTPResponseOwner( + ctx context.Context, + groupID int64, + responseID string, + userID, apiKeyID int64, +) (bool, error) { + if s == nil || strings.TrimSpace(responseID) == "" || userID <= 0 || apiKeyID <= 0 { + return false, nil + } + ownerUserID, ownerAPIKeyID, found, err := s.getOpenAIWSStateStore().GetHTTPResponseOwner(ctx, groupID, responseID) + if err != nil || !found { + return false, err + } + return ownerUserID == userID || (ownerUserID <= 0 && ownerAPIKeyID == apiKeyID), nil +} + +// BindOpenAIHTTPResponseOwner records an HTTP continuation owner independently +// from the upstream account selected for that response. +func (s *OpenAIGatewayService) BindOpenAIHTTPResponseOwner( + ctx context.Context, + groupID int64, + responseID string, + userID, apiKeyID int64, +) error { + if s == nil { + return nil + } + return s.getOpenAIWSStateStore().BindHTTPResponseOwner( + ctx, groupID, responseID, userID, apiKeyID, s.openAIWSResponseStickyTTL(), + ) +} + func (s *OpenAIGatewayService) bindHTTPResponseAccount(ctx context.Context, c *gin.Context, account *Account, responseID string) { if s == nil || account == nil || account.ID <= 0 { return @@ -1383,6 +1452,21 @@ func (s *OpenAIGatewayService) bindHTTPResponseAccount(ctx context.Context, c *g groupID := getOpenAIGroupIDFromContext(c) ttl := s.openAIWSResponseStickyTTL() logOpenAIWSBindResponseAccountWarn(groupID, account.ID, responseID, store.BindResponseAccount(ctx, groupID, responseID, account.ID, ttl)) + if rawOwner, ok := c.Get(openAIHTTPResponseOwnerContextKey); ok { + if owner, ok := rawOwner.(openAIHTTPResponseOwner); ok && owner.userID > 0 && owner.apiKeyID > 0 { + if err := s.BindOpenAIHTTPResponseOwner(ctx, groupID, responseID, owner.userID, owner.apiKeyID); err != nil { + logger.L().Warn( + "openai.http_bind_response_owner_failed", + zap.Int64("group_id", groupID), + zap.Int64("account_id", account.ID), + zap.Int64("user_id", owner.userID), + zap.Int64("api_key_id", owner.apiKeyID), + zap.String("response_id", truncateOpenAIWSLogValue(responseID, openAIWSIDValueMaxLen)), + zap.Error(err), + ) + } + } + } } func openAIUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) { @@ -1665,9 +1749,6 @@ func extractOpenAISSETerminalEvent(body string) (string, []byte, bool) { var terminalType string var terminalPayload []byte forEachOpenAISSEFrame(body, func(eventType string, data []byte) { - if terminalPayload != nil { - return - } switch eventType { case "response.completed", "response.done", "response.failed", "response.incomplete", "response.cancelled", "response.canceled", "error": terminalType = eventType diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 122a542964..1e41298cc1 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -536,6 +536,7 @@ func TestOpenAIGatewayService_BindHTTPResponseAccount(t *testing.T) { c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) groupID := int64(4201) c.Set("api_key", &APIKey{ID: 501, GroupID: &groupID}) + SetOpenAIHTTPResponseOwner(c, 601, 501) svc := &OpenAIGatewayService{} account := &Account{ID: 37001, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} @@ -544,6 +545,22 @@ func TestOpenAIGatewayService_BindHTTPResponseAccount(t *testing.T) { got, err := svc.getOpenAIWSStateStore().GetResponseAccount(context.Background(), groupID, "resp_http_001") require.NoError(t, err) require.Equal(t, account.ID, got) + + owned, err := svc.ValidateOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_http_001", 601, 501) + require.NoError(t, err) + require.True(t, owned) + + owned, err = svc.ValidateOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_http_001", 601, 502) + require.NoError(t, err) + require.True(t, owned, "API keys owned by the same downstream user remain interoperable") + + owned, err = svc.ValidateOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_http_001", 602, 501) + require.NoError(t, err) + require.False(t, owned) + + owned, err = svc.ValidateOpenAIHTTPResponseOwner(context.Background(), groupID, "resp_unknown", 601, 501) + require.NoError(t, err) + require.False(t, owned) } func TestOpenAIGatewayService_GenerateExplicitSessionHash_SkipsContentFallback(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_upstream_errors.go b/backend/internal/service/openai_gateway_upstream_errors.go index 62add58e24..4a4e9ab7fc 100644 --- a/backend/internal/service/openai_gateway_upstream_errors.go +++ b/backend/internal/service/openai_gateway_upstream_errors.go @@ -260,7 +260,7 @@ func (s *OpenAIGatewayService) shouldFailoverOpenAIUpstreamResponse(statusCode i if isOpenAIContextWindowError(upstreamMsg, upstreamBody) { return false } - if isOpenAIUpstreamAccessStateError(upstreamMsg, upstreamBody) { + if isOpenAIHTTPUpstreamAccessStateError(statusCode, upstreamMsg, upstreamBody) { return true } if isOpenAIRequestBodyTooLargeError(statusCode, upstreamMsg, upstreamBody) { @@ -306,7 +306,7 @@ func newOpenAIUpstreamFailoverError( failoverErr.ClientStatusCode = http.StatusRequestEntityTooLarge failoverErr.ClientMessage = OpenAIRequestBodyTooLargeClientMessage } - if isOpenAIUpstreamAccessStateError(upstreamMsg, responseBody) { + if isOpenAIHTTPUpstreamAccessStateError(statusCode, upstreamMsg, responseBody) { failoverErr.RetryableOnSameAccount = false failoverErr.RequestScopedTransient = false failoverErr.Stage = GatewayFailureStageAccountAuth @@ -359,59 +359,42 @@ const ( ) // isOpenAIUpstreamAccessStateError recognizes provider-side credential state -// failures from explicit structured fields. Valid JSON is never scanned as a -// blob because it may contain echoed user input with the same words. -func isOpenAIUpstreamAccessStateError(upstreamMsg string, body []byte) bool { - matchCode := func(value string) bool { - value = strings.ToLower(strings.TrimSpace(value)) - if value == "deactivated_workspace" { - return true - } - for _, subject := range []string{"workspace", "account", "organization", "org"} { - for _, state := range []string{"deactivated", "disabled", "suspended"} { - if value == subject+"_"+state || value == state+"_"+subject { - return true - } - } - } +// failures only from explicit structured codes. Free-form messages may contain +// echoed user input, including inside stream terminal error.message fields. +func isOpenAIUpstreamAccessStateError(_ string, body []byte) bool { + if len(body) == 0 || !gjson.ValidBytes(body) { return false } - matchMessage := func(value string) bool { - value = strings.ToLower(strings.TrimSpace(value)) - for _, subject := range []string{"workspace", "account", "organization", "org"} { - for _, state := range []string{"deactivated", "disabled", "suspended"} { - if strings.Contains(value, subject+" is "+state) || - strings.Contains(value, subject+" has been "+state) || - strings.Contains(value, subject+" "+state) { - return true - } - } - } - return false - } - - if matchMessage(upstreamMsg) { - return true - } - if len(body) == 0 { - return false - } - if !gjson.ValidBytes(body) { - return matchMessage(string(body)) || matchCode(string(body)) - } for _, path := range []string{"error.code", "response.error.code", "detail.code", "code"} { - if matchCode(gjson.GetBytes(body, path).String()) { - return true - } - } - for _, path := range []string{"error.message", "response.error.message", "detail.message", "detail", "message"} { - if matchMessage(gjson.GetBytes(body, path).String()) { + if isOpenAIUpstreamAccessStateCode(gjson.GetBytes(body, path).String()) { return true } } return false } +func isOpenAIUpstreamAccessStateCode(value string) bool { + value = strings.ToLower(strings.TrimSpace(value)) + if value == "deactivated_workspace" { + return true + } + for _, subject := range []string{"workspace", "account", "organization", "org"} { + for _, state := range []string{"deactivated", "disabled", "suspended"} { + if value == subject+"_"+state || value == state+"_"+subject { + return true + } + } + } + return false +} + +// isOpenAIHTTPUpstreamAccessStateError is deliberately status-independent: +// known provider codes are durable evidence, while 401/403 messages without +// such a code must flow through the existing authentication/403 policies. +func isOpenAIHTTPUpstreamAccessStateError(_ int, _ string, body []byte) bool { + return isOpenAIUpstreamAccessStateError("", body) +} + func openAICapacityShedClientMessage(upstreamMsg string, body []byte) string { for _, candidate := range []string{ upstreamMsg, diff --git a/backend/internal/service/openai_images.go b/backend/internal/service/openai_images.go index c29af50016..08f94fca67 100644 --- a/backend/internal/service/openai_images.go +++ b/backend/internal/service/openai_images.go @@ -663,7 +663,7 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesAPIKey( if account.IsOpenAIOAuthLike() && resp.StatusCode == http.StatusTooManyRequests { return nil, s.newOpenAIAccountFailoverError(account, resp.StatusCode, resp.Header, respBody, upstreamMsg, shouldDisable, retryableOnSameAccount) } - if isOpenAIUpstreamAccessStateError(upstreamMsg, respBody) { + if isOpenAIHTTPUpstreamAccessStateError(resp.StatusCode, upstreamMsg, respBody) { return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, retryableOnSameAccount) } return nil, &UpstreamFailoverError{StatusCode: resp.StatusCode, ResponseBody: respBody, RetryableOnSameAccount: retryableOnSameAccount} diff --git a/backend/internal/service/openai_response_terminal_usage_compat_test.go b/backend/internal/service/openai_response_terminal_usage_compat_test.go index f6dc899e09..5a24f0addb 100644 --- a/backend/internal/service/openai_response_terminal_usage_compat_test.go +++ b/backend/internal/service/openai_response_terminal_usage_compat_test.go @@ -9,6 +9,7 @@ import ( "github.com/Wei-Shaw/sub2api/internal/pkg/apicompat" "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" ) func TestEffectiveOpenAISSEEventTypePrefersPayload(t *testing.T) { @@ -31,6 +32,17 @@ func TestExtractOpenAISSETerminalEventUsesEventField(t *testing.T) { require.Equal(t, "provider failed", extractOpenAISSEErrorMessage(payload)) } +func TestExtractOpenAISSETerminalEventUsesFinalAuthoritativeTerminal(t *testing.T) { + t.Parallel() + + body := "data: {\"type\":\"error\",\"error\":{\"message\":\"recovering\"}}\n\n" + + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_1\",\"status\":\"completed\"}}\n\n" + eventType, payload, ok := extractOpenAISSETerminalEvent(body) + require.True(t, ok) + require.Equal(t, "response.completed", eventType) + require.Equal(t, "resp_1", gjson.GetBytes(payload, "response.id").String()) +} + func TestParseSSEUsageEffectiveTerminalRules(t *testing.T) { t.Parallel() @@ -47,6 +59,17 @@ func TestParseSSEUsageEffectiveTerminalRules(t *testing.T) { require.Equal(t, OpenAIUsage{InputTokens: 2}, *usage) } +func BenchmarkParseSSEUsageNoUsageDelta(b *testing.B) { + svc := &OpenAIGatewayService{} + usage := &OpenAIUsage{} + payload := []byte(`{"type":"response.output_text.delta","delta":"hello"}`) + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + svc.parseSSEUsageBytesWithType(payload, "response.output_text.delta", usage) + } +} + func TestOpenAICompatTerminalResponseSynthesizesBareError(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/openai_responses_ingress_compat.go b/backend/internal/service/openai_responses_ingress_compat.go index c9ae11c0ae..f3162b3b02 100644 --- a/backend/internal/service/openai_responses_ingress_compat.go +++ b/backend/internal/service/openai_responses_ingress_compat.go @@ -46,11 +46,17 @@ func normalizeOpenAIResponsesLegacyIngress(body []byte) ([]byte, bool, error) { } if prompt, hasPrompt := request["prompt"]; hasPrompt { - if input, hasInput := request["input"]; (!hasInput || input == nil) && prompt != nil { - request["input"] = prompt + // Only the legacy string alias is unambiguously equivalent to Responses + // input. Objects are native reusable prompt templates and must remain in + // prompt; arrays and other shapes are left for upstream validation rather + // than being relabeled as a structurally different input value. + if promptText, isLegacyString := prompt.(string); isLegacyString { + if input, hasInput := request["input"]; !hasInput || input == nil { + request["input"] = promptText + } + delete(request, "prompt") + changed = true } - delete(request, "prompt") - changed = true } if _, hasCommands := request["commands"]; hasCommands { delete(request, "commands") diff --git a/backend/internal/service/openai_responses_ingress_compat_test.go b/backend/internal/service/openai_responses_ingress_compat_test.go index 3e898f03e6..acd2c202c0 100644 --- a/backend/internal/service/openai_responses_ingress_compat_test.go +++ b/backend/internal/service/openai_responses_ingress_compat_test.go @@ -70,3 +70,25 @@ func TestNormalizeOpenAIResponsesLegacyIngressKeepsPromptAliasAndDropsCommands(t require.False(t, gjson.GetBytes(normalized, "prompt").Exists()) require.False(t, gjson.GetBytes(normalized, "commands").Exists()) } + +func TestNormalizeOpenAIResponsesLegacyIngressPreservesNativePromptTemplate(t *testing.T) { + body := []byte(`{"model":"gpt-5.4","prompt":{"id":"pmpt_abc","version":"7","variables":{"topic":"ownership"}}}`) + + normalized, changed, err := normalizeOpenAIResponsesLegacyIngress(body) + require.NoError(t, err) + require.False(t, changed) + require.JSONEq(t, string(body), string(normalized)) + require.Equal(t, "pmpt_abc", gjson.GetBytes(normalized, "prompt.id").String()) + require.False(t, gjson.GetBytes(normalized, "input").Exists()) +} + +func TestNormalizeOpenAIResponsesLegacyIngressPreservesUnknownPromptShapeWhileDroppingCommands(t *testing.T) { + body := []byte(`{"model":"gpt-5.4","prompt":["one","two"],"commands":[{"name":"legacy"}]}`) + + normalized, changed, err := normalizeOpenAIResponsesLegacyIngress(body) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, int64(2), gjson.GetBytes(normalized, "prompt.#").Int()) + require.False(t, gjson.GetBytes(normalized, "input").Exists()) + require.False(t, gjson.GetBytes(normalized, "commands").Exists()) +} diff --git a/backend/internal/service/openai_responses_item_id.go b/backend/internal/service/openai_responses_item_id.go index f4621aa395..3a10966e89 100644 --- a/backend/internal/service/openai_responses_item_id.go +++ b/backend/internal/service/openai_responses_item_id.go @@ -70,35 +70,20 @@ func sanitizeOpenAIResponsesInputItemIDs(body []byte) ([]byte, bool, error) { type inputItem struct { body []byte - itemType string - id string - callID string stripID bool stripCallID bool - drop bool - isObject bool } items := make([]inputItem, 0) - strippedIDs := make(map[string]struct{}) - validCallIDs := make(map[string]struct{}) input.ForEach(func(_, item gjson.Result) bool { - parsed := inputItem{body: []byte(item.Raw), isObject: item.IsObject()} + parsed := inputItem{body: []byte(item.Raw)} if item.IsObject() { itemType := item.Get("type") id := item.Get("id") - parsed.itemType = strings.TrimSpace(itemType.String()) - parsed.callID = strings.TrimSpace(item.Get("call_id").String()) - parsed.stripCallID = item.Get("call_id").Exists() && shouldStripOpenAIResponsesNonPairCallID(parsed.itemType) + trimmedItemType := strings.TrimSpace(itemType.String()) + parsed.stripCallID = item.Get("call_id").Exists() && shouldStripOpenAIResponsesNonPairCallID(trimmedItemType) if id.Type == gjson.String { - parsed.id = id.String() - parsed.stripID = shouldStripOpenAIResponsesInputItemID(parsed.itemType, parsed.id) - if parsed.stripID && parsed.id != "" { - strippedIDs[parsed.id] = struct{}{} - } - } - if isCodexToolCallContextItemType(parsed.itemType) && parsed.callID != "" { - validCallIDs[parsed.callID] = struct{}{} + parsed.stripID = shouldStripOpenAIResponsesInputItemID(trimmedItemType, id.String()) } } items = append(items, parsed) @@ -115,51 +100,8 @@ func sanitizeOpenAIResponsesInputItemIDs(body []byte) ([]byte, bool, error) { return body, false, nil } - // First decide which outputs become dangling because their call_id points at - // an item ID that is being removed and no call item owns that call_id. This - // must happen before computing retained IDs: a dropped output cannot keep an - // item_reference alive merely because it used to have the same id. - for index := range items { - item := &items[index] - if !item.isObject || !isCodexToolCallOutputItemType(item.itemType) { - continue - } - if _, pointsAtStrippedID := strippedIDs[item.callID]; !pointsAtStrippedID { - continue - } - if _, hasMatchingCallID := validCallIDs[item.callID]; !hasMatchingCallID { - item.drop = true - } - } - - removedItemIDs := make(map[string]struct{}, len(strippedIDs)) - retainedItemIDs := make(map[string]struct{}, len(items)) - for _, item := range items { - if item.id == "" || item.itemType == "item_reference" { - continue - } - if item.stripID || item.drop { - removedItemIDs[item.id] = struct{}{} - continue - } - retainedItemIDs[item.id] = struct{}{} - } - for id := range retainedItemIDs { - delete(removedItemIDs, id) - } - rebuiltItems := make([][]byte, 0, len(items)) for index, item := range items { - if item.isObject { - if item.itemType == "item_reference" { - if _, dangling := removedItemIDs[item.id]; dangling { - continue - } - } - if item.drop { - continue - } - } itemBody := item.body if item.stripID { var err error diff --git a/backend/internal/service/openai_responses_item_id_test.go b/backend/internal/service/openai_responses_item_id_test.go index 17a0b2f9bc..b300b22a2e 100644 --- a/backend/internal/service/openai_responses_item_id_test.go +++ b/backend/internal/service/openai_responses_item_id_test.go @@ -45,11 +45,11 @@ func TestOpenAIResponsesInputItemIDPrefixUsesObservedOutputContracts(t *testing. } } -func TestSanitizeOpenAIResponsesInputItemIDsKeepsReferenceGraphConsistent(t *testing.T) { +func TestSanitizeOpenAIResponsesInputItemIDsDoesNotCascadeAcrossIDNamespaces(t *testing.T) { body := []byte(`{"input":[ {"type":"function_call","id":"item_bad_call","call_id":"call_valid","name":"lookup","arguments":"{}"}, {"type":"function_call_output","call_id":"call_valid","output":"preserve paired output"}, - {"type":"function_call_output","call_id":"item_bad_call","output":"drop dangling output"}, + {"type":"function_call_output","call_id":"item_bad_call","output":"preserve opaque output"}, {"type":"item_reference","id":"item_bad_call"}, {"type":"item_reference","id":"remote_valid"}, {"type":"custom_tool_call","id":"ctc_valid","call_id":"ctco_bad_output","name":"apply_patch","input":"patch"}, @@ -61,14 +61,16 @@ func TestSanitizeOpenAIResponsesInputItemIDsKeepsReferenceGraphConsistent(t *tes require.NoError(t, err) require.True(t, changed) items := gjson.GetBytes(sanitized, "input").Array() - require.Len(t, items, 5) + require.Len(t, items, 7) require.False(t, items[0].Get("id").Exists()) require.Equal(t, "call_valid", items[0].Get("call_id").String()) require.Equal(t, "preserve paired output", items[1].Get("output").String()) - require.Equal(t, "remote_valid", items[2].Get("id").String()) - require.Equal(t, "ctc_valid", items[3].Get("id").String()) - require.False(t, items[4].Get("id").Exists()) - require.Equal(t, "ctco_bad_output", items[4].Get("call_id").String()) + require.Equal(t, "preserve opaque output", items[2].Get("output").String()) + require.Equal(t, "item_bad_call", items[3].Get("id").String()) + require.Equal(t, "remote_valid", items[4].Get("id").String()) + require.Equal(t, "ctc_valid", items[5].Get("id").String()) + require.False(t, items[6].Get("id").Exists()) + require.Equal(t, "ctco_bad_output", items[6].Get("call_id").String()) } func TestSanitizeOpenAIResponsesInputItemIDsLeavesUnrelatedReferencesUntouched(t *testing.T) { @@ -93,7 +95,7 @@ func TestSanitizeOpenAIResponsesInputItemIDsPreservesReferenceToDuplicateRetaine require.Equal(t, "ctc_shared", gjson.GetBytes(sanitized, "input.2.id").String()) } -func TestSanitizeOpenAIResponsesInputItemIDsClosesReferencesAfterDroppingOutput(t *testing.T) { +func TestSanitizeOpenAIResponsesInputItemIDsPreservesOpaqueOutputsAndReferences(t *testing.T) { body := []byte(`{"input":[ {"type":"function_call","id":"item_shared","call_id":"call_real"}, {"type":"function_call_output","id":"item_shared","call_id":"item_shared","output":"dangling"}, @@ -106,10 +108,12 @@ func TestSanitizeOpenAIResponsesInputItemIDsClosesReferencesAfterDroppingOutput( require.NoError(t, err) require.True(t, changed) - require.Len(t, gjson.GetBytes(sanitized, "input").Array(), 3) + require.Len(t, gjson.GetBytes(sanitized, "input").Array(), 5) require.False(t, gjson.GetBytes(sanitized, "input.0.id").Exists()) - require.Equal(t, "kept_output", gjson.GetBytes(sanitized, "input.1.id").String()) - require.Equal(t, "kept_output", gjson.GetBytes(sanitized, "input.2.id").String()) + require.Equal(t, "dangling", gjson.GetBytes(sanitized, "input.1.output").String()) + require.Equal(t, "item_shared", gjson.GetBytes(sanitized, "input.2.id").String()) + require.Equal(t, "kept_output", gjson.GetBytes(sanitized, "input.3.id").String()) + require.Equal(t, "kept_output", gjson.GetBytes(sanitized, "input.4.id").String()) second, changedAgain, err := sanitizeOpenAIResponsesInputItemIDs(sanitized) require.NoError(t, err) diff --git a/backend/internal/service/openai_responses_tool_schema.go b/backend/internal/service/openai_responses_tool_schema.go index 08c9ecc2d6..c77640d0e3 100644 --- a/backend/internal/service/openai_responses_tool_schema.go +++ b/backend/internal/service/openai_responses_tool_schema.go @@ -26,13 +26,45 @@ const ( var errOpenAIResponsesToolSchemaLimit = errors.New("OpenAI Responses tool schema safety limit exceeded") -// shouldSanitizeOpenAIResponsesToolSchemas centralizes the platform boundary -// for callers. These rewrites describe OpenAI Responses constraints, not the -// behavior of every provider routed through the generic OpenAI gateway. -func shouldSanitizeOpenAIResponsesToolSchemas(platform string) bool { +// shouldRepairOpenAIResponsesNullToolSchemaType reports whether the upstream +// path requires a concrete object type at a function tool's parameter root. +// This defect is shared by the OpenAI, Anthropic, and CN-compatible paths. +func shouldRepairOpenAIResponsesNullToolSchemaType(platform string) bool { + return platform == PlatformOpenAI || platform == PlatformAnthropic || IsCNProvider(platform) +} + +// shouldSanitizeOpenAIResponsesToolSchemaPatterns is intentionally narrower: +// regex lookaround rejection is an OpenAI-specific schema constraint. +func shouldSanitizeOpenAIResponsesToolSchemaPatterns(platform string) bool { return platform == PlatformOpenAI } +func sanitizeOpenAIResponsesToolSchemasForPlatform(body []byte, platform string) ([]byte, bool, error) { + normalized := body + changed := false + if shouldRepairOpenAIResponsesNullToolSchemaType(platform) { + next, repaired, err := sanitizeOpenAIResponsesToolParameterTypes(normalized) + if err != nil { + return body, false, fmt.Errorf("sanitize OpenAI Responses tool parameters: %w", err) + } + if repaired { + normalized = next + changed = true + } + } + if shouldSanitizeOpenAIResponsesToolSchemaPatterns(platform) { + next, sanitized, err := sanitizeOpenAIResponsesToolSchemaPatterns(normalized) + if err != nil { + return body, false, fmt.Errorf("sanitize OpenAI Responses tool schema patterns: %w", err) + } + if sanitized { + normalized = next + changed = true + } + } + return normalized, changed, nil +} + // sanitizeOpenAIResponsesToolSchemaPatterns removes only schema constraints // containing regex lookaround, which OpenAI rejects. It deliberately does not // descend into instance-valued keywords such as default, examples, const, or diff --git a/backend/internal/service/openai_responses_tool_schema_test.go b/backend/internal/service/openai_responses_tool_schema_test.go index eeb9cf5dab..cf676645e1 100644 --- a/backend/internal/service/openai_responses_tool_schema_test.go +++ b/backend/internal/service/openai_responses_tool_schema_test.go @@ -239,11 +239,59 @@ func TestSanitizeOpenAIResponsesToolSchemas_InvalidAndTrailingJSON(t *testing.T) } } -func TestShouldSanitizeOpenAIResponsesToolSchemas_PlatformBoundary(t *testing.T) { - require.True(t, shouldSanitizeOpenAIResponsesToolSchemas(PlatformOpenAI)) - for _, platform := range []string{PlatformGrok, PlatformKimi, PlatformZhipu, PlatformDeepseek, PlatformComposite, ""} { - require.False(t, shouldSanitizeOpenAIResponsesToolSchemas(platform), platform) +func TestOpenAIResponsesToolSchemaCapabilities_PlatformBoundary(t *testing.T) { + tests := []struct { + platform string + repairNullType bool + removeLookaround bool + }{ + {PlatformOpenAI, true, true}, + {PlatformAnthropic, true, false}, + {PlatformKimi, true, false}, + {PlatformZhipu, true, false}, + {PlatformDeepseek, true, false}, + {PlatformGrok, false, false}, + {PlatformGemini, false, false}, + {PlatformAntigravity, false, false}, + {PlatformComposite, false, false}, + {"", false, false}, } + for _, tt := range tests { + t.Run(tt.platform, func(t *testing.T) { + require.Equal(t, tt.repairNullType, shouldRepairOpenAIResponsesNullToolSchemaType(tt.platform)) + require.Equal(t, tt.removeLookaround, shouldSanitizeOpenAIResponsesToolSchemaPatterns(tt.platform)) + }) + } +} + +func TestSanitizeOpenAIResponsesToolSchemasForPlatform_ReplayBoundary(t *testing.T) { + body := []byte(`{"tools":[{"type":"function","parameters":{"type":null,"properties":{"query":{"type":"string","pattern":"(?=keep)"}}}}]}`) + + // A malformed tool definition may be replayed after account failover. Every + // compatible account must repair it, while non-OpenAI providers retain their + // supported regex semantics. + for _, platform := range []string{PlatformAnthropic, PlatformKimi, PlatformZhipu, PlatformDeepseek} { + t.Run(platform, func(t *testing.T) { + for attempt := 0; attempt < 2; attempt++ { + normalized, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, platform) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(normalized, "tools.0.parameters.type").String()) + require.Equal(t, "(?=keep)", gjson.GetBytes(normalized, "tools.0.parameters.properties.query.pattern").String()) + } + }) + } + + openAI, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, PlatformOpenAI) + require.NoError(t, err) + require.True(t, changed) + require.Equal(t, "object", gjson.GetBytes(openAI, "tools.0.parameters.type").String()) + require.False(t, gjson.GetBytes(openAI, "tools.0.parameters.properties.query.pattern").Exists()) + + unsupported, changed, err := sanitizeOpenAIResponsesToolSchemasForPlatform(body, PlatformGrok) + require.NoError(t, err) + require.False(t, changed) + require.Equal(t, string(body), string(unsupported)) } // 索引映射:只有坏条目被改,前后兄弟条目按原下标保持不变。 diff --git a/backend/internal/service/openai_upstream_client_error_test.go b/backend/internal/service/openai_upstream_client_error_test.go index 2b6ed51967..291c1c491e 100644 --- a/backend/internal/service/openai_upstream_client_error_test.go +++ b/backend/internal/service/openai_upstream_client_error_test.go @@ -164,6 +164,7 @@ func TestHandleErrorResponse_NonDeterministicStatusesKeepGeneric502(t *testing.T {"unprocessable", http.StatusUnprocessableEntity, `{"error":{"message":"Invalid schema for field messages"}}`, http.StatusBadGateway, "upstream_error", "Upstream request failed"}, // 401/402/403 是网关运营方的凭据/账单问题,必须继续对客户端屏蔽上游账号状态。 + // 403 的自由文本不能升级成 durable access-state typed failover;只有明确结构化 code 才可以。 {"unauthorized", http.StatusUnauthorized, `{"error":{"message":"Incorrect API key provided: sk-abc"}}`, http.StatusBadGateway, "upstream_error", "Upstream authentication failed, please contact administrator"}, {"forbidden", http.StatusForbidden, `{"error":{"message":"Your account is deactivated"}}`, @@ -183,15 +184,11 @@ func TestHandleErrorResponse_NonDeterministicStatusesKeepGeneric502(t *testing.T newOpenAIUpstreamErrorResponse(tc.statusCode, tc.body), c, newOpenAIUpstreamErrorTestAccount(), nil, ) + require.Error(t, err) if tc.name == "forbidden" { var failoverErr *UpstreamFailoverError - require.ErrorAs(t, err, &failoverErr) - require.Equal(t, http.StatusForbidden, failoverErr.StatusCode) - require.False(t, c.Writer.Written()) - return + require.False(t, errors.As(err, &failoverErr)) } - - require.Error(t, err) require.Equal(t, tc.wantStatus, rec.Code) require.Equal(t, tc.wantType, gjson.Get(rec.Body.String(), "error.type").String()) require.Equal(t, tc.wantMsg, gjson.Get(rec.Body.String(), "error.message").String()) diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index ae41c3bbbd..3f5d1e57b5 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -489,6 +489,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( pendingClientMessageBytes := int64(0) capacityFailoverSuppressedLogged := false clientDisconnected := false + officialOpenAIResponses := account != nil && account.Platform == PlatformOpenAI bareErrorPending := false var bareErrorPayload []byte bareErrorMessage := "" @@ -642,7 +643,15 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( replayCollector.AddEvent(eventType, upstreamMessage) var upstreamEventErr error - suppressClientMessage := bareErrorPending && eventType != "response.failed" + if officialOpenAIResponses && bareErrorPending && (eventType == "response.completed" || eventType == "response.done") { + // Some upstreams emit a recoverable bare error before the authoritative + // successful terminal. Do not replace that terminal with a synthetic + // failure or retain side effects from the superseded error. + bareErrorPending = false + bareErrorPayload = nil + bareErrorMessage = "" + } + suppressClientMessage := officialOpenAIResponses && bareErrorPending && eventType != "response.failed" if eventType == "error" || eventType == "response.failed" { errMessage := extractOpenAISSEErrorMessage(upstreamMessage) if errMessage == "" { @@ -677,7 +686,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( return nil, s.newOpenAIStreamFailoverError(c, account, true, resp.Header.Get("x-request-id"), upstreamMessage, errMessage, resp.Header) } if account.Platform != PlatformGrok && !failureAccountSideEffectsApplied { - if eventType == "response.failed" || (shouldFailover && !requestScopedCapacity) { + if eventType == "response.failed" || (!officialOpenAIResponses && shouldFailover && !requestScopedCapacity) { failureAccountSideEffectsApplied = s.handleOpenAIWSFailureAccountSideEffects(ctx, account, mappedModel, resp.Header, upstreamMessage) } } @@ -685,7 +694,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( logOpenAICapacityFailoverSuppressed(ctx, account, "ws_http_bridge", resp.Header.Get("x-request-id"), eventType) capacityFailoverSuppressedLogged = true } - if eventType == "error" && account.Platform == PlatformGrok { + if eventType == "error" && !officialOpenAIResponses { upstreamEventErr = errors.New(errMessage) } else if eventType == "error" { bareErrorPending = true diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index a8c10cf893..4aec9d7d8a 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -758,6 +758,40 @@ func TestProxyOpenAIWSHTTPBridgeTurnBareErrorEOFSynthesizesFailed(t *testing.T) require.Equal(t, "resp_eof", gjson.GetBytes(writes[1], "response.id").String()) } +func TestProxyOpenAIWSHTTPBridgeTurnBareErrorFollowedByCompletedUsesCompleted(t *testing.T) { + gin.SetMode(gin.TestMode) + body := strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"resp_recovered","status":"in_progress"}}`, + ``, + `data: {"type":"error","error":{"code":"transient","message":"retrying"}}`, + ``, + `data: {"type":"response.completed","response":{"id":"resp_recovered","status":"completed","output":[],"usage":{"input_tokens":8,"output_tokens":4}}}`, + ``, + }, "\n") + upstream := &httpUpstreamRecorder{resp: &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(body))}} + svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ID: 113, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1} + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) + var writes [][]byte + + result, err := svc.proxyOpenAIWSHTTPBridgeTurn(context.Background(), c, account, "sk-test", payload, len(payload), "gpt-5", "", "", "", "", 2, func(message []byte) error { + writes = append(writes, append([]byte(nil), message...)) + return nil + }) + + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, "response.completed", result.UpstreamTerminalEvent) + require.Equal(t, 8, result.Usage.InputTokens) + require.Equal(t, 4, result.Usage.OutputTokens) + require.Len(t, writes, 2) + require.Equal(t, "response.created", gjson.GetBytes(writes[0], "type").String()) + require.Equal(t, "response.completed", gjson.GetBytes(writes[1], "type").String()) +} + func TestProxyOpenAIWSHTTPBridgeTurnStagesMetadataBeforeCapacityFailover(t *testing.T) { gin.SetMode(gin.TestMode) body := strings.Join([]string{ diff --git a/backend/internal/service/openai_ws_state_store.go b/backend/internal/service/openai_ws_state_store.go index d3b6891b0a..fe4ee54926 100644 --- a/backend/internal/service/openai_ws_state_store.go +++ b/backend/internal/service/openai_ws_state_store.go @@ -13,6 +13,8 @@ import ( const ( openAIWSResponseAccountCachePrefix = "openai:response:" + openAIHTTPResponseOwnerUserPrefix = "openai:http-response-owner:user:" + openAIHTTPResponseOwnerKeyPrefix = "openai:http-response-owner:key:" openAIWSStateStoreCleanupInterval = time.Minute openAIWSStateStoreCleanupMaxPerMap = 512 openAIWSStateStoreMaxEntriesPerMap = 65536 @@ -24,6 +26,12 @@ type openAIWSAccountBinding struct { expiresAt time.Time } +type openAIHTTPResponseOwnerBinding struct { + userID int64 + apiKeyID int64 + expiresAt time.Time +} + type openAIWSConnBinding struct { connID string expiresAt time.Time @@ -49,6 +57,8 @@ type OpenAIWSStateStore interface { BindResponseAccount(ctx context.Context, groupID int64, responseID string, accountID int64, ttl time.Duration) error GetResponseAccount(ctx context.Context, groupID int64, responseID string) (int64, error) DeleteResponseAccount(ctx context.Context, groupID int64, responseID string) error + BindHTTPResponseOwner(ctx context.Context, groupID int64, responseID string, userID, apiKeyID int64, ttl time.Duration) error + GetHTTPResponseOwner(ctx context.Context, groupID int64, responseID string) (userID, apiKeyID int64, found bool, err error) BindResponseConn(responseID, connID string, ttl time.Duration) GetResponseConn(responseID string) (string, bool) @@ -68,6 +78,8 @@ type defaultOpenAIWSStateStore struct { responseToAccountMu sync.RWMutex responseToAccount map[string]openAIWSAccountBinding + responseOwnerMu sync.RWMutex + responseOwners map[string]openAIHTTPResponseOwnerBinding responseToConnMu sync.RWMutex responseToConn map[string]openAIWSConnBinding sessionToTurnStateMu sync.RWMutex @@ -83,6 +95,7 @@ func NewOpenAIWSStateStore(cache GatewayCache) OpenAIWSStateStore { store := &defaultOpenAIWSStateStore{ cache: cache, responseToAccount: make(map[string]openAIWSAccountBinding, 256), + responseOwners: make(map[string]openAIHTTPResponseOwnerBinding, 256), responseToConn: make(map[string]openAIWSConnBinding, 256), sessionToTurnState: make(map[string]openAIWSTurnStateBinding, 256), sessionToConn: make(map[string]openAIWSSessionConnBinding, 256), @@ -91,6 +104,72 @@ func NewOpenAIWSStateStore(cache GatewayCache) OpenAIWSStateStore { return store } +func (s *defaultOpenAIWSStateStore) BindHTTPResponseOwner(ctx context.Context, groupID int64, responseID string, userID, apiKeyID int64, ttl time.Duration) error { + id := normalizeOpenAIWSResponseID(responseID) + if id == "" || userID <= 0 || apiKeyID <= 0 { + return nil + } + ttl = normalizeOpenAIWSTTL(ttl) + s.maybeCleanup() + + mapKey := openAIWSResponseAccountMapKey(groupID, id) + s.responseOwnerMu.Lock() + ensureBindingCapacity(s.responseOwners, mapKey, openAIWSStateStoreMaxEntriesPerMap) + s.responseOwners[mapKey] = openAIHTTPResponseOwnerBinding{ + userID: userID, apiKeyID: apiKeyID, expiresAt: time.Now().Add(ttl), + } + s.responseOwnerMu.Unlock() + + if s.cache == nil { + return nil + } + cacheCtx, cancel := withOpenAIWSStateStoreRedisTimeout(ctx) + defer cancel() + if err := s.cache.SetSessionAccountID(cacheCtx, groupID, openAIHTTPResponseOwnerCacheKey(openAIHTTPResponseOwnerUserPrefix, id), userID, ttl); err != nil { + return err + } + return s.cache.SetSessionAccountID(cacheCtx, groupID, openAIHTTPResponseOwnerCacheKey(openAIHTTPResponseOwnerKeyPrefix, id), apiKeyID, ttl) +} + +func (s *defaultOpenAIWSStateStore) GetHTTPResponseOwner(ctx context.Context, groupID int64, responseID string) (int64, int64, bool, error) { + id := normalizeOpenAIWSResponseID(responseID) + if id == "" { + return 0, 0, false, nil + } + s.maybeCleanup() + + now := time.Now() + mapKey := openAIWSResponseAccountMapKey(groupID, id) + s.responseOwnerMu.RLock() + if binding, ok := s.responseOwners[mapKey]; ok && now.Before(binding.expiresAt) { + s.responseOwnerMu.RUnlock() + return binding.userID, binding.apiKeyID, true, nil + } + s.responseOwnerMu.RUnlock() + + if s.cache == nil { + return 0, 0, false, nil + } + cacheCtx, cancel := withOpenAIWSStateStoreRedisTimeout(ctx) + defer cancel() + userID, err := s.cache.GetSessionAccountID(cacheCtx, groupID, openAIHTTPResponseOwnerCacheKey(openAIHTTPResponseOwnerUserPrefix, id)) + if err != nil || userID <= 0 { + return 0, 0, false, err + } + apiKeyID, err := s.cache.GetSessionAccountID(cacheCtx, groupID, openAIHTTPResponseOwnerCacheKey(openAIHTTPResponseOwnerKeyPrefix, id)) + if err != nil || apiKeyID <= 0 { + return 0, 0, false, err + } + + s.responseOwnerMu.Lock() + ensureBindingCapacity(s.responseOwners, mapKey, openAIWSStateStoreMaxEntriesPerMap) + s.responseOwners[mapKey] = openAIHTTPResponseOwnerBinding{ + userID: userID, apiKeyID: apiKeyID, expiresAt: now.Add(time.Minute), + } + s.responseOwnerMu.Unlock() + return userID, apiKeyID, true, nil +} + func (s *defaultOpenAIWSStateStore) BindResponseAccount(ctx context.Context, groupID int64, responseID string, accountID int64, ttl time.Duration) error { id := normalizeOpenAIWSResponseID(responseID) if id == "" || accountID <= 0 { @@ -115,6 +194,22 @@ func (s *defaultOpenAIWSStateStore) BindResponseAccount(ctx context.Context, gro return s.cache.SetSessionAccountID(cacheCtx, groupID, cacheKey, accountID, ttl) } +func cleanupExpiredHTTPResponseOwnerBindings(bindings map[string]openAIHTTPResponseOwnerBinding, now time.Time, maxScan int) { + if len(bindings) == 0 || maxScan <= 0 { + return + } + scanned := 0 + for key, binding := range bindings { + if now.After(binding.expiresAt) { + delete(bindings, key) + } + scanned++ + if scanned >= maxScan { + break + } + } +} + func (s *defaultOpenAIWSStateStore) GetResponseAccount(ctx context.Context, groupID int64, responseID string) (int64, error) { id := normalizeOpenAIWSResponseID(responseID) if id == "" { @@ -319,6 +414,10 @@ func (s *defaultOpenAIWSStateStore) maybeCleanup() { cleanupExpiredAccountBindings(s.responseToAccount, now, openAIWSStateStoreCleanupMaxPerMap) s.responseToAccountMu.Unlock() + s.responseOwnerMu.Lock() + cleanupExpiredHTTPResponseOwnerBindings(s.responseOwners, now, openAIWSStateStoreCleanupMaxPerMap) + s.responseOwnerMu.Unlock() + s.responseToConnMu.Lock() cleanupExpiredConnBindings(s.responseToConn, now, openAIWSStateStoreCleanupMaxPerMap) s.responseToConnMu.Unlock() @@ -419,6 +518,11 @@ func openAIWSResponseAccountCacheKey(responseID string) string { return openAIWSResponseAccountCachePrefix + hex.EncodeToString(sum[:]) } +func openAIHTTPResponseOwnerCacheKey(prefix, responseID string) string { + sum := sha256.Sum256([]byte(responseID)) + return prefix + hex.EncodeToString(sum[:]) +} + // openAIWSResponseAccountMapKey 本地热缓存按分组隔离的 key,与 Redis 层保持一致,避免跨组命中。 func openAIWSResponseAccountMapKey(groupID int64, responseID string) string { return fmt.Sprintf("%d:%s", groupID, responseID) diff --git a/backend/internal/service/openai_ws_state_store_test.go b/backend/internal/service/openai_ws_state_store_test.go index 869e5e0fd6..6b20a95507 100644 --- a/backend/internal/service/openai_ws_state_store_test.go +++ b/backend/internal/service/openai_ws_state_store_test.go @@ -28,6 +28,27 @@ func TestOpenAIWSStateStore_BindGetDeleteResponseAccount(t *testing.T) { require.Zero(t, accountID) } +func TestOpenAIWSStateStore_HTTPResponseOwnerPersistsAcrossStoreInstances(t *testing.T) { + cache := &stubGatewayCache{} + ctx := context.Background() + groupID := int64(8) + writer := NewOpenAIWSStateStore(cache) + + require.NoError(t, writer.BindHTTPResponseOwner(ctx, groupID, "resp_owned", 201, 301, time.Minute)) + userID, apiKeyID, found, err := writer.GetHTTPResponseOwner(ctx, groupID, "resp_owned") + require.NoError(t, err) + require.True(t, found) + require.Equal(t, int64(201), userID) + require.Equal(t, int64(301), apiKeyID) + + reader := NewOpenAIWSStateStore(cache) + userID, apiKeyID, found, err = reader.GetHTTPResponseOwner(ctx, groupID, "resp_owned") + require.NoError(t, err) + require.True(t, found) + require.Equal(t, int64(201), userID) + require.Equal(t, int64(301), apiKeyID) +} + func TestOpenAIWSStateStore_ResponseConnTTL(t *testing.T) { store := NewOpenAIWSStateStore(nil) store.BindResponseConn("resp_conn", "conn_1", 30*time.Millisecond) diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay.go b/backend/internal/service/openai_ws_v2/passthrough_relay.go index e5e2633644..67c26f3e96 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay.go @@ -270,30 +270,34 @@ func Relay( if !options.StartClientAfterFirstDownstream { startClientReader() } - go runUpstreamToClient( - relayCtx, - upstreamConn, - writeClient, - startAt, - nowFn, - state, - options.OnUsageParseFailure, - options.OnTurnComplete, - options.BeforeWriteClient, - options.BeforeClientWrite, - options.AfterClientWrite, - func(msgType coderws.MessageType, payload []byte) { - if options.StartClientAfterFirstDownstream { - startClientReader() - } - }, - &dropDownstreamWrites, - upstreamToClientFrames, - droppedDownstreamFrames, - markActivity, - onTrace, - exitCh, - ) + upstreamDone := make(chan struct{}) + go func() { + defer close(upstreamDone) + runUpstreamToClient( + relayCtx, + upstreamConn, + writeClient, + startAt, + nowFn, + state, + options.OnUsageParseFailure, + options.OnTurnComplete, + options.BeforeWriteClient, + options.BeforeClientWrite, + options.AfterClientWrite, + func(msgType coderws.MessageType, payload []byte) { + if options.StartClientAfterFirstDownstream { + startClientReader() + } + }, + &dropDownstreamWrites, + upstreamToClientFrames, + droppedDownstreamFrames, + markActivity, + onTrace, + exitCh, + ) + }() go runIdleWatchdog(relayCtx, nowFn, options.IdleTimeout, &lastActivity, onTrace, exitCh) firstExit := <-exitCh @@ -347,6 +351,10 @@ func Relay( relayCancel() _ = upstreamConn.Close() + // ReadFrame observes relayCtx cancellation and Close is the transport-level + // fallback. Join the reader before touching relayState or firing the final + // turn callback; otherwise a late read can race Relay's result settlement. + <-upstreamDone emitTurnComplete(options.OnTurnComplete, state, finalizePendingBareError(state, nowFn())) enrichResult(&result, state, nowFn().Sub(startAt)) diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay_test.go b/backend/internal/service/openai_ws_v2/passthrough_relay_test.go index 3bca8bf0ab..28b896f248 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay_test.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay_test.go @@ -41,6 +41,16 @@ type closeSpyFrameConn struct { closeCalls atomic.Int32 } +type cancelJoinProbeFrameConn struct { + readStarted chan struct{} + readCanceled chan struct{} + allowReturn chan struct{} + readReturned chan struct{} + startOnce sync.Once + cancelOnce sync.Once + returnOnce sync.Once +} + func newPassthroughTestFrameConn(frames []passthroughTestFrame, autoClose bool) *passthroughTestFrameConn { c := &passthroughTestFrameConn{ readCh: make(chan passthroughTestFrame, len(frames)+1), @@ -179,6 +189,38 @@ func (c *closeSpyFrameConn) CloseCalls() int32 { return c.closeCalls.Load() } +func newCancelJoinProbeFrameConn() *cancelJoinProbeFrameConn { + return &cancelJoinProbeFrameConn{ + readStarted: make(chan struct{}), + readCanceled: make(chan struct{}), + allowReturn: make(chan struct{}), + readReturned: make(chan struct{}), + } +} + +func (c *cancelJoinProbeFrameConn) ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) { + c.startOnce.Do(func() { close(c.readStarted) }) + <-ctx.Done() + c.cancelOnce.Do(func() { close(c.readCanceled) }) + <-c.allowReturn + c.returnOnce.Do(func() { close(c.readReturned) }) + return coderws.MessageText, nil, ctx.Err() +} + +func (c *cancelJoinProbeFrameConn) WriteFrame(ctx context.Context, _ coderws.MessageType, _ []byte) error { + if ctx == nil { + ctx = context.Background() + } + select { + case <-ctx.Done(): + return ctx.Err() + default: + return nil + } +} + +func (c *cancelJoinProbeFrameConn) Close() error { return nil } + func TestRelay_BasicRelayAndUsage(t *testing.T) { t.Parallel() @@ -372,6 +414,58 @@ func TestRelay_IdleTimeoutDoesNotCloseClientOnError(t *testing.T) { require.GreaterOrEqual(t, upstreamConn.CloseCalls(), int32(1)) } +func TestRelay_JoinsUpstreamReaderBeforeReturning(t *testing.T) { + t.Parallel() + + clientConn := &closeSpyFrameConn{} + upstreamConn := newCancelJoinProbeFrameConn() + ctx, cancel := context.WithCancel(context.Background()) + resultCh := make(chan *RelayExit, 1) + + go func() { + _, relayExit := Relay( + ctx, + clientConn, + upstreamConn, + []byte(`{"type":"response.create","model":"gpt-4o","input":[]}`), + RelayOptions{}, + ) + resultCh <- relayExit + }() + + select { + case <-upstreamConn.readStarted: + case <-time.After(time.Second): + t.Fatal("upstream reader did not start") + } + cancel() + select { + case <-upstreamConn.readCanceled: + case <-time.After(time.Second): + t.Fatal("upstream reader did not observe relay cancellation") + } + select { + case <-resultCh: + t.Fatal("Relay returned before the upstream reader exited") + case <-time.After(50 * time.Millisecond): + } + + close(upstreamConn.allowReturn) + select { + case relayExit := <-resultCh: + require.NotNil(t, relayExit) + require.ErrorIs(t, relayExit.Err, context.Canceled) + case <-time.After(time.Second): + t.Fatal("Relay did not return after the upstream reader exited") + } + select { + case <-upstreamConn.readReturned: + default: + t.Fatal("Relay returned before the upstream reader completion signal") + } + require.Zero(t, clientConn.CloseCalls(), "错误路径不应提前关闭客户端连接") +} + func TestRelay_NilConnections(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/ops_upstream_context.go b/backend/internal/service/ops_upstream_context.go index b4919b7809..509620fa1b 100644 --- a/backend/internal/service/ops_upstream_context.go +++ b/backend/internal/service/ops_upstream_context.go @@ -388,7 +388,9 @@ type OpsUpstreamErrorEvent struct { Detail string `json:"detail,omitempty"` // SkipMonitoring is request-local rule state. It is intentionally excluded - // from persisted attempt JSON and only lets the final attempt control Ops. + // from persisted attempt JSON. The logger consults it only when this event is + // the final client-visible failure; recovered attempts remain provider-health + // telemetry and do not count as failed requests. SkipMonitoring bool `json:"-"` } @@ -427,11 +429,10 @@ func appendOpsUpstreamError(c *gin.Context, ev OpsUpstreamErrorEvent) { checkSkipMonitoringForUpstreamEvent(c, &evCopy) } -// checkSkipMonitoringForUpstreamEvent checks whether the upstream error event -// matches a passthrough rule with skip_monitoring=true and, if so, sets the -// OpsSkipPassthroughKey on the context. This ensures intermediate retry / -// failover errors (which never go through the final applyErrorPassthroughRule -// path) can still suppress ops_error_logs recording. +// checkSkipMonitoringForUpstreamEvent snapshots whether this attempt matches a +// skip_monitoring passthrough rule. The final failure decides whether the +// request error is hidden; an intermediate recovered attempt cannot suppress a +// later client-visible failure. func checkSkipMonitoringForUpstreamEvent(c *gin.Context, ev *OpsUpstreamErrorEvent) { if ev.UpstreamStatusCode == 0 { return diff --git a/backend/internal/service/temp_unsched.go b/backend/internal/service/temp_unsched.go index 9b120c5775..92b7da3f72 100644 --- a/backend/internal/service/temp_unsched.go +++ b/backend/internal/service/temp_unsched.go @@ -29,7 +29,6 @@ type TempUnschedCache interface { // aggregate pool API-key failures across gateway instances. type OpenAIAPIKeyHealthCache interface { RecordOpenAIAPIKeyHealthFailure(ctx context.Context, accountID int64, windowMinutes, threshold int) (count int64, tripped bool, err error) - ResetOpenAIAPIKeyHealthFailures(ctx context.Context, accountID int64) error } // TimeoutCounterCache 超时计数器缓存接口