Merge pull request #795 from Routstr/fix/concise-upstream-errors

fix: shorten litellm errors and normalize buffered stream failures
This commit is contained in:
9qeklajc
2026-10-01 15:27:56 +02:00
committed by GitHub
3 changed files with 266 additions and 56 deletions
+33 -18
View File
@@ -3076,25 +3076,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",
+81 -38
View File
@@ -485,6 +485,76 @@ 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.removeprefix("litellm.")
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 = "<unreadable>"
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 +676,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 = "<unreadable>"
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 +697,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={
+152
View File
@@ -0,0 +1,152 @@
import os
from typing import Any, AsyncIterator
from unittest.mock import AsyncMock, patch
import litellm
import pytest
from litellm.exceptions import MidStreamFallbackError
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.core.exceptions import UpstreamError # noqa: E402
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
from routstr.upstream.messages_dispatch import ( # noqa: E402
collapse_litellm_message,
)
_MIDSTREAM_FAILURE = MidStreamFallbackError(
message="No credits.",
model="x",
llm_provider="openai",
original_exception=litellm.APIError(
status_code=500, message="No credits.", llm_provider="openai", model="x"
),
)
@pytest.mark.parametrize(
("message", "expected"),
[
("You have no credits remaining.", "You have no credits remaining."),
# upstream_error_from_exception reads `.message`, which omits the
# "Original exception:" chain that only `str()` appends.
(_MIDSTREAM_FAILURE.message, "No credits."),
(str(_MIDSTREAM_FAILURE), "No credits."),
("x" * 301, "x" * 299 + "…"),
],
)
def test_collapse_litellm_message(message: str, expected: str) -> None:
assert collapse_litellm_message(message) == expected
_RATE_LIMIT = litellm.RateLimitError(
message=(
"Rate limit reached for gpt-4o on tokens per min (TPM): Limit 30000, "
"Used 29000, Requested 2000. Please try again in 1.2s."
),
llm_provider="openai",
model="gpt-4o",
)
_BAD_REQUEST = litellm.BadRequestError(
message="context length exceeded", model="gpt-4o", llm_provider="openai"
)
_MID_STREAM_CASES = [
pytest.param(_RATE_LIMIT, 429, "UPSTREAM_RATE_LIMIT", id="rate-limit"),
pytest.param(_BAD_REQUEST, 400, None, id="bad-request"),
pytest.param(_MIDSTREAM_FAILURE, 500, None, id="midstream-fallback"),
]
def _make_model() -> Model:
return Model(
id="gpt-4o",
name="gpt-4o",
created=0,
description="",
context_length=8192,
architecture=Architecture(
modality="text",
input_modalities=["text"],
output_modalities=["text"],
tokenizer="x",
instruct_type=None,
),
pricing=Pricing(
prompt=0.0,
completion=0.0,
request=0.0,
image=0.0,
web_search=0.0,
internal_reasoning=0.0,
max_cost=0.0,
),
)
def _failing_stream(exc: Exception) -> AsyncIterator[dict]:
async def gen() -> AsyncIterator[dict]:
yield {
"type": "message_start",
"message": {"id": "msg_1", "model": "gpt-4o", "usage": {}},
}
raise exc
return gen()
def _assert_upstream_error(
err: UpstreamError, status_code: int, code: str | None
) -> None:
assert err.status_code == status_code
assert err.code == code
assert err.from_upstream_response is True
assert "litellm." not in str(err)
@pytest.mark.asyncio
@pytest.mark.parametrize(("exc", "status_code", "code"), _MID_STREAM_CASES)
async def test_non_streaming_aggregation_surfaces_mid_stream_failure(
exc: Exception, status_code: int, code: str | None
) -> None:
async def fake_acreate(**kwargs: Any) -> AsyncIterator[dict]:
return _failing_stream(exc)
with (
patch(
"litellm.anthropic.messages.acreate",
new=AsyncMock(side_effect=fake_acreate),
),
pytest.raises(UpstreamError) as exc_info,
):
await BaseUpstreamProvider(
base_url="http://test", api_key="k"
)._dispatch_anthropic_messages(
request_body=b'{"messages": [], "max_tokens": 8, "stream": false}',
model_obj=_make_model(),
)
_assert_upstream_error(exc_info.value, status_code, code)
@pytest.mark.asyncio
@pytest.mark.parametrize(("exc", "status_code", "code"), _MID_STREAM_CASES)
async def test_x_cashu_buffered_stream_surfaces_mid_stream_failure(
exc: Exception, status_code: int, code: str | None
) -> None:
provider = BaseUpstreamProvider(base_url="http://test", api_key="k")
with pytest.raises(UpstreamError) as exc_info:
await provider._stream_x_cashu_litellm_messages(
_failing_stream(exc),
amount=5_000,
unit="sat",
max_cost_for_model=10_000,
requested_model="gpt-4o",
mint=None,
request_id="req-test",
)
_assert_upstream_error(exc_info.value, status_code, code)