mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-06 04:38:22 +00:00
merge: sync latest main into upstream timeout PR
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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}"
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user