mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
add provider field to response
This commit is contained in:
+79
-25
@@ -197,6 +197,22 @@ class BaseUpstreamProvider:
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
|
||||
def _apply_provider_field(self, response_json: object) -> None:
|
||||
"""Stamp the routstr ``provider`` field onto an upstream response payload.
|
||||
|
||||
Format is ``"<provider_type>:<upstream_provider>"`` when the upstream
|
||||
already reported its own provider (e.g. OpenRouter returns
|
||||
``"provider": "Fireworks"``), otherwise just ``"<provider_type>"``
|
||||
for direct upstreams.
|
||||
"""
|
||||
if not isinstance(response_json, dict):
|
||||
return
|
||||
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
|
||||
|
||||
def inject_cost_metadata(
|
||||
self,
|
||||
response_json: dict,
|
||||
@@ -204,6 +220,7 @@ class BaseUpstreamProvider:
|
||||
key: ApiKey,
|
||||
) -> None:
|
||||
"""Unifies the injection of cost and usage metadata across all completion types."""
|
||||
self._apply_provider_field(response_json)
|
||||
if isinstance(cost_data, dict):
|
||||
total_msats = cost_data.get("total_msats", 0)
|
||||
total_usd = cost_data.get("total_usd", 0.0)
|
||||
@@ -723,6 +740,7 @@ class BaseUpstreamProvider:
|
||||
):
|
||||
obj = json.loads(part)
|
||||
if isinstance(obj, dict):
|
||||
self._apply_provider_field(obj)
|
||||
if obj.get("model"):
|
||||
last_model_seen = str(obj.get("model"))
|
||||
if requested_model:
|
||||
@@ -889,6 +907,7 @@ class BaseUpstreamProvider:
|
||||
try:
|
||||
content = await response.aread()
|
||||
response_json = json.loads(content)
|
||||
self._apply_provider_field(response_json)
|
||||
|
||||
logger.debug(
|
||||
"Parsed response JSON",
|
||||
@@ -1068,6 +1087,7 @@ class BaseUpstreamProvider:
|
||||
try:
|
||||
obj = json.loads(part)
|
||||
if isinstance(obj, dict):
|
||||
self._apply_provider_field(obj)
|
||||
if obj.get("model"):
|
||||
last_model_seen = str(obj.get("model"))
|
||||
if requested_model:
|
||||
@@ -1261,6 +1281,7 @@ class BaseUpstreamProvider:
|
||||
try:
|
||||
content = await response.aread()
|
||||
response_json = json.loads(content)
|
||||
self._apply_provider_field(response_json)
|
||||
|
||||
logger.debug(
|
||||
"Parsed Responses API response JSON",
|
||||
@@ -1499,6 +1520,11 @@ class BaseUpstreamProvider:
|
||||
if msg and msg.get("model"):
|
||||
last_model_seen = str(msg.get("model"))
|
||||
|
||||
provider_added = (
|
||||
"provider" not in data
|
||||
)
|
||||
self._apply_provider_field(data)
|
||||
|
||||
if requested_model:
|
||||
# Apply requested_model override
|
||||
model_updated = False
|
||||
@@ -1509,9 +1535,12 @@ class BaseUpstreamProvider:
|
||||
data["model"] = requested_model
|
||||
model_updated = True
|
||||
|
||||
if model_updated:
|
||||
if model_updated or provider_added:
|
||||
line = "data: " + json.dumps(data)
|
||||
changed = True
|
||||
elif provider_added:
|
||||
line = "data: " + json.dumps(data)
|
||||
changed = True
|
||||
|
||||
if usage := msg.get("usage"):
|
||||
input_tokens += usage.get("input_tokens", 0)
|
||||
@@ -1833,6 +1862,7 @@ class BaseUpstreamProvider:
|
||||
)
|
||||
|
||||
response_json = messages_dispatch.coerce_litellm_payload(result)
|
||||
self._apply_provider_field(response_json)
|
||||
if requested_model and "model" in response_json:
|
||||
response_json["model"] = requested_model
|
||||
|
||||
@@ -3145,18 +3175,29 @@ class BaseUpstreamProvider:
|
||||
},
|
||||
)
|
||||
|
||||
if cost_data:
|
||||
for i, line in enumerate(lines):
|
||||
if line.startswith("data: "):
|
||||
try:
|
||||
data_json = json.loads(line[6:])
|
||||
if "usage" in data_json and data_json["usage"]:
|
||||
data_json["usage"]["cost_sats"] = (
|
||||
cost_data.total_msats // 1000
|
||||
)
|
||||
lines[i] = "data: " + json.dumps(data_json)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
for i, line in enumerate(lines):
|
||||
if line.startswith("data: "):
|
||||
try:
|
||||
data_json = json.loads(line[6:])
|
||||
if not isinstance(data_json, dict):
|
||||
continue
|
||||
changed = False
|
||||
if "provider" not in data_json:
|
||||
self._apply_provider_field(data_json)
|
||||
changed = True
|
||||
if (
|
||||
cost_data
|
||||
and "usage" in data_json
|
||||
and data_json["usage"]
|
||||
):
|
||||
data_json["usage"]["cost_sats"] = (
|
||||
cost_data.total_msats // 1000
|
||||
)
|
||||
changed = True
|
||||
if changed:
|
||||
lines[i] = "data: " + json.dumps(data_json)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
async def generate() -> AsyncGenerator[bytes, None]:
|
||||
for line in lines:
|
||||
@@ -3200,6 +3241,7 @@ class BaseUpstreamProvider:
|
||||
|
||||
try:
|
||||
response_json = json.loads(content_str)
|
||||
self._apply_provider_field(response_json)
|
||||
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
|
||||
|
||||
if cost_data and "usage" in response_json:
|
||||
@@ -4121,18 +4163,29 @@ class BaseUpstreamProvider:
|
||||
},
|
||||
)
|
||||
|
||||
if cost_data:
|
||||
for i, line in enumerate(lines):
|
||||
if line.startswith("data: "):
|
||||
try:
|
||||
data_json = json.loads(line[6:])
|
||||
if "usage" in data_json and data_json["usage"]:
|
||||
data_json["usage"]["cost_sats"] = (
|
||||
cost_data.total_msats // 1000
|
||||
)
|
||||
lines[i] = "data: " + json.dumps(data_json)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
for i, line in enumerate(lines):
|
||||
if line.startswith("data: "):
|
||||
try:
|
||||
data_json = json.loads(line[6:])
|
||||
if not isinstance(data_json, dict):
|
||||
continue
|
||||
changed = False
|
||||
if "provider" not in data_json:
|
||||
self._apply_provider_field(data_json)
|
||||
changed = True
|
||||
if (
|
||||
cost_data
|
||||
and "usage" in data_json
|
||||
and data_json["usage"]
|
||||
):
|
||||
data_json["usage"]["cost_sats"] = (
|
||||
cost_data.total_msats // 1000
|
||||
)
|
||||
changed = True
|
||||
if changed:
|
||||
lines[i] = "data: " + json.dumps(data_json)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
async def generate() -> AsyncGenerator[bytes, None]:
|
||||
for line in lines:
|
||||
@@ -4164,6 +4217,7 @@ class BaseUpstreamProvider:
|
||||
|
||||
try:
|
||||
response_json = json.loads(content_str)
|
||||
self._apply_provider_field(response_json)
|
||||
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
|
||||
|
||||
if cost_data and "usage" in response_json:
|
||||
|
||||
@@ -0,0 +1,105 @@
|
||||
from routstr.upstream.anthropic import AnthropicUpstreamProvider
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
from routstr.upstream.openrouter import OpenRouterUpstreamProvider
|
||||
|
||||
|
||||
def _make_provider(cls: type, provider_type: str) -> BaseUpstreamProvider:
|
||||
p = cls(api_key="test_key")
|
||||
assert p.provider_type == provider_type
|
||||
return p
|
||||
|
||||
|
||||
def test_apply_provider_field_direct_upstream() -> None:
|
||||
"""For a direct upstream (no upstream-reported provider), the field
|
||||
is just the provider_type string."""
|
||||
p = _make_provider(AnthropicUpstreamProvider, "anthropic")
|
||||
data: dict = {"id": "msg_1", "model": "claude-3-5-sonnet"}
|
||||
p._apply_provider_field(data)
|
||||
assert data["provider"] == "anthropic"
|
||||
|
||||
|
||||
def test_apply_provider_field_openrouter_passthrough() -> None:
|
||||
"""OpenRouter responses include an upstream ``provider`` string —
|
||||
routstr should prefix with its own provider_type."""
|
||||
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
|
||||
data: dict = {
|
||||
"id": "gen-abc",
|
||||
"model": "anthropic/claude-3.5-sonnet",
|
||||
"provider": "Anthropic",
|
||||
}
|
||||
p._apply_provider_field(data)
|
||||
assert data["provider"] == "openrouter:Anthropic"
|
||||
|
||||
|
||||
def test_apply_provider_field_openrouter_no_upstream_provider() -> None:
|
||||
"""If OpenRouter omits the provider field, fall back to provider_type."""
|
||||
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
|
||||
data: dict = {"id": "gen-abc"}
|
||||
p._apply_provider_field(data)
|
||||
assert data["provider"] == "openrouter"
|
||||
|
||||
|
||||
def test_apply_provider_field_strips_whitespace() -> None:
|
||||
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
|
||||
data: dict = {"provider": " Fireworks "}
|
||||
p._apply_provider_field(data)
|
||||
assert data["provider"] == "openrouter:Fireworks"
|
||||
|
||||
|
||||
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"
|
||||
|
||||
|
||||
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"
|
||||
|
||||
|
||||
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."""
|
||||
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.
|
||||
|
||||
|
||||
def test_apply_provider_field_ignores_non_dict() -> None:
|
||||
"""Lists / primitives must be skipped silently."""
|
||||
p = _make_provider(AnthropicUpstreamProvider, "anthropic")
|
||||
# Should not raise.
|
||||
p._apply_provider_field([1, 2, 3]) # type: ignore[arg-type]
|
||||
p._apply_provider_field("hello") # type: ignore[arg-type]
|
||||
p._apply_provider_field(None) # type: ignore[arg-type]
|
||||
|
||||
|
||||
def test_inject_cost_metadata_sets_provider() -> None:
|
||||
"""``inject_cost_metadata`` is the unified injection point and must
|
||||
also stamp the provider field."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
|
||||
key = MagicMock(spec=ApiKey)
|
||||
key.balance = 1000
|
||||
|
||||
response_json: dict = {
|
||||
"model": "anthropic/claude-3.5-sonnet",
|
||||
"provider": "Anthropic",
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 5},
|
||||
}
|
||||
cost_data = {"total_msats": 2500, "total_usd": 0.0025}
|
||||
p.inject_cost_metadata(response_json, cost_data, key)
|
||||
|
||||
assert response_json["provider"] == "openrouter:Anthropic"
|
||||
Reference in New Issue
Block a user