diff --git a/internal/runtime/executor/codex_executor_stream_output_test.go b/internal/runtime/executor/codex_executor_stream_output_test.go index 46a43de3..19b38fb5 100644 --- a/internal/runtime/executor/codex_executor_stream_output_test.go +++ b/internal/runtime/executor/codex_executor_stream_output_test.go @@ -506,6 +506,11 @@ func TestCodexTerminalFailureErrClassifiesStatus(t *testing.T) { event: `{"type":"error","error":{"type":"invalid_request_error","code":"invalid_value","message":"Invalid input."}}`, wantStatus: http.StatusBadRequest, }, + { + name: "cyber policy", + event: `{"type":"error","error":{"type":"invalid_request","code":"cyber_policy","message":"This content was flagged for possible cybersecurity risk."}}`, + wantStatus: http.StatusBadRequest, + }, { name: "authentication", event: `{"type":"response.failed","response":{"error":{"type":"authentication_error","code":"invalid_api_key","message":"Invalid token."}}}`, diff --git a/internal/runtime/executor/codex_executor_terminal.go b/internal/runtime/executor/codex_executor_terminal.go index b2845f92..f05727dc 100644 --- a/internal/runtime/executor/codex_executor_terminal.go +++ b/internal/runtime/executor/codex_executor_terminal.go @@ -160,6 +160,8 @@ func codexTerminalFailureStatus(body []byte) int { errorType := strings.ToLower(strings.TrimSpace(gjson.GetBytes(body, "error.type").String())) errorCode := strings.ToLower(strings.TrimSpace(gjson.GetBytes(body, "error.code").String())) switch { + case errorCode == "cyber_policy": + return http.StatusBadRequest case errorType == "invalid_request_error", errorType == "bad_request_error": return http.StatusBadRequest case errorType == "authentication_error", errorCode == "invalid_api_key", errorCode == "unauthorized": diff --git a/sdk/api/handlers/handlers_execution.go b/sdk/api/handlers/handlers_execution.go index 7c25ab14..b2deb73f 100644 --- a/sdk/api/handlers/handlers_execution.go +++ b/sdk/api/handlers/handlers_execution.go @@ -291,6 +291,10 @@ func (h *BaseAPIHandler) applyRequestInterceptorsAfterPluginExecutorRoute(ctx co return req, opts, nil } +func ExecutionErrorMessage(err error) *interfaces.ErrorMessage { + return executionErrorMessage(err) +} + func executionErrorMessage(err error) *interfaces.ErrorMessage { var terminated *coreexecutor.RequestTerminatedError if errors.As(err, &terminated) && terminated != nil { diff --git a/sdk/api/handlers/handlers_routing.go b/sdk/api/handlers/handlers_routing.go index c590f415..fe79b412 100644 --- a/sdk/api/handlers/handlers_routing.go +++ b/sdk/api/handlers/handlers_routing.go @@ -197,7 +197,17 @@ func (h *BaseAPIHandler) getRequestDetailsWithOptions(modelName string, allowIma } if len(providers) == 0 { - return nil, "", &interfaces.ErrorMessage{StatusCode: http.StatusBadGateway, Error: fmt.Errorf("unknown provider for model %s", modelName)} + // The client asked for a model this proxy cannot route. Report it as a request + // error so streaming clients receive an actionable message instead of a + // gateway failure they would keep retrying. 400 is used rather than 404 to keep + // it distinguishable from an unregistered HTTP route. + return nil, "", &interfaces.ErrorMessage{ + StatusCode: http.StatusBadRequest, + Error: fmt.Errorf( + `{"error":{"message":"unknown provider for model %s","type":"invalid_request_error","code":"model_not_found","param":"model"}}`, + modelName, + ), + } } // The thinking suffix is preserved in the model name itself, so no diff --git a/sdk/api/handlers/openai/openai_responses_handlers.go b/sdk/api/handlers/openai/openai_responses_handlers.go index cb45b959..ba78a447 100644 --- a/sdk/api/handlers/openai/openai_responses_handlers.go +++ b/sdk/api/handlers/openai/openai_responses_handlers.go @@ -583,7 +583,7 @@ func (h *OpenAIResponsesAPIHandler) forwardResponsesStream(c *gin.Context, flush }, WriteTerminalError: func(errMsg *interfaces.ErrorMessage) { framer.Flush(c.Writer) - if errMsg == nil { + if !shouldExposeResponsesUpstreamError(errMsg) { return } status := http.StatusInternalServerError diff --git a/sdk/api/handlers/openai/openai_responses_handlers_stream_error_test.go b/sdk/api/handlers/openai/openai_responses_handlers_stream_error_test.go index 54d14675..f416dcd0 100644 --- a/sdk/api/handlers/openai/openai_responses_handlers_stream_error_test.go +++ b/sdk/api/handlers/openai/openai_responses_handlers_stream_error_test.go @@ -13,31 +13,78 @@ import ( sdkconfig "github.com/router-for-me/CLIProxyAPI/v7/sdk/config" ) -func TestForwardResponsesStreamTerminalErrorUsesResponsesErrorChunk(t *testing.T) { +// TestForwardResponsesStreamExposesOnlyClientErrors pins the SSE side: only +// request-shape failures reach the client. Credential, quota and transport +// failures end the stream silently so the client retries on its own. +func TestForwardResponsesStreamExposesOnlyClientErrors(t *testing.T) { gin.SetMode(gin.TestMode) - base := handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, nil) - h := NewOpenAIResponsesAPIHandler(base) - recorder := httptest.NewRecorder() - c, _ := gin.CreateTestContext(recorder) - c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) - - flusher, ok := c.Writer.(http.Flusher) - if !ok { - t.Fatalf("expected gin writer to implement http.Flusher") + tests := []struct { + name string + status int + message string + wantExposed bool + }{ + { + name: "bad request", + status: http.StatusBadRequest, + message: `{"error":{"type":"invalid_request","code":"cyber_policy","message":"blocked"}}`, + wantExposed: true, + }, + { + // Observed in production: the same cyber_policy rejection arrives with 502 + // when it is surfaced through the websocket disconnect channel. + name: "cyber policy behind bad gateway status", + status: http.StatusBadGateway, + message: `{"error":{"type":"invalid_request","code":"cyber_policy","message":"This content was flagged for possible cybersecurity risk.","param":null}}`, + wantExposed: true, + }, + { + name: "context length exceeded behind bad gateway status", + status: http.StatusBadGateway, + message: `{"error":{"type":"invalid_request_error","code":"context_length_exceeded","message":"Your input exceeds the context window."}}`, + wantExposed: true, + }, + {name: "conflict", status: http.StatusConflict, message: "conflict", wantExposed: true}, + {name: "message too big", status: http.StatusRequestEntityTooLarge, message: "too large", wantExposed: true}, + {name: "unprocessable entity", status: http.StatusUnprocessableEntity, message: "invalid input", wantExposed: true}, + {name: "authentication", status: http.StatusUnauthorized, message: "invalid credential"}, + {name: "payment required", status: http.StatusPaymentRequired, message: "insufficient credits"}, + {name: "quota error", status: http.StatusTooManyRequests, message: "usage limit reached"}, + {name: "request timeout", status: http.StatusRequestTimeout, message: "upstream timeout"}, + {name: "transport error", status: http.StatusInternalServerError, message: "unexpected EOF"}, + {name: "upstream websocket drop", status: http.StatusInternalServerError, + message: `{"error":{"message":"websocket: close 1006 (abnormal closure): unexpected EOF","type":"server_error","code":"internal_server_error"}}`}, } - data := make(chan []byte) - errs := make(chan *interfaces.ErrorMessage, 1) - errs <- &interfaces.ErrorMessage{StatusCode: http.StatusInternalServerError, Error: errors.New("unexpected EOF")} - close(errs) + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + base := handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, nil) + h := NewOpenAIResponsesAPIHandler(base) - h.forwardResponsesStream(c, flusher, func(error) {}, data, errs, nil) - body := recorder.Body.String() - if !strings.Contains(body, `"type":"error"`) { - t.Fatalf("expected responses error chunk, got: %q", body) - } - if strings.Contains(body, `"error":{`) { - t.Fatalf("expected streaming error chunk (top-level type), got HTTP error body: %q", body) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + + flusher, ok := c.Writer.(http.Flusher) + if !ok { + t.Fatal("expected gin writer to implement http.Flusher") + } + + data := make(chan []byte) + errs := make(chan *interfaces.ErrorMessage, 1) + errs <- &interfaces.ErrorMessage{StatusCode: tc.status, Error: errors.New(tc.message)} + close(errs) + + h.forwardResponsesStream(c, flusher, func(error) {}, data, errs, nil) + body := recorder.Body.String() + exposed := strings.Contains(body, `"type":"error"`) + if exposed != tc.wantExposed { + t.Fatalf("error exposed = %t, want %t: %q", exposed, tc.wantExposed, body) + } + if exposed && strings.Contains(body, `"error":{`) { + t.Fatalf("expected streaming error chunk, got HTTP error body: %q", body) + } + }) } } diff --git a/sdk/api/handlers/openai/openai_responses_websocket.go b/sdk/api/handlers/openai/openai_responses_websocket.go index fce86e82..2a46510e 100644 --- a/sdk/api/handlers/openai/openai_responses_websocket.go +++ b/sdk/api/handlers/openai/openai_responses_websocket.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "errors" + "net" "net/http" "strings" "sync" @@ -139,6 +140,36 @@ func (w *responsesWebsocketWriter) closeForUpstreamError(err error) (bool, error return true, errClose } +func (w *responsesWebsocketWriter) closeWithoutError() (bool, error) { + if w == nil || w.conn == nil { + return false, nil + } + if !w.closing.CompareAndSwap(false, true) { + return false, nil + } + return true, w.conn.Close() +} + +func (w *responsesWebsocketWriter) closeWithPayload(payload []byte) (bool, error) { + if w == nil || w.conn == nil { + return false, nil + } + if !w.closing.CompareAndSwap(false, true) { + return false, nil + } + if !w.writeMu.TryLock() { + return false, w.conn.Close() + } + defer w.writeMu.Unlock() + + errWrite := w.conn.WriteMessage(websocket.TextMessage, payload) + errClose := w.conn.Close() + if errWrite != nil { + return false, errWrite + } + return true, errClose +} + func (w *responsesWebsocketWriter) closeForUpstreamDisconnect(err error) { if w == nil || w.conn == nil { return @@ -146,7 +177,42 @@ func (w *responsesWebsocketWriter) closeForUpstreamDisconnect(err error) { if matched, _ := w.closeForUpstreamError(err); matched { return } - _ = w.conn.Close() + + errMsg := handlers.ExecutionErrorMessage(err) + if !shouldExposeResponsesUpstreamError(errMsg) { + _, _ = w.closeWithoutError() + return + } + payload, errBuild := buildResponsesWebsocketErrorPayload(errMsg) + if errBuild != nil { + _, _ = w.closeWithoutError() + return + } + wrote, errClose := w.closeWithPayload(payload) + if wrote { + log.Infof( + "responses websocket: downstream_out disconnect_error event=%s payload=%s", + websocketPayloadEventType(payload), + websocketPayloadPreview(payload), + ) + } + if errClose != nil && !errors.Is(errClose, websocket.ErrCloseSent) { + log.Debugf("responses websocket: upstream disconnect close failed: %v", errClose) + } +} + +// isWebsocketConnectionClosedError reports whether the error only means the +// connection was already torn down. These are expected during shutdown races +// (the proxy closes after sending a terminal frame, or the client hangs up mid +// write) and must not be logged as proxy failures. +func isWebsocketConnectionClosedError(err error) bool { + if err == nil { + return false + } + if errors.Is(err, net.ErrClosed) || errors.Is(err, websocket.ErrCloseSent) { + return true + } + return strings.Contains(err.Error(), "use of closed network connection") } func truncateWebsocketCloseReason(reason string, maxBytes int) string { @@ -243,7 +309,7 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) { log.Infof("responses websocket: upstream execution session closed id=%s", passthroughSessionID) } wsTimelineLog.SetContext(c) - if errClose := conn.Close(); errClose != nil { + if errClose := conn.Close(); errClose != nil && !isWebsocketConnectionClosedError(errClose) { log.Warnf("responses websocket: close connection error: %v", errClose) } }() @@ -549,7 +615,13 @@ func (h *OpenAIResponsesAPIHandler) ResponsesWebsocket(c *gin.Context) { ) if errForward != nil { wsTerminateErr = errForward - if !errors.Is(errForward, websocket.ErrCloseSent) { + switch { + case errors.Is(errForward, websocket.ErrCloseSent): + case isWebsocketConnectionClosedError(errForward): + // The client hung up while a downstream write was in flight. This is a + // normal shutdown race, not a proxy failure. + log.Debugf("responses websocket: client closed during forward id=%s error=%v", passthroughSessionID, errForward) + default: log.Warnf("responses websocket: forward failed id=%s error=%v", passthroughSessionID, errForward) } return diff --git a/sdk/api/handlers/openai/openai_responses_websocket_forward.go b/sdk/api/handlers/openai/openai_responses_websocket_forward.go index b465201b..3d7d31ab 100644 --- a/sdk/api/handlers/openai/openai_responses_websocket_forward.go +++ b/sdk/api/handlers/openai/openai_responses_websocket_forward.go @@ -60,21 +60,27 @@ func (h *OpenAIResponsesAPIHandler) forwardResponsesWebsocket( errs = nil continue } - if errMsg != nil { - h.LoggingAPIResponseError(context.WithValue(context.Background(), "gin", c), errMsg) - if opts.suppressError != nil && opts.suppressError(errMsg) { - cancel(errMsg.Error) - return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, nil - } - markAPIResponseTimestamp(c) - if matched, errClose := writer.closeForUpstreamError(errMsg.Error); matched { - cancel(errMsg.Error) - if errClose != nil { - return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, errClose - } - return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, websocket.ErrCloseSent + if errMsg == nil { + cancel(nil) + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), nil, nil + } + + h.LoggingAPIResponseError(context.WithValue(context.Background(), "gin", c), errMsg) + if opts.suppressError != nil && opts.suppressError(errMsg) { + cancel(errMsg.Error) + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, nil + } + markAPIResponseTimestamp(c) + if matched, errClose := writer.closeForUpstreamError(errMsg.Error); matched { + cancel(errMsg.Error) + if errClose != nil { + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, errClose } - errorPayload, errWrite := writeResponsesWebsocketError(writer, wsTimelineLog, errMsg) + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, websocket.ErrCloseSent + } + + errorPayload, wrote, errTerminate := writeResponsesWebsocketTerminalError(writer, wsTimelineLog, errMsg, nil) + if wrote { log.Infof( "responses websocket: downstream_out id=%s type=%d event=%s payload=%s", sessionID, @@ -82,23 +88,9 @@ func (h *OpenAIResponsesAPIHandler) forwardResponsesWebsocket( websocketPayloadEventType(errorPayload), websocketPayloadPreview(errorPayload), ) - if errWrite != nil { - // log.Warnf( - // "responses websocket: downstream_out write failed id=%s event=%s error=%v", - // sessionID, - // websocketPayloadEventType(errorPayload), - // errWrite, - // ) - cancel(errMsg.Error) - return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, errWrite - } - } - if errMsg != nil { - cancel(errMsg.Error) - } else { - cancel(nil) } - return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, nil + cancel(errMsg.Error) + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, errTerminate case chunk, ok := <-data: if !ok { if !completed { @@ -108,26 +100,12 @@ func (h *OpenAIResponsesAPIHandler) forwardResponsesWebsocket( } h.LoggingAPIResponseError(context.WithValue(context.Background(), "gin", c), errMsg) markAPIResponseTimestamp(c) - errorPayload, errWrite := writeResponsesWebsocketError(writer, wsTimelineLog, errMsg) - log.Infof( - "responses websocket: downstream_out id=%s type=%d event=%s payload=%s", - sessionID, - websocket.TextMessage, - websocketPayloadEventType(errorPayload), - websocketPayloadPreview(errorPayload), - ) - if errWrite != nil { - log.Warnf( - "responses websocket: downstream_out write failed id=%s event=%s error=%v", - sessionID, - websocketPayloadEventType(errorPayload), - errWrite, - ) - cancel(errMsg.Error) - return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, errWrite - } + _, errClose := writer.closeWithoutError() cancel(errMsg.Error) - return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, nil + if errClose != nil { + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, errClose + } + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), errMsg, websocket.ErrCloseSent } cancel(nil) return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), nil, nil @@ -162,6 +140,27 @@ func (h *OpenAIResponsesAPIHandler) forwardResponsesWebsocket( completedResponseID = responseCompletedIDFromPayload(payloads[i]) } markAPIResponseTimestamp(c) + if payloadErrMsg != nil { + if matched, errClose := writer.closeForUpstreamError(payloadErrMsg.Error); matched { + cancel(payloadErrMsg.Error) + if errClose != nil { + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), payloadErrMsg, errClose + } + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), payloadErrMsg, websocket.ErrCloseSent + } + errorPayload, wrote, errTerminate := writeResponsesWebsocketTerminalError(writer, wsTimelineLog, payloadErrMsg, payloads[i]) + if wrote { + log.Infof( + "responses websocket: downstream_out id=%s type=%d event=%s payload=%s", + sessionID, + websocket.TextMessage, + websocketPayloadEventType(errorPayload), + websocketPayloadPreview(errorPayload), + ) + } + cancel(payloadErrMsg.Error) + return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), payloadErrMsg, errTerminate + } // log.Infof( // "responses websocket: downstream_out id=%s type=%d event=%s payload=%s", // sessionID, @@ -179,10 +178,6 @@ func (h *OpenAIResponsesAPIHandler) forwardResponsesWebsocket( cancel(errWrite) return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), nil, errWrite } - if payloadErrMsg != nil { - cancel(payloadErrMsg.Error) - return completedOutput, completedResponseID, sortedStringSet(pendingToolCallIDs), payloadErrMsg, nil - } } } } @@ -201,6 +196,130 @@ func responsesWebsocketErrorStatus(errMsg *interfaces.ErrorMessage) int { return status } +// responsesClientFaultErrorCodes lists upstream error codes caused by the request +// payload itself. These must reach the client verbatim regardless of the HTTP +// status the upstream attached, because retrying or rotating credentials cannot +// change the outcome. +var responsesClientFaultErrorCodes = map[string]struct{}{ + "cyber_policy": {}, + "context_length_exceeded": {}, + "message_too_big": {}, + "string_above_max_length": {}, + "invalid_prompt": {}, + "invalid_value": {}, + "unsupported_value": {}, + "invalid_request_error": {}, + "previous_response_not_found": {}, +} + +// responsesClientFaultErrorTypes mirrors responsesClientFaultErrorCodes for +// upstreams that only classify the failure through `error.type`. +var responsesClientFaultErrorTypes = map[string]struct{}{ + "invalid_request": {}, + "invalid_request_error": {}, + "bad_request_error": {}, + "invalid_prompt": {}, +} + +// isResponsesClientFaultError reports whether the upstream error body identifies +// a request-shape failure. Upstreams are inconsistent about the status paired +// with these bodies: Codex reports `cyber_policy` as 400 on the stream error path +// but as 502 when the same rejection arrives through the websocket disconnect +// channel, so the body is authoritative here rather than the status. +func isResponsesClientFaultError(errMsg *interfaces.ErrorMessage) bool { + if errMsg == nil || errMsg.Error == nil { + return false + } + body := strings.TrimSpace(errMsg.Error.Error()) + if body == "" || !json.Valid([]byte(body)) { + return false + } + for _, path := range []string{"error.code", "code", "response.error.code"} { + code := strings.ToLower(strings.TrimSpace(gjson.Get(body, path).String())) + if code == "" { + continue + } + if _, ok := responsesClientFaultErrorCodes[code]; ok { + return true + } + } + for _, path := range []string{"error.type", "type", "response.error.type"} { + errType := strings.ToLower(strings.TrimSpace(gjson.Get(body, path).String())) + if errType == "" { + continue + } + if _, ok := responsesClientFaultErrorTypes[errType]; ok { + return true + } + } + return false +} + +// shouldExposeResponsesUpstreamError reports whether a terminal upstream error +// must reach the downstream client. +// +// Only request-shape failures are exposed: the client can act on them and no +// credential rotation or retry can make the request succeed. Credential, quota +// and transport failures stay silent so the client simply reconnects and retries; +// a fresh connection carries no server-side transcript, so reconnecting already +// implies a full context resend. +func shouldExposeResponsesUpstreamError(errMsg *interfaces.ErrorMessage) bool { + if errMsg == nil { + return false + } + if isResponsesClientFaultError(errMsg) { + return true + } + switch responsesWebsocketErrorStatus(errMsg) { + case http.StatusBadRequest, + http.StatusConflict, + http.StatusRequestEntityTooLarge, + http.StatusUnprocessableEntity: + return true + default: + return false + } +} + +func writeResponsesWebsocketTerminalError( + writer *responsesWebsocketWriter, + wsTimelineLog websocketTimelineAppender, + errMsg *interfaces.ErrorMessage, + payload []byte, +) ([]byte, bool, error) { + if !shouldExposeResponsesUpstreamError(errMsg) { + // Keep the upstream reason in the request-log timeline even though the client + // only observes a closed connection, otherwise silent failures are + // undiagnosable after the fact. + if wsTimelineLog != nil && errMsg != nil { + appendWebsocketTimelineDisconnect(wsTimelineLog, errMsg.Error, time.Now()) + } + _, errClose := writer.closeWithoutError() + if errClose != nil { + return nil, false, errClose + } + return nil, false, websocket.ErrCloseSent + } + + if len(payload) == 0 { + var errBuild error + payload, errBuild = buildResponsesWebsocketErrorPayload(errMsg) + if errBuild != nil { + _, _ = writer.closeWithoutError() + return nil, false, errBuild + } + } + + wrote, errClose := writer.closeWithPayload(payload) + if wrote && wsTimelineLog != nil { + wsTimelineLog.Append("response", payload, time.Now()) + } + if errClose != nil { + return payload, wrote, errClose + } + return payload, wrote, websocket.ErrCloseSent +} + func shouldReplayResponsesWebsocketPinnedAuthFailure(errMsg *interfaces.ErrorMessage) bool { switch responsesWebsocketErrorStatus(errMsg) { case http.StatusUnauthorized, http.StatusTooManyRequests: @@ -478,7 +597,7 @@ func websocketJSONPayloadsFromChunk(chunk []byte) [][]byte { return payloads } -func writeResponsesWebsocketError(writer *responsesWebsocketWriter, wsTimelineLog websocketTimelineAppender, errMsg *interfaces.ErrorMessage) ([]byte, error) { +func buildResponsesWebsocketErrorPayload(errMsg *interfaces.ErrorMessage) ([]byte, error) { status := http.StatusInternalServerError errText := http.StatusText(status) if errMsg != nil { @@ -548,5 +667,13 @@ func writeResponsesWebsocketError(writer *responsesWebsocketWriter, wsTimelineLo } } + return payload, nil +} + +func writeResponsesWebsocketError(writer *responsesWebsocketWriter, wsTimelineLog websocketTimelineAppender, errMsg *interfaces.ErrorMessage) ([]byte, error) { + payload, errBuild := buildResponsesWebsocketErrorPayload(errMsg) + if errBuild != nil { + return nil, errBuild + } return payload, writeResponsesWebsocketPayload(writer, wsTimelineLog, payload, time.Now()) } diff --git a/sdk/api/handlers/openai/openai_responses_websocket_test.go b/sdk/api/handlers/openai/openai_responses_websocket_test.go index f1257817..9c7fe991 100644 --- a/sdk/api/handlers/openai/openai_responses_websocket_test.go +++ b/sdk/api/handlers/openai/openai_responses_websocket_test.go @@ -340,6 +340,54 @@ func TestResponsesWebsocketWriterCloseDoesNotWaitForActiveDataWriter(t *testing. } } +func TestResponsesWebsocketGenericDisconnectDoesNotWaitForActiveDataWriter(t *testing.T) { + serverErrCh := make(chan error, 1) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := responsesWebsocketUpgrader.Upgrade(w, r, nil) + if err != nil { + serverErrCh <- err + return + } + writer := newResponsesWebsocketWriter(conn) + + writer.writeMu.Lock() + closeDone := make(chan struct{}) + go func() { + writer.closeForUpstreamDisconnect(&websocket.CloseError{ + Code: websocket.CloseAbnormalClosure, + Text: "unexpected EOF", + }) + close(closeDone) + }() + + select { + case <-closeDone: + writer.writeMu.Unlock() + serverErrCh <- nil + case <-time.After(time.Second): + writer.writeMu.Unlock() + serverErrCh <- errors.New("generic disconnect waited behind active data writer") + } + })) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer func() { _ = conn.Close() }() + if err = conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatalf("set read deadline: %v", err) + } + if _, _, err = conn.ReadMessage(); err == nil { + t.Fatal("client read succeeded, want connection closure") + } + if errServer := <-serverErrCh; errServer != nil { + t.Fatalf("server error: %v", errServer) + } +} + func TestTruncateWebsocketCloseReason(t *testing.T) { tests := []struct { name string @@ -476,6 +524,71 @@ func TestForwardResponsesWebsocketMirrorsMappedMessageTooBig(t *testing.T) { } } +func TestForwardResponsesWebsocketMirrorsPayloadMessageTooBig(t *testing.T) { + gin.SetMode(gin.TestMode) + + serverErrCh := make(chan error, 1) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := responsesWebsocketUpgrader.Upgrade(w, r, nil) + if err != nil { + serverErrCh <- err + return + } + defer func() { _ = conn.Close() }() + + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Request = r + data := make(chan []byte, 1) + errCh := make(chan *interfaces.ErrorMessage) + data <- []byte(`{"type":"error","status":413,"error":{"message":"upstream websocket message too big","code":"message_too_big"}}`) + close(data) + close(errCh) + + h := NewOpenAIResponsesAPIHandler(handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, nil)) + _, _, _, errMsg, errForward := h.forwardResponsesWebsocket( + ctx, + newResponsesWebsocketWriter(conn), + func(...interface{}) {}, + data, + errCh, + newInMemoryWebsocketTimelineLog(), + "session-1", + ) + if errMsg == nil || errMsg.StatusCode != http.StatusRequestEntityTooLarge { + serverErrCh <- fmt.Errorf("forward error message = %#v, want status %d", errMsg, http.StatusRequestEntityTooLarge) + return + } + if !errors.Is(errForward, websocket.ErrCloseSent) { + serverErrCh <- fmt.Errorf("forward error = %v, want %v", errForward, websocket.ErrCloseSent) + return + } + serverErrCh <- nil + })) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer func() { _ = conn.Close() }() + + if err = conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatalf("set read deadline: %v", err) + } + _, _, err = conn.ReadMessage() + var closeErr *websocket.CloseError + if !errors.As(err, &closeErr) { + t.Fatalf("expected websocket close error, got %v", err) + } + if closeErr.Code != websocket.CloseMessageTooBig { + t.Fatalf("close code = %d, want %d", closeErr.Code, websocket.CloseMessageTooBig) + } + if err = <-serverErrCh; err != nil { + t.Fatalf("server error: %v", err) + } +} + type websocketCaptureExecutor struct { streamCalls int payloads [][]byte @@ -2253,6 +2366,91 @@ func TestForwardResponsesWebsocketTreatsResponseDoneAsTerminalWithoutRewriting(t } } +func TestShouldExposeResponsesUpstreamError(t *testing.T) { + tests := []struct { + status int + want bool + }{ + {status: http.StatusBadRequest, want: true}, + {status: http.StatusConflict, want: true}, + {status: http.StatusRequestEntityTooLarge, want: true}, + {status: http.StatusUnprocessableEntity, want: true}, + {status: http.StatusUnauthorized}, + {status: http.StatusRequestTimeout}, + {status: http.StatusTooManyRequests}, + {status: http.StatusInternalServerError}, + } + + for _, tc := range tests { + t.Run(strconv.Itoa(tc.status), func(t *testing.T) { + errMsg := &interfaces.ErrorMessage{StatusCode: tc.status, Error: errors.New(http.StatusText(tc.status))} + if got := shouldExposeResponsesUpstreamError(errMsg); got != tc.want { + t.Fatalf("shouldExposeResponsesUpstreamError(%d) = %t, want %t", tc.status, got, tc.want) + } + }) + } +} + +// TestResponsesUpstreamErrorBodyDrivesExposure pins that the error body, not the +// attached status, decides whether a request-shape failure is exposed. Codex +// reports the same cyber_policy rejection as 400 on the stream error path and as +// 502 through the websocket disconnect channel. +func TestResponsesUpstreamErrorBodyDrivesExposure(t *testing.T) { + tests := []struct { + name string + status int + body string + want bool + }{ + {name: "bad request", status: http.StatusBadRequest, body: "bad request", want: true}, + {name: "conflict", status: http.StatusConflict, body: "conflict", want: true}, + {name: "entity too large", status: http.StatusRequestEntityTooLarge, body: "too large", want: true}, + {name: "unprocessable", status: http.StatusUnprocessableEntity, body: "unprocessable", want: true}, + { + name: "cyber policy at 502", + status: http.StatusBadGateway, + body: `{"error":{"type":"invalid_request","code":"cyber_policy","message":"flagged"}}`, + want: true, + }, + { + name: "context length exceeded at 500", + status: http.StatusInternalServerError, + body: `{"error":{"type":"invalid_request_error","code":"context_length_exceeded","message":"too long"}}`, + want: true, + }, + // Credential, quota and transport failures stay silent: the client just + // reconnects, and a fresh socket already implies a full context resend. + {name: "unauthorized", status: http.StatusUnauthorized, body: "invalid token"}, + {name: "payment required", status: http.StatusPaymentRequired, body: "insufficient credits"}, + {name: "forbidden", status: http.StatusForbidden, body: "forbidden"}, + {name: "too many requests", status: http.StatusTooManyRequests, body: "usage limit reached"}, + {name: "request timeout", status: http.StatusRequestTimeout, body: "timeout"}, + {name: "bad gateway", status: http.StatusBadGateway, body: "bad gateway"}, + { + name: "upstream websocket drop", + status: http.StatusInternalServerError, + body: `{"error":{"message":"websocket: close 1006 (abnormal closure): unexpected EOF","type":"server_error","code":"internal_server_error"}}`, + }, + {name: "no error message", status: 0, body: ""}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + errMsg := &interfaces.ErrorMessage{StatusCode: tc.status} + if tc.body != "" { + errMsg.Error = errors.New(tc.body) + } + if got := shouldExposeResponsesUpstreamError(errMsg); got != tc.want { + t.Fatalf("shouldExposeResponsesUpstreamError = %t, want %t", got, tc.want) + } + }) + } + + if shouldExposeResponsesUpstreamError(nil) { + t.Fatal("nil error message must not be exposed") + } +} + func TestForwardResponsesWebsocketTreatsErrorPayloadAsTerminal(t *testing.T) { gin.SetMode(gin.TestMode) @@ -2263,19 +2461,14 @@ func TestForwardResponsesWebsocketTreatsErrorPayloadAsTerminal(t *testing.T) { serverErrCh <- err return } - defer func() { - errClose := conn.Close() - if errClose != nil { - serverErrCh <- errClose - } - }() + defer func() { _ = conn.Close() }() ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) ctx.Request = r data := make(chan []byte, 1) errCh := make(chan *interfaces.ErrorMessage) - data <- []byte(`{"type":"error","status":429,"error":{"message":"upstream failed"}}`) + data <- []byte(`{"type":"error","status":400,"error":{"type":"invalid_request_error","message":"invalid request"}}`) close(data) close(errCh) @@ -2288,7 +2481,7 @@ func TestForwardResponsesWebsocketTreatsErrorPayloadAsTerminal(t *testing.T) { newInMemoryWebsocketTimelineLog(), "session-1", ) - if err != nil { + if err != nil && !errors.Is(err, websocket.ErrCloseSent) { serverErrCh <- err return } @@ -2296,12 +2489,12 @@ func TestForwardResponsesWebsocketTreatsErrorPayloadAsTerminal(t *testing.T) { serverErrCh <- errors.New("expected websocket error message") return } - if errMsg.StatusCode != http.StatusTooManyRequests { - serverErrCh <- fmt.Errorf("websocket error status = %d, want %d", errMsg.StatusCode, http.StatusTooManyRequests) + if errMsg.StatusCode != http.StatusBadRequest { + serverErrCh <- fmt.Errorf("websocket error status = %d, want %d", errMsg.StatusCode, http.StatusBadRequest) return } - if errMsg.Error == nil || !strings.Contains(errMsg.Error.Error(), "upstream failed") { - serverErrCh <- fmt.Errorf("websocket error = %v, want upstream failed", errMsg.Error) + if errMsg.Error == nil || !strings.Contains(errMsg.Error.Error(), "invalid request") { + serverErrCh <- fmt.Errorf("websocket error = %v, want invalid request", errMsg.Error) return } serverErrCh <- nil @@ -2523,6 +2716,275 @@ func TestResponsesWebsocketMirrorsUpstreamMessageTooBigDisconnect(t *testing.T) } } +func TestResponsesWebsocketSendsJSONErrorOnUpstreamCyberPolicyDisconnect(t *testing.T) { + gin.SetMode(gin.TestMode) + + executor := &websocketUpstreamDisconnectExecutor{provider: "codex", subscribed: make(chan string, 1)} + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(executor) + base := handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, manager) + h := NewOpenAIResponsesAPIHandler(base) + + router := gin.New() + router.GET("/v1/responses/ws", h.ResponsesWebsocket) + server := httptest.NewServer(router) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/v1/responses/ws" + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer func() { _ = conn.Close() }() + + var sessionID string + select { + case sessionID = <-executor.subscribed: + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for upstream disconnect subscription") + } + + cyberPolicyErr := websocketPinnedFailoverStatusError{ + status: http.StatusBadRequest, + msg: `{"error":{"type":"invalid_request","code":"cyber_policy","message":"This content was flagged for possible cybersecurity risk. If this seems wrong, try rephrasing your request. To get authorized for security work, join the Trusted Access for Cyber program: https://chatgpt.com/cyber","param":null}}`, + } + executor.TriggerDisconnect(sessionID, cyberPolicyErr) + + _ = conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + msgType, payload, err := conn.ReadMessage() + if err != nil { + t.Fatalf("expected downstream text error payload before socket close, got read error: %v", err) + } + if msgType != websocket.TextMessage { + t.Fatalf("msgType = %d, want TextMessage (%d)", msgType, websocket.TextMessage) + } + + if gjson.GetBytes(payload, "type").String() != "error" { + t.Fatalf("payload type = %q, want %q", gjson.GetBytes(payload, "type").String(), "error") + } + if status := int(gjson.GetBytes(payload, "status").Int()); status != http.StatusBadRequest { + t.Fatalf("status = %d, want %d", status, http.StatusBadRequest) + } + if gjson.GetBytes(payload, "error.code").String() != "cyber_policy" { + t.Fatalf("error.code = %q, want %q", gjson.GetBytes(payload, "error.code").String(), "cyber_policy") + } + if !strings.Contains(gjson.GetBytes(payload, "error.message").String(), "cybersecurity risk") { + t.Fatalf("error.message = %q, want cybersecurity risk text", gjson.GetBytes(payload, "error.message").String()) + } + if _, duplicate, errRead := conn.ReadMessage(); errRead == nil { + t.Fatalf("received duplicate error frame: %s", duplicate) + } +} + +func TestResponsesWebsocketHidesNonClientUpstreamDisconnectErrors(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + err error + }{ + { + name: "abnormal closure", + err: &websocket.CloseError{ + Code: websocket.CloseAbnormalClosure, + Text: "unexpected EOF", + }, + }, + { + name: "upstream read timeout", + err: errors.New("read tcp 198.18.0.1:53030->145.223.58.12:6281: i/o timeout"), + }, + { + // Credential failover already ran and lost; the client only needs to + // reconnect, so no downstream error is produced. + name: "quota exhausted", + err: websocketPinnedFailoverStatusError{ + status: http.StatusTooManyRequests, + msg: `{"error":{"type":"usage_limit_reached","message":"The usage limit has been reached"}}`, + }, + }, + { + name: "credential rejected", + err: websocketPinnedFailoverStatusError{ + status: http.StatusUnauthorized, + msg: `{"error":{"type":"authentication_error","message":"Invalid token"}}`, + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + executor := &websocketUpstreamDisconnectExecutor{provider: "codex", subscribed: make(chan string, 1)} + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(executor) + base := handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, manager) + h := NewOpenAIResponsesAPIHandler(base) + + router := gin.New() + router.GET("/v1/responses/ws", h.ResponsesWebsocket) + server := httptest.NewServer(router) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/v1/responses/ws" + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer func() { _ = conn.Close() }() + + var sessionID string + select { + case sessionID = <-executor.subscribed: + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for upstream disconnect subscription") + } + executor.TriggerDisconnect(sessionID, tc.err) + + _ = conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, payload, errRead := conn.ReadMessage() + if errRead == nil { + t.Fatalf("non-client upstream error was exposed: %s", payload) + } + // Nothing may be written downstream: no error frame and no close frame + // carrying proxy-internal detail, so the client just reconnects. + var closeErr *websocket.CloseError + if errors.As(errRead, &closeErr) && closeErr.Code != websocket.CloseAbnormalClosure { + t.Fatalf("non-client upstream error produced a close frame: %#v", closeErr) + } + }) + } +} + +// TestResponsesWebsocketExposesCyberPolicyRegardlessOfStatus pins the other +// production shape from main.log: the identical cyber_policy rejection arrives +// with status 400 on the stream path and 502 through the disconnect channel. Both +// must reach the client, because no credential rotation can satisfy the request. +func TestResponsesWebsocketExposesCyberPolicyRegardlessOfStatus(t *testing.T) { + gin.SetMode(gin.TestMode) + + const cyberPolicyBody = `{"error":{"type":"invalid_request","code":"cyber_policy","message":"This content was flagged for possible cybersecurity risk.","param":null}}` + + for _, status := range []int{http.StatusBadRequest, http.StatusBadGateway, http.StatusInternalServerError} { + t.Run(strconv.Itoa(status), func(t *testing.T) { + executor := &websocketUpstreamDisconnectExecutor{provider: "codex", subscribed: make(chan string, 1)} + manager := coreauth.NewManager(nil, nil, nil) + manager.RegisterExecutor(executor) + base := handlers.NewBaseAPIHandlers(&sdkconfig.SDKConfig{}, manager) + h := NewOpenAIResponsesAPIHandler(base) + + router := gin.New() + router.GET("/v1/responses/ws", h.ResponsesWebsocket) + server := httptest.NewServer(router) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/v1/responses/ws" + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer func() { _ = conn.Close() }() + + var sessionID string + select { + case sessionID = <-executor.subscribed: + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for upstream disconnect subscription") + } + executor.TriggerDisconnect(sessionID, websocketPinnedFailoverStatusError{status: status, msg: cyberPolicyBody}) + + _ = conn.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, payload, errRead := conn.ReadMessage() + if errRead != nil { + t.Fatalf("cyber_policy rejection was hidden at status %d: %v", status, errRead) + } + if got := gjson.GetBytes(payload, "error.code").String(); got != "cyber_policy" { + t.Fatalf("error.code = %q, want cyber_policy: %s", got, payload) + } + }) + } +} + +func TestResponsesWebsocketTerminalErrorWrittenOnceAcrossForwardAndDisconnect(t *testing.T) { + serverErrCh := make(chan error, 1) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := responsesWebsocketUpgrader.Upgrade(w, r, nil) + if err != nil { + serverErrCh <- err + return + } + writer := newResponsesWebsocketWriter(conn) + cyberPolicyErr := websocketPinnedFailoverStatusError{ + status: http.StatusBadRequest, + msg: `{"error":{"type":"invalid_request","code":"cyber_policy","message":"blocked"}}`, + } + errMsg := &interfaces.ErrorMessage{ + StatusCode: http.StatusBadRequest, + Error: cyberPolicyErr, + } + + start := make(chan struct{}) + resultCh := make(chan error, 2) + var wg sync.WaitGroup + wg.Add(3) + go func() { + defer wg.Done() + <-start + payload, _, errWrite := writeResponsesWebsocketTerminalError(writer, nil, errMsg, nil) + if !errors.Is(errWrite, websocket.ErrCloseSent) || gjson.GetBytes(payload, "error.code").String() != "cyber_policy" { + resultCh <- fmt.Errorf("err-channel terminal write failed: err=%v payload=%s", errWrite, payload) + } + }() + go func() { + defer wg.Done() + <-start + payload := []byte(`{"type":"error","status":400,"error":{"type":"invalid_request","code":"cyber_policy","message":"blocked"}}`) + writtenPayload, _, errWrite := writeResponsesWebsocketTerminalError(writer, nil, errMsg, payload) + if !errors.Is(errWrite, websocket.ErrCloseSent) || gjson.GetBytes(writtenPayload, "error.code").String() != "cyber_policy" { + resultCh <- fmt.Errorf("payload terminal write failed: err=%v payload=%s", errWrite, writtenPayload) + } + }() + go func() { + defer wg.Done() + <-start + writer.closeForUpstreamDisconnect(errMsg.Error) + }() + close(start) + wg.Wait() + select { + case errResult := <-resultCh: + serverErrCh <- errResult + default: + serverErrCh <- nil + } + })) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil) + if err != nil { + t.Fatalf("dial websocket: %v", err) + } + defer func() { _ = conn.Close() }() + + textFrames := 0 + for { + _, payload, errRead := conn.ReadMessage() + if errRead != nil { + break + } + textFrames++ + if got := gjson.GetBytes(payload, "error.code").String(); got != "cyber_policy" { + t.Fatalf("terminal error code = %q, want cyber_policy: %s", got, payload) + } + } + if textFrames != 1 { + t.Fatalf("terminal error frame count = %d, want 1", textFrames) + } + if errServer := <-serverErrCh; errServer != nil { + t.Fatalf("server error: %v", errServer) + } +} + func TestResponsesWebsocketCodexWebsocketPassthroughPassesCompactedRequestWithoutTranscriptMerge(t *testing.T) { gin.SetMode(gin.TestMode) @@ -2766,7 +3228,7 @@ func TestResponsesWebsocketFullRequestCanRouteFromNativeWebsocketToBuiltInProvid } } -func TestResponsesWebsocketFailedProviderRoutePreservesNativeWebsocketPin(t *testing.T) { +func TestResponsesWebsocketHidesProviderRouteAuthFailure(t *testing.T) { gin.SetMode(gin.TestMode) const sourceModel = "codex-provider-route-failure-source" @@ -2805,31 +3267,28 @@ func TestResponsesWebsocketFailedProviderRoutePreservesNativeWebsocketPin(t *tes } defer func() { _ = conn.Close() }() - requests := []string{ - fmt.Sprintf(`{"type":"response.create","model":%q,"input":[{"type":"message","id":"msg-1"}]}`, sourceModel), - fmt.Sprintf(`{"type":"response.create","model":%q,"route_to_claude":true,"input":[{"type":"message","id":"msg-routed"}]}`, sourceModel), - `{"type":"response.create","previous_response_id":"resp-1","input":[{"type":"message","id":"msg-2"}]}`, + firstRequest := fmt.Sprintf(`{"type":"response.create","model":%q,"input":[{"type":"message","id":"msg-1"}]}`, sourceModel) + if errWrite := conn.WriteMessage(websocket.TextMessage, []byte(firstRequest)); errWrite != nil { + t.Fatalf("write first request: %v", errWrite) } - wantTypes := []string{wsEventTypeCompleted, wsEventTypeError, wsEventTypeCompleted} - for i, request := range requests { - if errWrite := conn.WriteMessage(websocket.TextMessage, []byte(request)); errWrite != nil { - t.Fatalf("write request %d: %v", i+1, errWrite) - } - _, response, errRead := conn.ReadMessage() - if errRead != nil { - t.Fatalf("read response %d: %v", i+1, errRead) - } - if got := gjson.GetBytes(response, "type").String(); got != wantTypes[i] { - t.Fatalf("response %d type = %q, want %q: %s", i+1, got, wantTypes[i], response) - } + _, firstResponse, errRead := conn.ReadMessage() + if errRead != nil { + t.Fatalf("read first response: %v", errRead) + } + if got := gjson.GetBytes(firstResponse, "type").String(); got != wsEventTypeCompleted { + t.Fatalf("first response type = %q, want %q: %s", got, wsEventTypeCompleted, firstResponse) } - codexPayloads := codexExecutor.Payloads() - if len(codexPayloads) != 2 { - t.Fatalf("codex payload count = %d, want 2", len(codexPayloads)) + routedRequest := fmt.Sprintf(`{"type":"response.create","model":%q,"route_to_claude":true,"input":[{"type":"message","id":"msg-routed"}]}`, sourceModel) + if errWrite := conn.WriteMessage(websocket.TextMessage, []byte(routedRequest)); errWrite != nil { + t.Fatalf("write routed request: %v", errWrite) + } + if _, response, errRead := conn.ReadMessage(); errRead == nil { + t.Fatalf("credential error was exposed to the client: %s", response) } - if got := gjson.GetBytes(codexPayloads[1], "previous_response_id").String(); got != "resp-1" { - t.Fatalf("resumed codex previous_response_id = %q, want resp-1: %s", got, codexPayloads[1]) + + if got := len(codexExecutor.Payloads()); got != 1 { + t.Fatalf("codex payload count = %d, want 1", got) } if got := len(claudeExecutor.Payloads()); got != 1 { t.Fatalf("claude payload count = %d, want 1", got) @@ -3106,7 +3565,7 @@ func TestResponsesWebsocketRejectsUnknownPreviousResponseOnNewSocket(t *testing. } } -func TestResponsesWebsocketRollsBackCanonicalTranscriptAfterNonRetryableError(t *testing.T) { +func TestResponsesWebsocketClosesAfterNonRetryableClientError(t *testing.T) { gin.SetMode(gin.TestMode) modelName := "xai-websocket-rollback-model" @@ -3141,48 +3600,35 @@ func TestResponsesWebsocketRollsBackCanonicalTranscriptAfterNonRetryableError(t } defer func() { _ = conn.Close() }() - requests := []string{ - fmt.Sprintf(`{"type":"response.create","model":%q,"input":[{"type":"message","id":"msg-1"}]}`, modelName), - `{"type":"response.create","previous_response_id":"resp-1","input":[{"type":"function_call","id":"fc-failed","call_id":"failed-call","name":"failed_tool","arguments":"{}"},{"type":"function_call_output","id":"fco-failed","call_id":"failed-call","output":"must-not-survive"}]}`, - `{"type":"response.create","previous_response_id":"resp-1","input":[{"type":"message","id":"msg-3"},{"type":"function_call","id":"fc-retry","call_id":"failed-call","name":"failed_tool","arguments":"{}"}]}`, + firstRequest := fmt.Sprintf(`{"type":"response.create","model":%q,"input":[{"type":"message","id":"msg-1"}]}`, modelName) + if errWrite := conn.WriteMessage(websocket.TextMessage, []byte(firstRequest)); errWrite != nil { + t.Fatalf("write first websocket message: %v", errWrite) } - wantTypes := []string{wsEventTypeCompleted, wsEventTypeError, wsEventTypeCompleted} - for i := range requests { - if errWrite := conn.WriteMessage(websocket.TextMessage, []byte(requests[i])); errWrite != nil { - t.Fatalf("write websocket message %d: %v", i+1, errWrite) - } - _, payload, errRead := conn.ReadMessage() - if errRead != nil { - t.Fatalf("read websocket response %d: %v", i+1, errRead) - } - if got := gjson.GetBytes(payload, "type").String(); got != wantTypes[i] { - t.Fatalf("response %d type = %q, want %q: %s", i+1, got, wantTypes[i], payload) - } + _, firstResponse, errRead := conn.ReadMessage() + if errRead != nil || gjson.GetBytes(firstResponse, "type").String() != wsEventTypeCompleted { + t.Fatalf("first websocket response = %s, err=%v", firstResponse, errRead) } - payloads := executor.Payloads() - if len(payloads) != 3 { - t.Fatalf("executor payload count = %d, want 3", len(payloads)) + failedRequest := `{"type":"response.create","previous_response_id":"resp-1","input":[{"type":"function_call","id":"fc-failed","call_id":"failed-call","name":"failed_tool","arguments":"{}"},{"type":"function_call_output","id":"fco-failed","call_id":"failed-call","output":"must-not-survive"}]}` + if errWrite := conn.WriteMessage(websocket.TextMessage, []byte(failedRequest)); errWrite != nil { + t.Fatalf("write failed websocket message: %v", errWrite) } - third := payloads[2] - if gjson.GetBytes(third, "previous_response_id").Exists() { - t.Fatalf("retry payload must not depend on previous_response_id: %s", third) + _, errorResponse, errRead := conn.ReadMessage() + if errRead != nil { + t.Fatalf("read client error response: %v", errRead) } - input := gjson.GetBytes(third, "input").Array() - if len(input) != 3 { - t.Fatalf("retry canonical input len = %d, want 3: %s", len(input), third) + if got := gjson.GetBytes(errorResponse, "type").String(); got != wsEventTypeError { + t.Fatalf("client error response type = %q, want %q: %s", got, wsEventTypeError, errorResponse) } - wantIDs := []string{"msg-1", "out-1", "msg-3"} - for i, wantID := range wantIDs { - if got := input[i].Get("id").String(); got != wantID { - t.Fatalf("retry canonical input[%d].id = %q, want %q: %s", i, got, wantID, third) - } + if got := int(gjson.GetBytes(errorResponse, "status").Int()); got != http.StatusBadRequest { + t.Fatalf("client error response status = %d, want %d: %s", got, http.StatusBadRequest, errorResponse) } - if bytes.Contains(third, []byte(`"id":"fc-failed"`)) || bytes.Contains(third, []byte(`"id":"fco-failed"`)) { - t.Fatalf("failed turn leaked into retry transcript: %s", third) + if _, duplicate, errRead := conn.ReadMessage(); errRead == nil { + t.Fatalf("received frame after terminal client error: %s", duplicate) } - if bytes.Contains(third, []byte(`"call_id":"failed-call"`)) || bytes.Contains(third, []byte("must-not-survive")) { - t.Fatalf("failed turn contaminated tool repair cache: %s", third) + + if got := len(executor.Payloads()); got != 2 { + t.Fatalf("executor payload count = %d, want 2", got) } } @@ -4294,29 +4740,12 @@ func TestResponsesWebsocketReleasesPinnedAuthAfterStreamClosed408(t *testing.T) for { _, payload, errRead := conn.ReadMessage() if errRead != nil { - t.Fatalf("read stream-closed response: %v", errRead) - } - eventType := gjson.GetBytes(payload, "type").String() - if eventType == wsEventTypeError { - if got := int(gjson.GetBytes(payload, "status").Int()); got != http.StatusRequestTimeout { - t.Fatalf("stream-closed status = %d, want %d: %s", got, http.StatusRequestTimeout, payload) - } break } - if eventType == wsEventTypeCompleted { - t.Fatalf("stream-closed turn unexpectedly completed: %s", payload) + if gjson.GetBytes(payload, "type").String() == wsEventTypeError { + t.Fatalf("stream transport failure was exposed to the client: %s", payload) } } - - thirdDelta := `{"type":"response.create","previous_response_id":"resp-auth-a-1","input":[{"type":"message","id":"msg-3"}]}` - if errWrite := conn.WriteMessage(websocket.TextMessage, []byte(thirdDelta)); errWrite != nil { - t.Fatalf("write third websocket message: %v", errWrite) - } - _, _, errReadClose := conn.ReadMessage() - var replayClose *websocket.CloseError - if !errors.As(errReadClose, &replayClose) || replayClose.Code != websocket.CloseServiceRestart { - t.Fatalf("third websocket response error = %v, want replay close", errReadClose) - } if got := executor.AuthIDs(); len(got) != 2 || got[0] != "auth-a" || got[1] != "auth-a" { t.Fatalf("selected auth IDs before replay = %v, want [auth-a auth-a]", got) } @@ -4326,7 +4755,7 @@ func TestResponsesWebsocketReleasesPinnedAuthAfterStreamClosed408(t *testing.T) t.Fatalf("dial replay websocket: %v", errDialReplay) } defer func() { _ = replayConn.Close() }() - fullReplay := `{"type":"response.create","model":"stream-model","input":[{"type":"message","id":"msg-1"},{"type":"message","id":"out-auth-a-1"},{"type":"message","id":"msg-3"}]}` + fullReplay := `{"type":"response.create","model":"stream-model","input":[{"type":"message","id":"msg-1"},{"type":"message","id":"out-auth-a-1"},{"type":"message","id":"msg-2"}]}` if errWrite := replayConn.WriteMessage(websocket.TextMessage, []byte(fullReplay)); errWrite != nil { t.Fatalf("write full replay: %v", errWrite) } diff --git a/sdk/api/handlers/openai/openai_responses_websocket_timeline.go b/sdk/api/handlers/openai/openai_responses_websocket_timeline.go index 8126857e..1be849ba 100644 --- a/sdk/api/handlers/openai/openai_responses_websocket_timeline.go +++ b/sdk/api/handlers/openai/openai_responses_websocket_timeline.go @@ -221,6 +221,25 @@ func isResponsesWebsocketCompletionEvent(eventType string) bool { return eventType == wsEventTypeCompleted || eventType == wsEventTypeDone } +type responsesWebsocketPayloadError struct { + status int + payload []byte +} + +func (e *responsesWebsocketPayloadError) Error() string { + if e == nil { + return "" + } + return string(e.payload) +} + +func (e *responsesWebsocketPayloadError) StatusCode() int { + if e == nil { + return 0 + } + return e.status +} + func responsesWebsocketErrorMessageFromPayload(payload []byte) *interfaces.ErrorMessage { status := int(gjson.GetBytes(payload, "status").Int()) if status <= 0 { @@ -230,17 +249,17 @@ func responsesWebsocketErrorMessageFromPayload(payload []byte) *interfaces.Error status = http.StatusInternalServerError } - errText := strings.TrimSpace(gjson.GetBytes(payload, "error.message").String()) - if errText == "" { - errText = strings.TrimSpace(gjson.GetBytes(payload, "message").String()) - } - if errText == "" { - errText = strings.TrimSpace(string(payload)) - } - if errText == "" { - errText = http.StatusText(status) + trimmedPayload := bytes.TrimSpace(payload) + if len(trimmedPayload) > 0 { + return &interfaces.ErrorMessage{ + StatusCode: status, + Error: &responsesWebsocketPayloadError{ + status: status, + payload: bytes.Clone(trimmedPayload), + }, + } } - return &interfaces.ErrorMessage{StatusCode: status, Error: fmt.Errorf("%s", errText)} + return &interfaces.ErrorMessage{StatusCode: status, Error: fmt.Errorf("%s", http.StatusText(status))} } func setWebsocketTimelineBody(c *gin.Context, body string) { diff --git a/sdk/cliproxy/auth/conductor_cooldown.go b/sdk/cliproxy/auth/conductor_cooldown.go index dd7ddc8a..2820ef4a 100644 --- a/sdk/cliproxy/auth/conductor_cooldown.go +++ b/sdk/cliproxy/auth/conductor_cooldown.go @@ -1532,17 +1532,53 @@ func isMissingModelPhrase(value string) bool { } } +func isCyberPolicyError(err error) bool { + if err == nil { + return false + } + var payload any + if errJSON := json.Unmarshal([]byte(strings.TrimSpace(err.Error())), &payload); errJSON != nil { + return false + } + return containsStructuredErrorCode(payload, "cyber_policy") +} + +func containsStructuredErrorCode(value any, target string) bool { + target = strings.ToLower(strings.TrimSpace(target)) + switch typed := value.(type) { + case map[string]any: + for key, item := range typed { + if strings.EqualFold(strings.TrimSpace(key), "code") { + if code, ok := item.(string); ok && strings.ToLower(strings.TrimSpace(code)) == target { + return true + } + } + if containsStructuredErrorCode(item, target) { + return true + } + } + case []any: + for _, item := range typed { + if containsStructuredErrorCode(item, target) { + return true + } + } + } + return false +} + // isRequestInvalidError returns true if the error represents a client request -// error that should not be retried. Specifically, it treats 400 responses with -// "invalid_request_error", request-scoped 404 item misses caused by `store=false`, -// and all 422 responses as request-shape failures, where switching auths or -// pooled upstream models will not help. Model-support errors are excluded so -// routing can fall through to another auth or upstream. +// error that should not be retried. Specifically, it treats cyber policy +// rejections, 400 responses with "invalid_request_error", request-scoped 404 +// item misses caused by `store=false`, 413 payload/frame size rejections, and all +// 422 responses as request-shape failures, where switching auths or pooled +// upstream models will not help. Model-support errors are excluded so routing can +// fall through to another auth or upstream. func isRequestInvalidError(err error) bool { if err == nil { return false } - if isRequestScopedError(err) { + if isRequestScopedError(err) || isCyberPolicyError(err) { return true } if isCloudflareChallengeError(err) { @@ -1564,6 +1600,11 @@ func isRequestInvalidError(err error) bool { strings.Contains(msg, "FAILED_PRECONDITION") case http.StatusNotFound: return isRequestScopedNotFoundMessage(err.Error()) + case http.StatusRequestEntityTooLarge: + // The request payload (or websocket frame) is too large for the upstream. + // Every other credential enforces the same limit, so retrying elsewhere only + // burns the pool and marks healthy credentials unavailable. + return true case http.StatusUnprocessableEntity: return true case http.StatusInternalServerError: diff --git a/sdk/cliproxy/auth/conductor_overrides_test.go b/sdk/cliproxy/auth/conductor_overrides_test.go index 110506b6..8a048434 100644 --- a/sdk/cliproxy/auth/conductor_overrides_test.go +++ b/sdk/cliproxy/auth/conductor_overrides_test.go @@ -1189,6 +1189,16 @@ func TestManager_RequestScopedErrorStopsCredentialFallbackWithoutSuspendingAuth( HTTPStatus: http.StatusBadRequest, Message: `{"error":{"type":"bad_request_error","code":"invalid_value","message":"Bad input."}}`, } + cyberPolicyErr := &Error{ + HTTPStatus: http.StatusBadGateway, + Message: `{"error":{"type":"invalid_request","code":"cyber_policy","message":"This content was flagged for possible cybersecurity risk."}}`, + } + // A frame/payload that exceeds the upstream size limit fails identically on + // every credential, so it must not rotate or punish the pool. + tooLargeErr := &Error{ + HTTPStatus: http.StatusRequestEntityTooLarge, + Message: `{"error":{"code":"message_too_big","message":"upstream websocket message too big"}}`, + } tests := []struct { name string provider string @@ -1204,6 +1214,9 @@ func TestManager_RequestScopedErrorStopsCredentialFallbackWithoutSuspendingAuth( {name: "streaming invalid request", stream: true, err: invalidRequestErr, wantStatus: http.StatusBadRequest}, {name: "non-streaming bad request", err: badRequestErr, wantStatus: http.StatusBadRequest}, {name: "streaming bad request", stream: true, err: badRequestErr, wantStatus: http.StatusBadRequest}, + {name: "streaming cyber policy", provider: "codex", stream: true, err: cyberPolicyErr, wantStatus: http.StatusBadGateway}, + {name: "non-streaming message too big", provider: "codex", err: tooLargeErr, wantStatus: http.StatusRequestEntityTooLarge}, + {name: "streaming message too big", provider: "codex", stream: true, err: tooLargeErr, wantStatus: http.StatusRequestEntityTooLarge}, } for _, tc := range tests { @@ -1283,6 +1296,16 @@ func TestManager_RequestScopedErrorStopsCredentialFallbackWithoutSuspendingAuth( if state := updatedBad.ModelStates[model]; state != nil { t.Fatalf("expected request-scoped error to avoid model cooldown state, got %#v", state) } + if updatedBad.Failed != 1 { + t.Fatalf("failed count = %d, want 1", updatedBad.Failed) + } + updatedGood, ok := m.GetByID(goodAuth.ID) + if !ok || updatedGood == nil { + t.Fatal("expected good auth to remain registered") + } + if updatedGood.Failed != 0 { + t.Fatalf("fallback auth failed count = %d, want 0", updatedGood.Failed) + } }) } }