Compare commits

..
Author SHA1 Message Date
9qeklajc 9e9c2bde57 fix: buffer-based SSE parser for all supported providers
Replace the fragile `re.split(b"data: ")` streaming parser with a
buffered, event-delimited parser in both the chat-completion and
Responses-API streamers. Events are accumulated until the SSE blank-line
delimiter, so parsing is independent of network chunk boundaries.

Fixes:
- OpenRouter `: OPENROUTER PROCESSING` keepalive comments no longer leak
  to clients as `data: : ...` (the `Unexpected token ':'` client crash).
- JSON payloads split across TCP reads are reassembled before parsing.
- CRLF framing (Gemini native alt=sse) handled.
- Combined content+usage chunks (Gemini thinking models over the
  OpenAI-compat endpoint) forward content once and still report usage in
  the cost trailer, instead of dropping the assistant message.
- Multi-line non-JSON `data` blocks are re-prefixed per line so they stay
  valid SSE framing for the client.

Adds tests/unit/test_streaming_sse_providers.py driving the real
generator against per-provider on-the-wire framing, and switches the
integration token-mint fallback to secrets.token_hex to kill a
PRNG/clock collision flake.
2026-06-07 12:58:54 +02:00
9qeklajcandGitHub 222fd6ed45 Merge pull request #542 from Routstr/fix-model-provider-field-response
make sure model provider field naming is correct
2026-06-07 11:32:01 +02:00
9qeklajc 4abd751f5f make sure model provider field naming is correct 2026-06-07 11:01:53 +02:00
7 changed files with 153 additions and 368 deletions
+38 -78
View File
@@ -28,7 +28,6 @@ from .payment.helpers import (
from .payment.models import Model
from .upstream import BaseUpstreamProvider
from .upstream.helpers import init_upstreams
from .upstream.request_correction import correct_request, extract_error_message
logger = get_logger(__name__)
proxy_router = APIRouter()
@@ -353,89 +352,50 @@ async def proxy(
if request_body_dict:
await pay_for_request(key, max_cost_for_model, session)
# Tracks request params already removed in response to upstream rejections,
# shared across providers so a stripped param stays stripped on failover and
# the reactive retry can never loop unboundedly.
already_stripped: set[str] = set()
for i, upstream in enumerate(upstreams):
headers = upstream.prepare_headers(dict(request.headers))
try:
while True:
try:
if is_responses_api:
response = await upstream.forward_responses_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
else:
response = await upstream.forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
except UpstreamError:
# Let the outer UpstreamError handler manage retry/revert
raise
except Exception as e:
# Unexpected error (not an upstream failure) — revert and propagate
logger.error(
"Unexpected error in upstream request, reverting payment",
extra={
"error": str(e),
"error_type": type(e).__name__,
"path": path,
"key_hash": key.hashed_key[:8] + "...",
"max_cost_for_model": max_cost_for_model,
},
)
await revert_pay_for_request(key, session, max_cost_for_model)
raise
# Reactive recovery: some models reject one specific request
# param (e.g. newer Anthropic models deprecating `temperature`).
# When the upstream 400s naming such a param, strip it from the
# body and retry the SAME upstream. ``already_stripped`` bounds
# this to one retry per distinct param so it always terminates.
if response.status_code == 400:
correction = correct_request(
try:
if is_responses_api:
response = await upstream.forward_responses_request(
request,
path,
headers,
request_body,
extract_error_message(response),
already_stripped,
key,
max_cost_for_model,
session,
model_obj,
)
if correction is not None:
request_body, bad_param = correction.body, correction.label
already_stripped.add(bad_param)
request_body_dict = parse_request_body_json(
request_body, path
)
logger.warning(
"Upstream %s rejected param '%s' for model=%s; "
"stripping and retrying same upstream",
upstream.provider_type,
bad_param,
model_id,
extra={
"provider": upstream.provider_type,
"model": model_id,
"stripped_param": bad_param,
"path": path,
},
)
continue
break
else:
response = await upstream.forward_request(
request,
path,
headers,
request_body,
key,
max_cost_for_model,
session,
model_obj,
)
except UpstreamError:
# Let the outer UpstreamError handler manage retry/revert
raise
except Exception as e:
# Unexpected error (not an upstream failure) — revert and propagate
logger.error(
"Unexpected error in upstream request, reverting payment",
extra={
"error": str(e),
"error_type": type(e).__name__,
"path": path,
"key_hash": key.hashed_key[:8] + "...",
"max_cost_for_model": max_cost_for_model,
},
)
await revert_pay_for_request(key, session, max_cost_for_model)
raise
if response.status_code != 200:
# Check if we should retry (502 Upstream Error or 429 Rate Limit)
+30 -7
View File
@@ -202,14 +202,26 @@ class BaseUpstreamProvider:
already reported its own provider (e.g. OpenRouter returns
``"provider": "Fireworks"``), otherwise just ``"<provider_type>"``
for direct upstreams.
Idempotent: re-stamping an already-stamped payload must not nest the
prefix repeatedly (e.g. never ``"anthropic:anthropic"``). This matters
because streaming paths can apply the field more than once per chunk.
"""
if not isinstance(response_json, dict):
return
provider_type = (self.provider_type or "").strip()
existing = response_json.get("provider")
if isinstance(existing, str) and existing.strip():
response_json["provider"] = f"{self.provider_type}:{existing.strip()}"
else:
response_json["provider"] = self.provider_type
existing_str = existing.strip() if isinstance(existing, str) else ""
if not existing_str:
response_json["provider"] = provider_type
return
# Already stamped by a previous pass — leave it untouched.
if existing_str == provider_type or existing_str.startswith(
f"{provider_type}:"
):
response_json["provider"] = existing_str
return
response_json["provider"] = f"{provider_type}:{existing_str}"
def inject_cost_metadata(
self,
@@ -822,8 +834,14 @@ class BaseUpstreamProvider:
yield prefix + b"data: " + json.dumps(obj).encode() + b"\n\n"
else:
# Non-JSON data payload (partial fragment already reassembled
# by buffering, or a provider control string) - forward as-is.
yield prefix + b"data: " + data + b"\n\n"
# by buffering, or a provider control string). Re-prefix each
# line so multi-line ``data`` stays valid SSE framing - a bare
# second line would otherwise reach the client without its
# ``data:`` field and break naive parsers.
body = b"".join(
b"data: " + ln + b"\n" for ln in data.split(b"\n")
)
yield prefix + body + b"\n"
try:
# Buffer bytes across network chunks and dispatch only on the SSE
@@ -1214,7 +1232,12 @@ class BaseUpstreamProvider:
yield prefix + b"data: " + json.dumps(obj).encode() + b"\n\n"
else:
yield prefix + b"data: " + data + b"\n\n"
# Re-prefix each line so multi-line ``data`` stays valid SSE
# framing for the client.
body = b"".join(
b"data: " + ln + b"\n" for ln in data.split(b"\n")
)
yield prefix + body + b"\n"
try:
# Buffer across network chunks; dispatch only on the SSE event
+27
View File
@@ -18,6 +18,33 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
supports_anthropic_messages = True
litellm_provider_prefix = "openrouter/"
def _apply_provider_field(self, response_json: object) -> None:
"""Stamp the ``provider`` field for OpenRouter responses.
OpenRouter is a router, not the real serving provider, so a bare
``"openrouter"`` value carries no useful information. Rules:
- Real upstream sub-provider (e.g. ``"GMICloud"``) -> ``"openrouter:GMICloud"``.
- Missing sub-provider, or one that merely echoes ``"openrouter"`` ->
``"unknown"``.
- Idempotent: re-stamping never produces ``"openrouter:openrouter:..."``;
the ``openrouter:`` prefix appears at most once.
"""
if not isinstance(response_json, dict):
return
provider_type = (self.provider_type or "").strip()
existing = response_json.get("provider")
sub = existing.strip() if isinstance(existing, str) else ""
# Strip any already-applied "openrouter:" prefixes (idempotency).
prefix = f"{provider_type}:"
while sub.lower().startswith(prefix.lower()):
sub = sub[len(prefix) :].strip()
# No real sub-provider, or it just echoes our own router name.
if not sub or sub.lower() == provider_type.lower():
response_json["provider"] = "unknown"
return
response_json["provider"] = f"{provider_type}:{sub}"
def __init__(self, api_key: str, provider_fee: float = 1.06):
"""Initialize OpenRouter provider with API key.
-142
View File
@@ -1,142 +0,0 @@
"""Reactive request-correction layer.
When an upstream rejects a request with a recoverable 4xx error, this layer
tries to *fix* the request body and let the caller retry the same upstream
instead of failing outright. It is provider-agnostic: correctors key off the
upstream's own error wording, so the same recovery works across every provider.
The layer is a small pipeline of :data:`Corrector` callables. Each corrector
inspects the parsed request body and the upstream error message and either
returns a corrected body (plus a short label identifying the fix) or declines
by returning ``None``. Adding a new reactive fix means writing one corrector
and adding it to :data:`DEFAULT_CORRECTORS` — no changes to the proxy loop.
All corrections are immutable: a corrector never mutates the body it is given,
it returns a new ``dict``. The proxy threads an ``applied`` set of fix labels
through retries so each distinct fix is applied at most once, guaranteeing the
retry loop always terminates.
"""
from __future__ import annotations
import json
import re
from collections.abc import Callable, Sequence
from dataclasses import dataclass
from fastapi.responses import Response
from ..core import get_logger
logger = get_logger(__name__)
# Matches upstream error text that names a single rejected request parameter,
# e.g. "`temperature` is deprecated for this model." or
# "parameter 'top_p' is not supported". Keys off the upstream's own wording so
# a 400 about an unsupported sampling/option field can be recovered by stripping
# that field and retrying the same upstream.
_UNSUPPORTED_PARAM_RE = re.compile(
r"[`'\"]?(?P<param>[a-zA-Z_][a-zA-Z0-9_]*)[`'\"]?\s+is\s+"
r"(?:deprecated|not\s+supported|unsupported|no\s+longer\s+supported)",
re.IGNORECASE,
)
# A corrector inspects the parsed request body and the upstream error message
# and returns ``(new_body_dict, label)`` for a fix it can apply, or ``None`` to
# decline. ``label`` identifies the fix so it is applied at most once per request.
Corrector = Callable[[dict, str], "tuple[dict, str] | None"]
@dataclass(frozen=True)
class Correction:
"""A successful request correction ready to retry.
``body`` is the corrected JSON body (encoded), ``label`` identifies the fix
that was applied (e.g. the stripped param name) so the caller can guard
against applying the same fix twice.
"""
body: bytes
label: str
def extract_error_message(response: Response) -> str:
"""Best-effort extraction of an error message string from a proxy Response."""
body_bytes = getattr(response, "body", None)
if not body_bytes:
return ""
try:
data = json.loads(body_bytes)
except Exception:
return body_bytes.decode("utf-8", errors="ignore")[:500]
if isinstance(data, dict):
err = data.get("error")
if isinstance(err, dict):
msg = err.get("message") or err.get("detail")
if isinstance(msg, str):
return msg
elif isinstance(err, str):
return err
if isinstance(data.get("message"), str):
return data["message"]
return ""
def strip_unsupported_param(
body: dict, error_message: str
) -> tuple[dict, str] | None:
"""Drop a top-level param the upstream named as unsupported/deprecated.
Returns ``(new_body, param)`` (a new dict, original untouched) when the
error names a top-level param present in the body, otherwise ``None``.
"""
match = _UNSUPPORTED_PARAM_RE.search(error_message)
if not match:
return None
param = match.group("param")
if param not in body:
return None
new_body = {k: v for k, v in body.items() if k != param}
return new_body, param
# Ordered pipeline of correctors tried on each recoverable rejection.
DEFAULT_CORRECTORS: tuple[Corrector, ...] = (strip_unsupported_param,)
def correct_request(
request_body: bytes,
error_message: str,
applied: set[str],
correctors: Sequence[Corrector] = DEFAULT_CORRECTORS,
) -> Correction | None:
"""Try to correct a rejected request body so it can be retried.
Runs each corrector in order against the parsed body and ``error_message``.
The first corrector that proposes a fix whose ``label`` is not already in
``applied`` wins; its result is returned as a :class:`Correction`. Returns
``None`` when nothing parses, nothing matches, or every proposed fix was
already applied — the caller then treats the response as a normal failure.
``applied`` is read-only here; the caller records the returned ``label`` to
bound retries and guarantee forward progress.
"""
if not request_body or not error_message:
return None
try:
data = json.loads(request_body)
except Exception:
return None
if not isinstance(data, dict):
return None
for corrector in correctors:
result = corrector(data, error_message)
if result is None:
continue
new_body, label = result
if label in applied:
continue
return Correction(body=json.dumps(new_body).encode(), label=label)
return None
+35 -11
View File
@@ -32,11 +32,39 @@ def test_apply_provider_field_openrouter_passthrough() -> None:
def test_apply_provider_field_openrouter_no_upstream_provider() -> None:
"""If OpenRouter omits the provider field, fall back to provider_type."""
"""If OpenRouter omits the provider field, the real serving provider is
unknown — a bare ``openrouter`` value carries no information."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"id": "gen-abc"}
p._apply_provider_field(data)
assert data["provider"] == "openrouter"
assert data["provider"] == "unknown"
def test_apply_provider_field_openrouter_echoes_router_name() -> None:
"""If OpenRouter reports its own name as the provider, treat as unknown."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": "openrouter"}
p._apply_provider_field(data)
assert data["provider"] == "unknown"
def test_apply_provider_field_openrouter_idempotent_no_double_prefix() -> None:
"""Re-stamping must never nest the prefix: openrouter only once."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": "GMICloud"}
p._apply_provider_field(data)
assert data["provider"] == "openrouter:GMICloud"
# Second pass (e.g. streaming) keeps a single prefix.
p._apply_provider_field(data)
assert data["provider"] == "openrouter:GMICloud"
def test_apply_provider_field_openrouter_collapses_existing_double_prefix() -> None:
"""A pre-existing double prefix is collapsed to a single one."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": "openrouter:openrouter:GMICloud"}
p._apply_provider_field(data)
assert data["provider"] == "openrouter:GMICloud"
def test_apply_provider_field_strips_whitespace() -> None:
@@ -50,28 +78,24 @@ def test_apply_provider_field_blank_upstream_treated_as_missing() -> None:
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": " "}
p._apply_provider_field(data)
assert data["provider"] == "openrouter"
assert data["provider"] == "unknown"
def test_apply_provider_field_non_string_upstream_treated_as_missing() -> None:
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": 42}
p._apply_provider_field(data)
assert data["provider"] == "openrouter"
assert data["provider"] == "unknown"
def test_apply_provider_field_idempotent_for_direct_upstream() -> None:
"""Calling twice on a direct upstream payload should keep the same
value, not nest the prefix repeatedly."""
"""Calling twice on a direct upstream payload keeps the same value and
never nests the prefix (no ``anthropic:anthropic``)."""
p = _make_provider(AnthropicUpstreamProvider, "anthropic")
data: dict = {}
p._apply_provider_field(data)
p._apply_provider_field(data)
assert data["provider"] == "anthropic:anthropic"
# Document current (deliberate) behavior: second pass treats the
# first-pass value as an upstream-reported provider. Callers should
# only invoke this once per chunk — guarded via the
# ``"provider" not in data`` checks in streaming paths.
assert data["provider"] == "anthropic"
def test_apply_provider_field_ignores_non_dict() -> None:
-130
View File
@@ -1,130 +0,0 @@
"""Unit tests for the reactive request-correction layer.
Covers the recovery path that lets a request survive a 400 where the upstream
names a single unsupported request param (e.g. newer Anthropic models
deprecating ``temperature``): the param is stripped from the JSON body and the
same upstream is retried, provider-agnostically, keyed off the error text.
"""
from __future__ import annotations
import json
from fastapi.responses import Response
from routstr.upstream.request_correction import (
Correction,
correct_request,
extract_error_message,
strip_unsupported_param,
)
def _body(**kwargs: object) -> bytes:
return json.dumps(kwargs).encode()
class TestCorrectRequest:
def test_strips_deprecated_temperature(self) -> None:
body = _body(model="claude-opus-4-8", temperature=1, messages=[])
result = correct_request(
body, "`temperature` is deprecated for this model.", set()
)
assert isinstance(result, Correction)
assert result.label == "temperature"
decoded = json.loads(result.body)
assert "temperature" not in decoded
assert decoded["model"] == "claude-opus-4-8"
def test_strips_not_supported_param(self) -> None:
body = _body(model="m", top_p=0.9, messages=[])
result = correct_request(body, "Parameter 'top_p' is not supported", set())
assert result is not None
assert result.label == "top_p"
assert "top_p" not in json.loads(result.body)
def test_returns_none_when_label_already_applied(self) -> None:
body = _body(model="m", temperature=1)
assert (
correct_request(body, "`temperature` is deprecated", {"temperature"})
is None
)
def test_returns_none_when_param_absent_from_body(self) -> None:
body = _body(model="m", messages=[])
assert correct_request(body, "`temperature` is deprecated", set()) is None
def test_returns_none_when_message_does_not_match(self) -> None:
body = _body(model="m", temperature=1)
assert correct_request(body, "Insufficient balance", set()) is None
def test_returns_none_on_empty_inputs(self) -> None:
assert correct_request(b"", "`temperature` is deprecated", set()) is None
assert correct_request(_body(temperature=1), "", set()) is None
def test_returns_none_on_non_object_body(self) -> None:
assert correct_request(b"[1, 2, 3]", "`temperature` is deprecated", set()) is None
class TestStripUnsupportedParam:
def test_does_not_mutate_input(self) -> None:
body = {"model": "m", "temperature": 1}
result = strip_unsupported_param(body, "`temperature` is deprecated")
assert result is not None
new_body, param = result
assert param == "temperature"
assert "temperature" not in new_body
# original untouched (immutability)
assert body == {"model": "m", "temperature": 1}
def test_declines_when_no_match(self) -> None:
assert strip_unsupported_param({"temperature": 1}, "nope") is None
class TestExtractErrorMessage:
def test_extracts_nested_error_message(self) -> None:
resp = Response(
content=json.dumps(
{"error": {"message": "`temperature` is deprecated", "type": "x"}}
).encode(),
status_code=400,
)
assert extract_error_message(resp) == "`temperature` is deprecated"
def test_extracts_string_error(self) -> None:
resp = Response(
content=json.dumps({"error": "bad request"}).encode(), status_code=400
)
assert extract_error_message(resp) == "bad request"
def test_extracts_top_level_message(self) -> None:
resp = Response(
content=json.dumps({"message": "nope"}).encode(), status_code=400
)
assert extract_error_message(resp) == "nope"
def test_empty_body_returns_empty_string(self) -> None:
assert extract_error_message(Response(status_code=400)) == ""
def test_non_json_body_returns_preview(self) -> None:
resp = Response(content=b"plain text error", status_code=400)
assert extract_error_message(resp) == "plain text error"
class TestEndToEndChaining:
def test_two_distinct_params_corrected_sequentially(self) -> None:
"""Simulates the proxy loop: each 400 fixes one param, set guards reuse."""
body = _body(model="m", temperature=1, top_p=0.5, messages=[])
applied: set[str] = set()
first = correct_request(body, "`temperature` is deprecated", applied)
assert first is not None
body, applied = first.body, applied | {first.label}
second = correct_request(body, "`top_p` is not supported", applied)
assert second is not None
body, applied = second.body, applied | {second.label}
decoded = json.loads(body)
assert "temperature" not in decoded and "top_p" not in decoded
assert applied == {"temperature", "top_p"}
@@ -314,3 +314,26 @@ async def test_requested_model_override_applied() -> None:
content_chunks = [o for o in objs if o.get("choices")]
assert content_chunks, "expected at least one forwarded content chunk"
assert all(o.get("model") == "routstr-model" for o in content_chunks)
@pytest.mark.asyncio
async def test_multiline_non_json_data_each_line_prefixed() -> None:
"""A multi-line non-JSON ``data`` block must keep a ``data:`` prefix per line.
Two ``data:`` lines in one event reassemble to ``line one\\nline two``, which
is not JSON, so it takes the raw-forward path. The parser must re-prefix each
line; a bare second line would reach the client without its ``data:`` field
and break naive SSE parsers.
"""
chunks = [
b"data: line one\ndata: line two\n\n",
b"data: [DONE]\n\n",
]
out = await _drive(chunks)
blob = b"".join(out)
for line in blob.split(b"\n"):
stripped = line.strip()
if not stripped or stripped == b"[DONE]":
continue
assert line.startswith(b"data: "), f"bare line leaked to client: {line!r}"
assert b"data: line one" in blob and b"data: line two" in blob