diff --git a/agent/error_classifier.py b/agent/error_classifier.py index d471e781b4..472e53d622 100644 --- a/agent/error_classifier.py +++ b/agent/error_classifier.py @@ -1399,9 +1399,32 @@ def _headers_of(exc: Any) -> Any: return headers if headers and hasattr(headers, "get") else None +def _status_code_from_body(body: Any) -> Optional[int]: + """Numeric HTTP status (100-599) from ``error.code``/``code`` in a structured body. + An aggregator/relay can deliver the upstream failure only this way — as an + error object inside an HTTP-200 SSE stream — leaving the SDK to raise a + status-less ``APIError`` whose ``body`` carries the status (#121270). String + codes stay symbolic (``_code_from_payload``'s ``"400" is not a code``).""" + if not isinstance(body, dict): + return None + candidates = [] + error_obj = body.get("error") + if isinstance(error_obj, dict): + candidates.extend(error_obj.get(k) for k in ("status_code", "status", "http_status", "code")) + candidates.append(body.get("code")) + return next( + (c for c in candidates if isinstance(c, int) and not isinstance(c, bool) and 100 <= c < 600), + None, + ) + + def _extract_status_code(error: Exception) -> Optional[int]: - """HTTP status code from the error or its cause chain.""" - return _from_cause_chain(error, _status_of, None) + """HTTP status code from the error or its cause chain; a body-carried numeric + ``code`` counts when the exception itself carries no status (#121270).""" + status = _from_cause_chain(error, _status_of, None) + if status is None: + status = _status_code_from_body(_from_cause_chain(error, _body_of, {})) + return status def _extract_error_body(error: Exception) -> dict: diff --git a/tests/agent/test_error_classifier_body_status.py b/tests/agent/test_error_classifier_body_status.py new file mode 100644 index 0000000000..26bd0daea7 --- /dev/null +++ b/tests/agent/test_error_classifier_body_status.py @@ -0,0 +1,84 @@ +"""Body-carried HTTP status extraction — an in-stream SSE error object's numeric +``code`` must classify like the equivalent HTTP response (#121270).""" + +from types import SimpleNamespace + +from agent.error_classifier import ( + FailoverReason, + classify_api_error, + _extract_status_code, +) + + +class MockAPIError(Exception): + """Simulates a status-less OpenAI SDK APIError raised mid-stream.""" + + def __init__(self, message, status_code=None, body=None, headers=None): + super().__init__(message) + self.status_code = status_code + self.body = body or {} + self.response = SimpleNamespace(headers=headers or {}) + + +_BAN_BODY = { + "error": { + "code": 403, + "message": "Your account has been banned by the upstream provider", + "metadata": {"provider_name": "acme"}, + } +} + + +class TestBodyCarriedStatusExtraction: + def test_in_stream_error_code_is_extracted_as_status(self): + assert ( + _extract_status_code(MockAPIError("Error code: 403", body=_BAN_BODY)) == 403 + ) + + def test_exception_status_still_wins_over_body_code(self): + err = MockAPIError("Too Many Requests", status_code=429, body=_BAN_BODY) + assert _extract_status_code(err) == 429 + + def test_symbolic_string_codes_are_not_statuses(self): + body = {"error": {"code": "insufficient_quota", "message": "quota exceeded"}} + assert _extract_status_code(MockAPIError("quota", body=body)) is None + + def test_out_of_range_and_non_int_codes_are_ignored(self): + for code in (99, 600, "403", True, 403.0): + body = {"error": {"code": code, "message": "x"}} + assert _extract_status_code(MockAPIError("x", body=body)) is None, code + + def test_top_level_code_and_status_keys_also_count(self): + assert _extract_status_code(MockAPIError("x", body={"code": 503})) == 503 + assert ( + _extract_status_code( + MockAPIError("x", body={"error": {"http_status": 502}}) + ) + == 502 + ) + + +class TestInStreamErrorClassification: + def test_403_ban_is_auth_not_transient_retry(self): + result = classify_api_error( + MockAPIError("Error code: 403", body=_BAN_BODY), provider="custom" + ) + assert result.status_code == 403 + assert result.reason == FailoverReason.auth + assert result.retryable is False + assert result.should_fallback is True + + def test_502_in_stream_error_stays_retryable(self): + body = {"error": {"code": 502, "message": "upstream connect error"}} + result = classify_api_error( + MockAPIError("Error code: 502", body=body), provider="custom" + ) + assert result.status_code == 502 + assert result.retryable is True + + def test_status_less_body_less_error_stays_unknown(self): + result = classify_api_error( + MockAPIError("weird failure", body={}), provider="custom" + ) + assert result.status_code is None + assert result.reason == FailoverReason.unknown