From 9e62ba302ec50b907a601b539c3f64b653899da2 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Thu, 1 Oct 2026 01:14:06 +0200 Subject: [PATCH] fix: shorten litellm errors and normalize buffered stream failures --- routstr/upstream/base.py | 51 +++++--- routstr/upstream/messages_dispatch.py | 122 ++++++++++++++------ tests/unit/test_messages_upstream_errors.py | 20 ++++ 3 files changed, 137 insertions(+), 56 deletions(-) create mode 100644 tests/unit/test_messages_upstream_errors.py diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 2f65ebd3..69bc1769 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -3024,25 +3024,40 @@ class BaseUpstreamProvider: input_cost = 0.0 output_cost = 0.0 - async for annotated in messages_dispatch.stream_annotated_events( - iterator, requested_model - ): - if annotated.model: - last_model_seen = annotated.model - # See _stream_litellm_messages for why this is max() not +=. - input_tokens = max(input_tokens, annotated.input_tokens) - output_tokens = max(output_tokens, annotated.output_tokens) - cache_read_input_tokens = max( - cache_read_input_tokens, annotated.cache_read_input_tokens + try: + annotated_events = messages_dispatch.stream_annotated_events( + iterator, requested_model ) - cache_creation_input_tokens = max( - cache_creation_input_tokens, - annotated.cache_creation_input_tokens, - ) - total_cost = max(total_cost, annotated.total_cost) - input_cost = max(input_cost, annotated.input_cost) - output_cost = max(output_cost, annotated.output_cost) - buffered.append(annotated) + async for annotated in annotated_events: + if annotated.model: + last_model_seen = annotated.model + # See _stream_litellm_messages for why this is max() not +=. + input_tokens = max(input_tokens, annotated.input_tokens) + output_tokens = max(output_tokens, annotated.output_tokens) + cache_read_input_tokens = max( + cache_read_input_tokens, annotated.cache_read_input_tokens + ) + cache_creation_input_tokens = max( + cache_creation_input_tokens, + annotated.cache_creation_input_tokens, + ) + total_cost = max(total_cost, annotated.total_cost) + input_cost = max(input_cost, annotated.input_cost) + output_cost = max(output_cost, annotated.output_cost) + buffered.append(annotated) + except Exception as exc: + # Buffering lets us return an HTTP error before sending headers. + if messages_dispatch.is_provider_exception(exc): + raise messages_dispatch.upstream_error_from_exception( + exc, + log_message="Upstream stream failed mid-flight", + log_extra={ + "model": last_model_seen or requested_model or "unknown", + "provider": self.provider_type or self.base_url, + "request_id": request_id, + }, + ) from exc + raise response_headers: dict[str, str] = { "Cache-Control": "no-cache", diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index d9df0277..85bfd312 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -485,6 +485,79 @@ def compute_refund(amount: int, unit: str, cost_msats: int) -> int: raise ValueError(f"Invalid unit: {unit}") +_MAX_UPSTREAM_MESSAGE_CHARS = 300 + + +def collapse_litellm_message(message: str) -> str: + """Keep the innermost provider message and cap its length.""" + tail = message.rsplit("Original exception:", 1)[-1].strip() + while True: + stripped = tail + for prefix in ("litellm.",): + if stripped.startswith(prefix): + stripped = stripped[len(prefix) :] + head, _, rest = stripped.partition(": ") + if rest and head.endswith(("Error", "Exception")): + stripped = rest.strip() + if stripped == tail: + break + tail = stripped + if len(tail) > _MAX_UPSTREAM_MESSAGE_CHARS: + tail = tail[: _MAX_UPSTREAM_MESSAGE_CHARS - 1].rstrip() + "…" + return tail + + +def is_provider_exception(exc: BaseException) -> bool: + """Distinguish SDK failures from bugs in our stream handling.""" + return type(exc).__module__.split(".", 1)[0] in {"litellm", "openai"} + + +def upstream_error_from_exception( + exc: Exception, + *, + log_message: str, + log_extra: dict[str, Any] | None = None, +) -> UpstreamError: + """Redact and classify provider failures, including mid-stream errors.""" + raw_message = getattr(exc, "message", None) or str(exc) or repr(exc) + # Redact provider account ids before the message reaches logs or the client. + exc_message = collapse_litellm_message(redact_org_ids(raw_message)) + exc_status = getattr(exc, "status_code", None) + exc_response = getattr(exc, "response", None) + response_text = None + if exc_response is not None: + try: + response_text = redact_org_ids( + getattr(exc_response, "text", str(exc_response)) + ) + except Exception: + response_text = "" + status_for_classify = exc_status if isinstance(exc_status, int) else 502 + rate_limit = classify_rate_limit( + status_for_classify, exc_message, getattr(exc, "headers", None) + ) + logger.error( + log_message, + extra={ + "error": exc_message, + "error_type": type(exc).__name__, + "status_code": exc_status, + "error_code": rate_limit.code if rate_limit else None, + "llm_provider": getattr(exc, "llm_provider", None), + "body": redact_org_ids(str(getattr(exc, "body", "") or "")) or None, + "response_text": response_text, + **(log_extra or {}), + }, + ) + return UpstreamError( + f"Upstream error via litellm: {exc_message}", + status_code=status_for_classify, + code=rate_limit.code if rate_limit else None, + details=rate_limit.as_details() if rate_limit else None, + from_upstream_response=True, + ) + + async def dispatch_anthropic_messages( *, request_body: bytes | None, @@ -606,44 +679,10 @@ async def dispatch_anthropic_messages( try: result = await litellm.anthropic.messages.acreate(**kwargs) except Exception as exc: - raw_message = getattr(exc, "message", None) or str(exc) or repr(exc) - # Redact provider account identifiers before the message reaches logs - # or the surfaced error. - exc_message = redact_org_ids(raw_message) - exc_status = getattr(exc, "status_code", None) - exc_response = getattr(exc, "response", None) - response_text = None - if exc_response is not None: - try: - response_text = redact_org_ids( - getattr(exc_response, "text", str(exc_response)) - ) - except Exception: - response_text = "" - status_for_classify = exc_status if isinstance(exc_status, int) else 502 - rate_limit = classify_rate_limit( - status_for_classify, exc_message, getattr(exc, "headers", None) - ) - logger.error( - "litellm dispatch failed", - extra={ - "error": exc_message, - "error_type": type(exc).__name__, - "status_code": exc_status, - "error_code": rate_limit.code if rate_limit else None, - "llm_provider": getattr(exc, "llm_provider", None), - "body": redact_org_ids(str(getattr(exc, "body", "") or "")) or None, - "response_text": response_text, - "model": litellm_model, - "api_base": base_url, - }, - ) - raise UpstreamError( - f"Upstream error via litellm: {exc_message}", - status_code=status_for_classify, - code=rate_limit.code if rate_limit else None, - details=rate_limit.as_details() if rate_limit else None, - from_upstream_response=True, + raise upstream_error_from_exception( + exc, + log_message="litellm dispatch failed", + log_extra={"model": litellm_model, "api_base": base_url}, ) from exc if transform_stream is not None and hasattr(result, "__aiter__"): @@ -661,6 +700,13 @@ async def dispatch_anthropic_messages( cast(AsyncIterator[Any], result) ) except Exception as exc: + if is_provider_exception(exc): + # Upstream failed part-way through, not an aggregation bug. + raise upstream_error_from_exception( + exc, + log_message="Upstream stream failed mid-flight", + log_extra={"model": litellm_model, "api_base": base_url}, + ) from exc logger.error( "Failed to aggregate streamed events into message", extra={ diff --git a/tests/unit/test_messages_upstream_errors.py b/tests/unit/test_messages_upstream_errors.py new file mode 100644 index 00000000..61b6febd --- /dev/null +++ b/tests/unit/test_messages_upstream_errors.py @@ -0,0 +1,20 @@ +import pytest + +from routstr.upstream.messages_dispatch import collapse_litellm_message + + +@pytest.mark.parametrize( + ("message", "expected"), + [ + ("You have no credits remaining.", "You have no credits remaining."), + ( + "litellm.MidStreamFallbackError: litellm.APIError: No credits. " + "Original exception: MidStreamFallbackError: No credits. " + "Original exception: APIError: litellm.APIError: No credits.", + "No credits.", + ), + ("x" * 301, "x" * 299 + "…"), + ], +) +def test_collapse_litellm_message(message: str, expected: str) -> None: + assert collapse_litellm_message(message) == expected