diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 7e5b67ef..b7976a0f 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -217,6 +217,20 @@ def _responses_usage_payload(data_json: dict) -> dict: 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: """Re-frame one parsed event, re-prefixing every line of a multi-line data.""" body = "".join(f"{line}\n" for line in field_lines) @@ -488,8 +502,7 @@ class BaseUpstreamProvider: return response_json["provider_url"] = public_provider_url(self.base_url) provider_type = (self.provider_type or "").strip() - existing = response_json.get("provider") - existing_str = existing.strip() if isinstance(existing, str) else "" + existing_str = _reported_provider(response_json) or "" if not existing_str: response_json["provider"] = provider_type return @@ -501,6 +514,17 @@ class BaseUpstreamProvider: return 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( self, *, @@ -1185,6 +1209,7 @@ class BaseUpstreamProvider: usage_finalized = False last_model_seen: str | None = None + provider_seen: str | None = None async def finalize_db_only() -> None: nonlocal usage_finalized @@ -1259,6 +1284,7 @@ class BaseUpstreamProvider: end of stream. """ nonlocal last_model_seen, usage_chunk_data, done_seen, stream_id + nonlocal provider_seen event = raw_event.strip(b"\r\n") if not event: @@ -1298,7 +1324,7 @@ class BaseUpstreamProvider: if isinstance(obj, dict): usage_estimator.observe(obj) - self._apply_provider_field(obj) + provider_seen = self._stamp_streamed_provider(obj, provider_seen) if obj.get("model"): last_model_seen = str(obj.get("model")) if requested_model: @@ -1424,6 +1450,7 @@ class BaseUpstreamProvider: if legacy_completion else "chat.completion.chunk", "model": last_model_seen or "unknown", + "provider": provider_seen, "choices": [], "usage": { "prompt_tokens": cost_data.get("input_tokens", 0), @@ -1672,6 +1699,7 @@ class BaseUpstreamProvider: usage_finalized = False last_model_seen: str | None = None + provider_seen: str | None = None async def finalize_db_only() -> None: nonlocal usage_finalized @@ -1735,7 +1763,7 @@ class BaseUpstreamProvider: and preserves ``event:``/``id:`` fields attached to their data 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 event = raw_event.strip(b"\r\n") @@ -1771,7 +1799,7 @@ class BaseUpstreamProvider: obj = json_codec.loads(data) if isinstance(obj, dict): - self._apply_provider_field(obj) + provider_seen = self._stamp_streamed_provider(obj, provider_seen) if obj.get("model"): last_model_seen = str(obj.get("model")) if requested_model: @@ -1862,6 +1890,7 @@ class BaseUpstreamProvider: "type": "response.failed" if guarded_chunks.timed_out else "response.completed", + "provider": provider_seen, "response": { "model": last_model_seen or "unknown", "usage": { @@ -2245,6 +2274,7 @@ class BaseUpstreamProvider: usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_finalized = False last_model_seen: str | None = None + provider_seen: str | None = None async def finalize_without_usage() -> bytes | None: nonlocal usage_finalized @@ -2294,7 +2324,7 @@ class BaseUpstreamProvider: async def stream_with_cost( max_cost_for_model: int, ) -> AsyncGenerator[bytes, None]: - nonlocal usage_finalized, last_model_seen + nonlocal usage_finalized, last_model_seen, provider_seen stored_chunks: list[bytes] = [] input_tokens: int = 0 output_tokens: int = 0 @@ -2351,7 +2381,9 @@ class BaseUpstreamProvider: last_model_seen = str(msg.get("model")) provider_added = "provider" not in data - self._apply_provider_field(data) + provider_seen = self._stamp_streamed_provider( + data, provider_seen + ) if requested_model: # Apply requested_model override @@ -2469,6 +2501,7 @@ class BaseUpstreamProvider: try: combined_data = { "model": last_model_seen or "unknown", + "provider": provider_seen, "usage": usage_data, } cost_data = await adjust_payment_for_tokens( @@ -4249,6 +4282,7 @@ class BaseUpstreamProvider: }, ) + provider_seen: str | None = None for i, line in enumerate(lines): if line.startswith("data: "): try: @@ -4256,7 +4290,9 @@ class BaseUpstreamProvider: if not isinstance(data_json, dict): continue 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 if cost_data and "usage" in data_json and data_json["usage"]: _inject_cost_into_usage(data_json, cost_data) @@ -5317,6 +5353,7 @@ class BaseUpstreamProvider: }, ) + provider_seen: str | None = None for i, (fields, data) in enumerate(events): if data.strip() == "[DONE]": continue @@ -5327,7 +5364,7 @@ class BaseUpstreamProvider: if not isinstance(data_json, dict): continue 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 payload = _responses_usage_payload(data_json) if cost_data and isinstance(payload.get("usage"), dict): diff --git a/routstr/upstream/generic.py b/routstr/upstream/generic.py index 03bfa015..c9edf109 100644 --- a/routstr/upstream/generic.py +++ b/routstr/upstream/generic.py @@ -5,7 +5,7 @@ from urllib.parse import urlparse import httpx -from .base import BaseUpstreamProvider +from .base import BaseUpstreamProvider, _reported_provider from .model_paths import public_provider_url from .pricing_resolver import ( FallbackPricingResolver, @@ -60,8 +60,7 @@ class GenericUpstreamProvider(BaseUpstreamProvider): """ if not isinstance(response_json, dict): return - existing = response_json.get("provider") - if not (isinstance(existing, str) and existing.strip()): + if _reported_provider(response_json) is None: response_json["provider"] = ( urlparse(public_provider_url(self.base_url)).hostname or self.upstream_name diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index 488129a8..1951768d 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -36,6 +36,12 @@ from .reasoning_effort import adapt_messages_body_for_litellm logger = get_logger(__name__) +# Sent in place of a blank upstream key. LiteLLM treats ``""`` as missing and +# falls back to the provider's env var (e.g. ``OPENAI_API_KEY``), failing with +# an AuthenticationError for keyless upstreams such as self-hosted +# OpenAI-compatible servers, which the chat path reaches without auth. +KEYLESS_UPSTREAM_API_KEY = "no-key" + # Anthropic-Messages-only fields that don't translate to OpenAI # Chat Completions. ``litellm.drop_params`` only filters *known* # unsupported params; these newer/extension fields get passed through @@ -507,6 +513,32 @@ async def dispatch_anthropic_messages( model_suffix = adapt_request(body) if adapt_request else "" + # LiteLLM turns Anthropic's server-side web_search tool into the OpenAI + # `web_search_options` parameter. Generic OpenAI-compatible chat endpoints + # (including those serving Claude through a proxy) may reject that field. + # Only a provider with an explicit adaptation (e.g. Venice's model suffix) + # can preserve search semantics; do not silently remove the tool and return + # an answer that never searched. Native /v1/messages providers bypass this + # dispatcher and receive the original tool unchanged. + tools = body.get("tools") + if provider_prefix == "openai/" and isinstance(tools, list) and any( + isinstance(tool, dict) + and ( + ( + isinstance(tool.get("type"), str) + and tool["type"].startswith("web_search") + ) + or tool.get("name") == "web_search" + ) + for tool in tools + ): + raise UpstreamError( + "This upstream does not support Anthropic web search through " + "OpenAI-compatible /v1/messages translation", + status_code=400, + code="UNSUPPORTED_WEB_SEARCH", + ) + # Convention: `model.id` is the canonical upstream model name; # `forwarded_model_id` is the public alias the internal API exposes # and echoes back to the client. @@ -519,7 +551,7 @@ async def dispatch_anthropic_messages( kwargs: dict = { "model": litellm_model, "api_base": base_url, - "api_key": api_key, + "api_key": api_key or KEYLESS_UPSTREAM_API_KEY, "stream": upstream_stream, **body, } diff --git a/routstr/upstream/openrouter.py b/routstr/upstream/openrouter.py index 9ca190ce..3ff13d90 100644 --- a/routstr/upstream/openrouter.py +++ b/routstr/upstream/openrouter.py @@ -2,13 +2,27 @@ from typing import TYPE_CHECKING import httpx +from ..core.logging import get_logger 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 if TYPE_CHECKING: from ..core.db import UpstreamProviderRow +logger = get_logger(__name__) + +_UNKNOWN_SUB_PROVIDER = "unknown" + + +def _carries_usage(payload: dict) -> bool: + """Whether a payload holds usage, at top level or in the Anthropic + ``message`` / Responses ``response`` envelope.""" + return any( + isinstance(obj, dict) and isinstance(obj.get("usage"), dict) + for obj in (payload, payload.get("message"), payload.get("response")) + ) + class OpenRouterUpstreamProvider(BaseUpstreamProvider): """Upstream provider specifically configured for OpenRouter API.""" @@ -27,7 +41,8 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): - Real upstream sub-provider (e.g. ``"GMICloud"``) -> ``"openrouter:GMICloud"``. - Missing sub-provider, or one that merely echoes ``"openrouter"`` -> - ``"unknown"``. + ``"openrouter:unknown"``: the router is still known even when the + serving provider is not (e.g. the Responses API never reports it). - Idempotent: re-stamping never produces ``"openrouter:openrouter:..."``; the ``openrouter:`` prefix appears at most once. """ @@ -35,15 +50,28 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider): return response_json["provider_url"] = public_provider_url(self.base_url) provider_type = (self.provider_type or "").strip() - existing = response_json.get("provider") - sub = existing.strip() if isinstance(existing, str) else "" + sub = _reported_provider(response_json) or "" # Strip any already-applied "openrouter:" prefixes (idempotency). prefix = f"{provider_type}:" while sub.lower().startswith(prefix.lower()): sub = sub[len(prefix) :].strip() + # Already stamped as unknown on an earlier pass; keep it without + # warning again. + if sub.lower() == _UNKNOWN_SUB_PROVIDER: + response_json["provider"] = f"{provider_type}:{_UNKNOWN_SUB_PROVIDER}" + return # 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" + # Warn only on the billed payload, not on every stream chunk. + if _carries_usage(response_json): + logger.warning( + "OpenRouter did not report the serving provider", + extra={ + "model": response_json.get("model"), + "response_id": response_json.get("id"), + }, + ) + response_json["provider"] = f"{provider_type}:{_UNKNOWN_SUB_PROVIDER}" return response_json["provider"] = f"{provider_type}:{sub}" diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index c9b13bd2..7e2da476 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -1755,3 +1755,31 @@ async def test_x_cashu_zero_value_rejected_not_forwarded( assert body["error"]["code"] == "cashu_token_zero_value" # Spent-to-zero token must not be echoed back for retry. assert "X-Cashu" not in response.headers + + +@pytest.mark.asyncio +async def test_dispatch_passes_placeholder_key_for_keyless_upstream() -> None: + """A blank upstream key must not reach litellm, which would fall back to + OPENAI_API_KEY and fail with an AuthenticationError.""" + provider = BaseUpstreamProvider(base_url="http://localhost:8000/v1", api_key="") + captured_kwargs: dict[str, Any] = {} + + async def fake_acreate(**kwargs: Any) -> AsyncIterator[dict]: + captured_kwargs.update(kwargs) + + async def no_events() -> AsyncIterator[dict]: + return + yield + + return no_events() + + with patch( + "litellm.anthropic.messages.acreate", + new=AsyncMock(side_effect=fake_acreate), + ): + await provider._dispatch_anthropic_messages( + request_body=_anthropic_request_body(stream=True), + model_obj=_make_model(), + ) + + assert captured_kwargs["api_key"] == "no-key" diff --git a/tests/unit/test_provider_field_injection.py b/tests/unit/test_provider_field_injection.py index 86942ae6..6caea063 100644 --- a/tests/unit/test_provider_field_injection.py +++ b/tests/unit/test_provider_field_injection.py @@ -1,3 +1,5 @@ +from unittest.mock import patch + from routstr.upstream.anthropic import AnthropicUpstreamProvider from routstr.upstream.base import BaseUpstreamProvider from routstr.upstream.generic import GenericUpstreamProvider @@ -33,12 +35,12 @@ def test_apply_provider_field_openrouter_passthrough() -> None: def test_apply_provider_field_openrouter_no_upstream_provider() -> None: - """If OpenRouter omits the provider field, the real serving provider is - unknown — a bare ``openrouter`` value carries no information.""" + """If OpenRouter omits the provider field, the serving provider is + unknown but the router is not.""" p = _make_provider(OpenRouterUpstreamProvider, "openrouter") data: dict = {"id": "gen-abc"} p._apply_provider_field(data) - assert data["provider"] == "unknown" + assert data["provider"] == "openrouter:unknown" def test_apply_provider_field_openrouter_echoes_router_name() -> None: @@ -46,7 +48,35 @@ def test_apply_provider_field_openrouter_echoes_router_name() -> None: p = _make_provider(OpenRouterUpstreamProvider, "openrouter") data: dict = {"provider": "openrouter"} p._apply_provider_field(data) - assert data["provider"] == "unknown" + assert data["provider"] == "openrouter:unknown" + + +def test_apply_provider_field_openrouter_unknown_is_idempotent() -> None: + """Re-stamping an unknown payload (e.g. in inject_cost_metadata) keeps + ``openrouter:unknown`` instead of reading ``unknown`` as a sub-provider.""" + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + data: dict = {"id": "gen-abc"} + p._apply_provider_field(data) + p._apply_provider_field(data) + assert data["provider"] == "openrouter:unknown" + + +def test_apply_provider_field_openrouter_warns_once_on_billed_payload() -> None: + """A missing provider is logged on the payload carrying usage, not on + every stream chunk or on a re-stamp.""" + p = _make_provider(OpenRouterUpstreamProvider, "openrouter") + chunk: dict = {"type": "response.output_text.delta", "delta": "hi"} + completed: dict = { + "type": "response.completed", + "response": {"id": "gen-abc", "usage": {"input_tokens": 1}}, + } + with patch("routstr.upstream.openrouter.logger.warning") as warning: + p._apply_provider_field(chunk) + p._apply_provider_field(completed) + p._apply_provider_field(completed) + + warning.assert_called_once() + assert chunk["provider"] == completed["provider"] == "openrouter:unknown" def test_apply_provider_field_openrouter_idempotent_no_double_prefix() -> None: @@ -79,14 +109,41 @@ 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"] == "unknown" + assert data["provider"] == "openrouter: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"] == "unknown" + assert data["provider"] == "openrouter: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: diff --git a/tests/unit/test_venice_web_search.py b/tests/unit/test_venice_web_search.py index a836e956..380ccc33 100644 --- a/tests/unit/test_venice_web_search.py +++ b/tests/unit/test_venice_web_search.py @@ -123,13 +123,35 @@ async def test_requests_without_web_search_are_untouched() -> None: @pytest.mark.asyncio -async def test_other_providers_keep_their_existing_behaviour() -> None: - """The base hook is a no-op, so no non-Venice upstream changes shape.""" +async def test_generic_openai_upstream_rejects_untranslatable_web_search() -> None: + """Do not let LiteLLM send unsupported web_search_options to a generic API.""" provider = BaseUpstreamProvider(base_url="http://test", api_key="k") + with pytest.raises(UpstreamError) as excinfo: + await _dispatch(provider, _body(tools=[WEB_SEARCH_TOOL, FUNCTION_TOOL])) + + assert excinfo.value.status_code == 400 + assert excinfo.value.code == "UNSUPPORTED_WEB_SEARCH" + + +@pytest.mark.asyncio +async def test_generic_openai_upstream_still_accepts_function_tools() -> None: + provider = BaseUpstreamProvider(base_url="http://test", api_key="k") + + kwargs = await _dispatch(provider, _body(tools=[FUNCTION_TOOL])) + + assert kwargs["model"] == "openai/deepseek-v4-flash-0731" + assert kwargs["tools"] == [FUNCTION_TOOL] + + +@pytest.mark.asyncio +async def test_non_openai_adapter_can_still_handle_search_tool() -> None: + provider = BaseUpstreamProvider( + base_url="https://openrouter.ai/api/v1", api_key="k" + ) + kwargs = await _dispatch(provider, _body(tools=[WEB_SEARCH_TOOL])) - assert kwargs["model"] == "openai/deepseek-v4-flash-0731" assert kwargs["tools"] == [WEB_SEARCH_TOOL] diff --git a/tests/unit/test_x_cashu_provider_path.py b/tests/unit/test_x_cashu_provider_path.py index 11b47d6f..6b4fdec2 100644 --- a/tests/unit/test_x_cashu_provider_path.py +++ b/tests/unit/test_x_cashu_provider_path.py @@ -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: ")) 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