mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: stop reporting OpenRouter provider as unknown on stream and envelope payloads
The OpenRouter stamper wrote "unknown" whenever a payload lacked a top-level provider. That hit every Anthropic /messages event, every Responses event, and the usage/cost payloads routstr synthesizes at the end of a stream. - Read the provider from the Anthropic `message` and Responses `response` envelopes as well as the top level. - Carry the provider reported earlier in a stream to later events and to the synthesized usage/cost payloads.
This commit is contained in:
@@ -214,6 +214,20 @@ def _responses_usage_payload(data_json: dict) -> dict:
|
|||||||
return nested if isinstance(nested, dict) else data_json
|
return nested if isinstance(nested, dict) else data_json
|
||||||
|
|
||||||
|
|
||||||
|
def _reported_provider(payload: dict) -> str | None:
|
||||||
|
"""Provider named by an upstream payload, if any.
|
||||||
|
|
||||||
|
Checked at top level first, then inside the Anthropic ``message`` and
|
||||||
|
Responses ``response`` envelopes, which is where those dialects nest it.
|
||||||
|
"""
|
||||||
|
for obj in (payload, payload.get("message"), payload.get("response")):
|
||||||
|
if isinstance(obj, dict):
|
||||||
|
value = obj.get("provider")
|
||||||
|
if isinstance(value, str) and value.strip():
|
||||||
|
return value.strip()
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _render_sse_event(field_lines: list[str], data: str) -> str:
|
def _render_sse_event(field_lines: list[str], data: str) -> str:
|
||||||
"""Re-frame one parsed event, re-prefixing every line of a multi-line data."""
|
"""Re-frame one parsed event, re-prefixing every line of a multi-line data."""
|
||||||
body = "".join(f"{line}\n" for line in field_lines)
|
body = "".join(f"{line}\n" for line in field_lines)
|
||||||
@@ -485,8 +499,7 @@ class BaseUpstreamProvider:
|
|||||||
return
|
return
|
||||||
response_json["provider_url"] = public_provider_url(self.base_url)
|
response_json["provider_url"] = public_provider_url(self.base_url)
|
||||||
provider_type = (self.provider_type or "").strip()
|
provider_type = (self.provider_type or "").strip()
|
||||||
existing = response_json.get("provider")
|
existing_str = _reported_provider(response_json) or ""
|
||||||
existing_str = existing.strip() if isinstance(existing, str) else ""
|
|
||||||
if not existing_str:
|
if not existing_str:
|
||||||
response_json["provider"] = provider_type
|
response_json["provider"] = provider_type
|
||||||
return
|
return
|
||||||
@@ -498,6 +511,17 @@ class BaseUpstreamProvider:
|
|||||||
return
|
return
|
||||||
response_json["provider"] = f"{provider_type}:{existing_str}"
|
response_json["provider"] = f"{provider_type}:{existing_str}"
|
||||||
|
|
||||||
|
def _stamp_streamed_provider(
|
||||||
|
self, payload: dict, carried: str | None
|
||||||
|
) -> str | None:
|
||||||
|
"""Stamp a streamed payload, falling back to a provider an earlier event
|
||||||
|
reported. Returns the provider to carry forward to later payloads."""
|
||||||
|
reported = _reported_provider(payload)
|
||||||
|
if reported is None and carried is not None:
|
||||||
|
payload["provider"] = carried
|
||||||
|
self._apply_provider_field(payload)
|
||||||
|
return reported or carried
|
||||||
|
|
||||||
def _log_full_refund(
|
def _log_full_refund(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -1169,6 +1193,7 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
usage_finalized = False
|
usage_finalized = False
|
||||||
last_model_seen: str | None = None
|
last_model_seen: str | None = None
|
||||||
|
provider_seen: str | None = None
|
||||||
|
|
||||||
async def finalize_db_only() -> None:
|
async def finalize_db_only() -> None:
|
||||||
nonlocal usage_finalized
|
nonlocal usage_finalized
|
||||||
@@ -1243,6 +1268,7 @@ class BaseUpstreamProvider:
|
|||||||
end of stream.
|
end of stream.
|
||||||
"""
|
"""
|
||||||
nonlocal last_model_seen, usage_chunk_data, done_seen, stream_id
|
nonlocal last_model_seen, usage_chunk_data, done_seen, stream_id
|
||||||
|
nonlocal provider_seen
|
||||||
|
|
||||||
event = raw_event.strip(b"\r\n")
|
event = raw_event.strip(b"\r\n")
|
||||||
if not event:
|
if not event:
|
||||||
@@ -1282,7 +1308,7 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
if isinstance(obj, dict):
|
if isinstance(obj, dict):
|
||||||
usage_estimator.observe(obj)
|
usage_estimator.observe(obj)
|
||||||
self._apply_provider_field(obj)
|
provider_seen = self._stamp_streamed_provider(obj, provider_seen)
|
||||||
if obj.get("model"):
|
if obj.get("model"):
|
||||||
last_model_seen = str(obj.get("model"))
|
last_model_seen = str(obj.get("model"))
|
||||||
if requested_model:
|
if requested_model:
|
||||||
@@ -1408,6 +1434,7 @@ class BaseUpstreamProvider:
|
|||||||
if legacy_completion
|
if legacy_completion
|
||||||
else "chat.completion.chunk",
|
else "chat.completion.chunk",
|
||||||
"model": last_model_seen or "unknown",
|
"model": last_model_seen or "unknown",
|
||||||
|
"provider": provider_seen,
|
||||||
"choices": [],
|
"choices": [],
|
||||||
"usage": {
|
"usage": {
|
||||||
"prompt_tokens": cost_data.get("input_tokens", 0),
|
"prompt_tokens": cost_data.get("input_tokens", 0),
|
||||||
@@ -1652,6 +1679,7 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
usage_finalized = False
|
usage_finalized = False
|
||||||
last_model_seen: str | None = None
|
last_model_seen: str | None = None
|
||||||
|
provider_seen: str | None = None
|
||||||
|
|
||||||
async def finalize_db_only() -> None:
|
async def finalize_db_only() -> None:
|
||||||
nonlocal usage_finalized
|
nonlocal usage_finalized
|
||||||
@@ -1715,7 +1743,7 @@ class BaseUpstreamProvider:
|
|||||||
and preserves ``event:``/``id:`` fields attached to their data
|
and preserves ``event:``/``id:`` fields attached to their data
|
||||||
line so Responses API event framing stays intact.
|
line so Responses API event framing stays intact.
|
||||||
"""
|
"""
|
||||||
nonlocal last_model_seen, usage_chunk_data, done_seen
|
nonlocal last_model_seen, usage_chunk_data, done_seen, provider_seen
|
||||||
nonlocal reasoning_tokens
|
nonlocal reasoning_tokens
|
||||||
|
|
||||||
event = raw_event.strip(b"\r\n")
|
event = raw_event.strip(b"\r\n")
|
||||||
@@ -1751,7 +1779,7 @@ class BaseUpstreamProvider:
|
|||||||
obj = json_codec.loads(data)
|
obj = json_codec.loads(data)
|
||||||
|
|
||||||
if isinstance(obj, dict):
|
if isinstance(obj, dict):
|
||||||
self._apply_provider_field(obj)
|
provider_seen = self._stamp_streamed_provider(obj, provider_seen)
|
||||||
if obj.get("model"):
|
if obj.get("model"):
|
||||||
last_model_seen = str(obj.get("model"))
|
last_model_seen = str(obj.get("model"))
|
||||||
if requested_model:
|
if requested_model:
|
||||||
@@ -1840,6 +1868,7 @@ class BaseUpstreamProvider:
|
|||||||
if usage_chunk_data is None:
|
if usage_chunk_data is None:
|
||||||
usage_chunk_data = {
|
usage_chunk_data = {
|
||||||
"type": "response.completed",
|
"type": "response.completed",
|
||||||
|
"provider": provider_seen,
|
||||||
"response": {
|
"response": {
|
||||||
"model": last_model_seen or "unknown",
|
"model": last_model_seen or "unknown",
|
||||||
"usage": {
|
"usage": {
|
||||||
@@ -2195,6 +2224,7 @@ class BaseUpstreamProvider:
|
|||||||
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
usage_estimator = MissingUsageEstimator(request_body, model_obj)
|
||||||
usage_finalized = False
|
usage_finalized = False
|
||||||
last_model_seen: str | None = None
|
last_model_seen: str | None = None
|
||||||
|
provider_seen: str | None = None
|
||||||
|
|
||||||
async def finalize_without_usage() -> bytes | None:
|
async def finalize_without_usage() -> bytes | None:
|
||||||
nonlocal usage_finalized
|
nonlocal usage_finalized
|
||||||
@@ -2244,7 +2274,7 @@ class BaseUpstreamProvider:
|
|||||||
async def stream_with_cost(
|
async def stream_with_cost(
|
||||||
max_cost_for_model: int,
|
max_cost_for_model: int,
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes, None]:
|
||||||
nonlocal usage_finalized, last_model_seen
|
nonlocal usage_finalized, last_model_seen, provider_seen
|
||||||
stored_chunks: list[bytes] = []
|
stored_chunks: list[bytes] = []
|
||||||
input_tokens: int = 0
|
input_tokens: int = 0
|
||||||
output_tokens: int = 0
|
output_tokens: int = 0
|
||||||
@@ -2301,7 +2331,9 @@ class BaseUpstreamProvider:
|
|||||||
last_model_seen = str(msg.get("model"))
|
last_model_seen = str(msg.get("model"))
|
||||||
|
|
||||||
provider_added = "provider" not in data
|
provider_added = "provider" not in data
|
||||||
self._apply_provider_field(data)
|
provider_seen = self._stamp_streamed_provider(
|
||||||
|
data, provider_seen
|
||||||
|
)
|
||||||
|
|
||||||
if requested_model:
|
if requested_model:
|
||||||
# Apply requested_model override
|
# Apply requested_model override
|
||||||
@@ -2419,6 +2451,7 @@ class BaseUpstreamProvider:
|
|||||||
try:
|
try:
|
||||||
combined_data = {
|
combined_data = {
|
||||||
"model": last_model_seen or "unknown",
|
"model": last_model_seen or "unknown",
|
||||||
|
"provider": provider_seen,
|
||||||
"usage": usage_data,
|
"usage": usage_data,
|
||||||
}
|
}
|
||||||
cost_data = await adjust_payment_for_tokens(
|
cost_data = await adjust_payment_for_tokens(
|
||||||
@@ -4197,6 +4230,7 @@ class BaseUpstreamProvider:
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
provider_seen: str | None = None
|
||||||
for i, line in enumerate(lines):
|
for i, line in enumerate(lines):
|
||||||
if line.startswith("data: "):
|
if line.startswith("data: "):
|
||||||
try:
|
try:
|
||||||
@@ -4204,7 +4238,9 @@ class BaseUpstreamProvider:
|
|||||||
if not isinstance(data_json, dict):
|
if not isinstance(data_json, dict):
|
||||||
continue
|
continue
|
||||||
provider_before = data_json.get("provider")
|
provider_before = data_json.get("provider")
|
||||||
self._apply_provider_field(data_json)
|
provider_seen = self._stamp_streamed_provider(
|
||||||
|
data_json, provider_seen
|
||||||
|
)
|
||||||
changed = data_json.get("provider") != provider_before
|
changed = data_json.get("provider") != provider_before
|
||||||
if cost_data and "usage" in data_json and data_json["usage"]:
|
if cost_data and "usage" in data_json and data_json["usage"]:
|
||||||
_inject_cost_into_usage(data_json, cost_data)
|
_inject_cost_into_usage(data_json, cost_data)
|
||||||
@@ -5265,6 +5301,7 @@ class BaseUpstreamProvider:
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
provider_seen: str | None = None
|
||||||
for i, (fields, data) in enumerate(events):
|
for i, (fields, data) in enumerate(events):
|
||||||
if data.strip() == "[DONE]":
|
if data.strip() == "[DONE]":
|
||||||
continue
|
continue
|
||||||
@@ -5275,7 +5312,7 @@ class BaseUpstreamProvider:
|
|||||||
if not isinstance(data_json, dict):
|
if not isinstance(data_json, dict):
|
||||||
continue
|
continue
|
||||||
provider_before = data_json.get("provider")
|
provider_before = data_json.get("provider")
|
||||||
self._apply_provider_field(data_json)
|
provider_seen = self._stamp_streamed_provider(data_json, provider_seen)
|
||||||
changed = data_json.get("provider") != provider_before
|
changed = data_json.get("provider") != provider_before
|
||||||
payload = _responses_usage_payload(data_json)
|
payload = _responses_usage_payload(data_json)
|
||||||
if cost_data and isinstance(payload.get("usage"), dict):
|
if cost_data and isinstance(payload.get("usage"), dict):
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from urllib.parse import urlparse
|
|||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from .base import BaseUpstreamProvider
|
from .base import BaseUpstreamProvider, _reported_provider
|
||||||
from .model_paths import public_provider_url
|
from .model_paths import public_provider_url
|
||||||
from .pricing_resolver import (
|
from .pricing_resolver import (
|
||||||
FallbackPricingResolver,
|
FallbackPricingResolver,
|
||||||
@@ -60,8 +60,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider):
|
|||||||
"""
|
"""
|
||||||
if not isinstance(response_json, dict):
|
if not isinstance(response_json, dict):
|
||||||
return
|
return
|
||||||
existing = response_json.get("provider")
|
if _reported_provider(response_json) is None:
|
||||||
if not (isinstance(existing, str) and existing.strip()):
|
|
||||||
response_json["provider"] = (
|
response_json["provider"] = (
|
||||||
urlparse(public_provider_url(self.base_url)).hostname
|
urlparse(public_provider_url(self.base_url)).hostname
|
||||||
or self.upstream_name
|
or self.upstream_name
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from typing import TYPE_CHECKING
|
|||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from ..payment.models import Model, async_fetch_openrouter_models
|
from ..payment.models import Model, async_fetch_openrouter_models
|
||||||
from .base import BaseUpstreamProvider
|
from .base import BaseUpstreamProvider, _reported_provider
|
||||||
from .model_paths import public_provider_url
|
from .model_paths import public_provider_url
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -35,8 +35,7 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
|
|||||||
return
|
return
|
||||||
response_json["provider_url"] = public_provider_url(self.base_url)
|
response_json["provider_url"] = public_provider_url(self.base_url)
|
||||||
provider_type = (self.provider_type or "").strip()
|
provider_type = (self.provider_type or "").strip()
|
||||||
existing = response_json.get("provider")
|
sub = _reported_provider(response_json) or ""
|
||||||
sub = existing.strip() if isinstance(existing, str) else ""
|
|
||||||
# Strip any already-applied "openrouter:" prefixes (idempotency).
|
# Strip any already-applied "openrouter:" prefixes (idempotency).
|
||||||
prefix = f"{provider_type}:"
|
prefix = f"{provider_type}:"
|
||||||
while sub.lower().startswith(prefix.lower()):
|
while sub.lower().startswith(prefix.lower()):
|
||||||
|
|||||||
@@ -89,6 +89,33 @@ def test_apply_provider_field_non_string_upstream_treated_as_missing() -> None:
|
|||||||
assert data["provider"] == "unknown"
|
assert data["provider"] == "unknown"
|
||||||
|
|
||||||
|
|
||||||
|
def test_apply_provider_field_openrouter_reads_nested_envelopes() -> None:
|
||||||
|
"""Anthropic ``message`` and Responses ``response`` envelopes nest the
|
||||||
|
upstream provider; it must not be reported as unknown."""
|
||||||
|
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
|
||||||
|
message_start: dict = {
|
||||||
|
"type": "message_start",
|
||||||
|
"message": {"provider": "Anthropic"},
|
||||||
|
}
|
||||||
|
p._apply_provider_field(message_start)
|
||||||
|
assert message_start["provider"] == "openrouter:Anthropic"
|
||||||
|
|
||||||
|
created: dict = {"type": "response.created", "response": {"provider": "OpenAI"}}
|
||||||
|
p._apply_provider_field(created)
|
||||||
|
assert created["provider"] == "openrouter:OpenAI"
|
||||||
|
|
||||||
|
|
||||||
|
def test_stamp_streamed_provider_carries_earlier_provider() -> None:
|
||||||
|
"""Events without their own provider inherit the one reported earlier in
|
||||||
|
the stream instead of becoming ``unknown``."""
|
||||||
|
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
|
||||||
|
first: dict = {"provider": "Fireworks"}
|
||||||
|
carried = p._stamp_streamed_provider(first, None)
|
||||||
|
delta: dict = {"type": "content_block_delta"}
|
||||||
|
assert p._stamp_streamed_provider(delta, carried) == "Fireworks"
|
||||||
|
assert first["provider"] == delta["provider"] == "openrouter:Fireworks"
|
||||||
|
|
||||||
|
|
||||||
def test_apply_provider_field_idempotent_for_direct_upstream() -> None:
|
def test_apply_provider_field_idempotent_for_direct_upstream() -> None:
|
||||||
"""Calling twice on a direct upstream payload keeps the same value and
|
"""Calling twice on a direct upstream payload keeps the same value and
|
||||||
never nests the prefix (no ``anthropic:anthropic``)."""
|
never nests the prefix (no ``anthropic:anthropic``)."""
|
||||||
|
|||||||
@@ -52,3 +52,50 @@ async def test_x_cashu_responses_stream_reports_complete_provider_path() -> None
|
|||||||
|
|
||||||
payload = json.loads((await _body(response)).decode().removeprefix("data: "))
|
payload = json.loads((await _body(response)).decode().removeprefix("data: "))
|
||||||
assert payload["provider"] == "openrouter:z.ai"
|
assert payload["provider"] == "openrouter:z.ai"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_x_cashu_messages_stream_carries_provider_to_later_events() -> None:
|
||||||
|
provider = OpenRouterUpstreamProvider(api_key="test-key")
|
||||||
|
events = [
|
||||||
|
{"type": "message_start", "message": {"provider": "Anthropic"}},
|
||||||
|
{"type": "content_block_delta", "delta": {"text": "hi"}},
|
||||||
|
]
|
||||||
|
content = "".join(f"data: {json.dumps(e)}\n" for e in events)
|
||||||
|
|
||||||
|
response = await provider.handle_x_cashu_streaming_response(
|
||||||
|
content,
|
||||||
|
httpx.Response(200, headers={"content-type": "text/event-stream"}),
|
||||||
|
amount=1,
|
||||||
|
unit="sat",
|
||||||
|
max_cost_for_model=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
lines = (await _body(response)).decode().splitlines()
|
||||||
|
stamped = [json.loads(line.removeprefix("data: ")) for line in lines if line]
|
||||||
|
assert [e["provider"] for e in stamped] == ["openrouter:Anthropic"] * 2
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_x_cashu_responses_stream_carries_nested_provider() -> None:
|
||||||
|
provider = OpenRouterUpstreamProvider(api_key="test-key")
|
||||||
|
events = [
|
||||||
|
{"type": "response.created", "response": {"provider": "OpenAI"}},
|
||||||
|
{"type": "response.output_text.delta", "delta": "hi"},
|
||||||
|
]
|
||||||
|
content = "".join(f"data: {json.dumps(e)}\n\n" for e in events)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
provider, "get_x_cashu_cost", new=AsyncMock(return_value=None)
|
||||||
|
):
|
||||||
|
response = await provider.handle_x_cashu_streaming_responses_response(
|
||||||
|
content,
|
||||||
|
httpx.Response(200, headers={"content-type": "text/event-stream"}),
|
||||||
|
amount=1,
|
||||||
|
unit="sat",
|
||||||
|
max_cost_for_model=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
lines = (await _body(response)).decode().splitlines()
|
||||||
|
stamped = [json.loads(line.removeprefix("data: ")) for line in lines if line]
|
||||||
|
assert [e["provider"] for e in stamped] == ["openrouter:OpenAI"] * 2
|
||||||
|
|||||||
Reference in New Issue
Block a user