mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: shorten litellm errors and normalize buffered stream failures
This commit is contained in:
+33
-18
@@ -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",
|
||||
|
||||
@@ -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 = "<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 +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 = "<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 +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={
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user