diff --git a/routstr/payment/models.py b/routstr/payment/models.py
index 95d7ad89..afed3238 100644
--- a/routstr/payment/models.py
+++ b/routstr/payment/models.py
@@ -7,7 +7,7 @@ from fastapi import APIRouter, Depends
from pydantic.v1 import BaseModel
from sqlmodel.ext.asyncio.session import AsyncSession
-from ..core.db import ModelRow, get_session
+from ..core.db import ModelRow, UpstreamProviderRow, get_session
from ..core.logging import get_logger
from ..core.settings import settings
from .price import sats_usd_price
@@ -405,6 +405,76 @@ async def update_sats_pricing() -> None:
logger.error(f"Error updating sats pricing: {e}")
+class ModelTestRequest(BaseModel):
+ model_id: str
+ endpoint_type: str
+ request_data: dict
+
+
+@models_router.post("/api/models/test")
+async def test_model(
+ payload: ModelTestRequest,
+ session: AsyncSession = Depends(get_session),
+) -> dict:
+ """Test a model by sending a request through its configured upstream provider."""
+ from sqlmodel import select
+
+ result = await session.execute(
+ select(ModelRow).where(ModelRow.id == payload.model_id)
+ )
+ model_row = result.scalars().first()
+
+ if not model_row:
+ return {
+ "success": False,
+ "error": f"Model '{payload.model_id}' not found in database",
+ "status_code": 404,
+ }
+
+ provider = await session.get(UpstreamProviderRow, model_row.upstream_provider_id)
+ if not provider:
+ return {
+ "success": False,
+ "error": "Upstream provider not found",
+ "status_code": 404,
+ }
+
+ base_url = provider.base_url.rstrip("/")
+ if payload.endpoint_type == "chat-completions":
+ url = f"{base_url}/chat/completions"
+ else:
+ url = f"{base_url}/{payload.endpoint_type}"
+
+ actual_model_id = model_row.forwarded_model_id or model_row.id
+ request_data = dict(payload.request_data)
+ request_data["model"] = actual_model_id
+
+ headers = {
+ "Content-Type": "application/json",
+ "Authorization": f"Bearer {provider.api_key}",
+ }
+
+ try:
+ async with httpx.AsyncClient(timeout=30.0) as client:
+ response = await client.post(url, json=request_data, headers=headers)
+ try:
+ response_data = response.json()
+ except Exception:
+ response_data = {"raw": response.text}
+
+ return {
+ "success": response.status_code < 400,
+ "data": response_data,
+ "status_code": response.status_code,
+ }
+ except Exception as e:
+ return {
+ "success": False,
+ "error": str(e),
+ "status_code": 500,
+ }
+
+
@models_router.get("/v1/models")
@models_router.get("/v1/models/", include_in_schema=False)
@models_router.get("/models")
diff --git a/routstr/proxy.py b/routstr/proxy.py
index ddab1c31..9d2dbf09 100644
--- a/routstr/proxy.py
+++ b/routstr/proxy.py
@@ -69,7 +69,25 @@ def get_upstreams() -> list[BaseUpstreamProvider]:
def get_model_instance(model_id: str) -> Model | None:
"""Get Model instance by ID from global cache."""
- return _model_instances.get(model_id.lower())
+ if not model_id:
+ return None
+
+ model_id_lower = model_id.lower()
+ # Try exact match first
+ if model := _model_instances.get(model_id_lower):
+ return model
+
+ # Try stripping common version suffixes (e.g., -20251222)
+ # This handles cases where upstream returns a specific version
+ # but we only track the base model name.
+ import re
+
+ base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
+ if base_model_id != model_id_lower:
+ if model := _model_instances.get(base_model_id):
+ return model
+
+ return None
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None:
diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py
index b506e388..458e5dfd 100644
--- a/routstr/upstream/base.py
+++ b/routstr/upstream/base.py
@@ -119,6 +119,51 @@ class BaseUpstreamProvider:
"can_show_balance": False,
}
+ def inject_cost_metadata(
+ self,
+ response_json: dict,
+ cost_data: CostData | MaxCostData | dict,
+ key: ApiKey,
+ ) -> None:
+ """Unifies the injection of cost and usage metadata across all completion types."""
+ if isinstance(cost_data, dict):
+ total_msats = cost_data.get("total_msats", 0)
+ total_usd = cost_data.get("total_usd", 0.0)
+ cost_dict = cost_data
+ else:
+ total_msats = cost_data.total_msats
+ total_usd = cost_data.total_usd
+ cost_dict = cost_data.dict()
+
+ sats_cost = total_msats // 1000
+
+ # Inject into top-level usage block (OpenAI/Anthropic style)
+ if "usage" in response_json:
+ response_json["usage"]["cost"] = total_usd
+ response_json["usage"]["cost_sats"] = sats_cost
+ response_json["usage"]["remaining_balance_msats"] = key.balance
+
+ # Inject into Anthropic nested usage block if present
+ if (
+ "message" in response_json
+ and isinstance(response_json["message"], dict)
+ and "usage" in response_json["message"]
+ ):
+ response_json["message"]["usage"]["sats_cost"] = sats_cost
+
+ # Unified Routstr metadata
+ response_json["metadata"] = response_json.get("metadata", {})
+ response_json["metadata"]["routstr"] = {
+ "cost": cost_dict,
+ "sats_cost": sats_cost,
+ "remaining_balance_msats": key.balance,
+ }
+
+ # Legacy/Compatibility fields
+ response_json["cost"] = cost_dict.copy()
+ response_json["cost"]["sats_cost"] = sats_cost
+ response_json["cost"]["remaining_balance_msats"] = key.balance
+
def prepare_headers(self, request_headers: dict) -> dict:
"""Prepare headers for upstream request by removing proxy-specific headers and adding authentication.
@@ -376,75 +421,83 @@ class BaseUpstreamProvider:
"""
pass
- async def map_upstream_error_response(
+ async def forward_upstream_error_response(
self, request: Request, path: str, upstream_response: httpx.Response
) -> Response:
- """Map upstream error responses to appropriate proxy error responses.
-
- Args:
- request: Original FastAPI request
- path: Request path
- upstream_response: Response from upstream service
-
- Returns:
- Mapped error response with appropriate status code and error type
- """
+ """Log upstream errors and forward the upstream response unchanged."""
status_code = upstream_response.status_code
headers = dict(upstream_response.headers)
- content_type = headers.get("content-type", "")
+ content_type = headers.get("content-type") or headers.get("Content-Type", "")
+ upstream_request_id = (
+ headers.get("request-id")
+ or headers.get("Request-Id")
+ or headers.get("x-request-id")
+ or headers.get("X-Request-Id")
+ or headers.get("anthropic-request-id")
+ or headers.get("openai-request-id")
+ )
+
+ body_read_error = None
try:
body_bytes = await upstream_response.aread()
- except Exception:
+ except Exception as exc:
body_bytes = b""
+ body_read_error = f"{type(exc).__name__}: {exc}"
message, upstream_code = self._extract_upstream_error_message(body_bytes)
- lowered_message = message.lower()
- lowered_code = (upstream_code or "").lower()
+ body_preview = body_bytes.decode("utf-8", errors="ignore").strip()[:500]
- error_type = "upstream_error"
- mapped_status = 502
-
- if status_code in (400, 422):
- error_type = "invalid_request_error"
- mapped_status = 400
- elif status_code in (401, 403):
- error_type = "upstream_auth_error"
- mapped_status = 502
- elif status_code == 404:
- if path.endswith("chat/completions"):
- error_type = "invalid_model"
- mapped_status = 400
- if not message or message == "Upstream request failed":
- message = "Requested model is not available upstream"
- elif "model" in lowered_message or "model" in lowered_code:
- error_type = "invalid_model"
- mapped_status = 400
- if not message or message == "Upstream request failed":
- message = "Requested model is not available upstream"
- else:
- error_type = "upstream_error"
- mapped_status = 502
- elif status_code == 429:
- error_type = "rate_limit_exceeded"
- mapped_status = 429
- elif status_code >= 500:
- error_type = "upstream_error"
- mapped_status = 502
-
- logger.debug(
- "Mapped upstream error",
+ logger.warning(
+ "Forwarding upstream error response as-is",
extra={
"path": path,
+ "provider": self.provider_type,
"upstream_status": status_code,
- "mapped_status": mapped_status,
- "error_type": error_type,
+ "upstream_code": upstream_code,
"upstream_content_type": content_type,
+ "upstream_request_id": upstream_request_id,
"message_preview": message[:200],
+ "body_preview": body_preview,
+ "body_read_error": body_read_error,
+ "method": request.method,
},
)
- return create_error_response(
- error_type, message, mapped_status, request=request
+ for header_name in (
+ "content-length",
+ "Content-Length",
+ "transfer-encoding",
+ "Transfer-Encoding",
+ "content-encoding",
+ "Content-Encoding",
+ "connection",
+ "Connection",
+ "keep-alive",
+ "Keep-Alive",
+ "proxy-authenticate",
+ "Proxy-Authenticate",
+ "proxy-authorization",
+ "Proxy-Authorization",
+ "te",
+ "TE",
+ "trailer",
+ "Trailer",
+ "upgrade",
+ "Upgrade",
+ ):
+ headers.pop(header_name, None)
+
+ if not content_type:
+ headers.pop("content-type", None)
+ headers.pop("Content-Type", None)
+
+ media_type = content_type or None
+
+ return Response(
+ content=body_bytes,
+ status_code=status_code,
+ headers=headers,
+ media_type=media_type,
)
async def handle_streaming_chat_completion(
@@ -1139,7 +1192,11 @@ class BaseUpstreamProvider:
)
async def handle_streaming_messages_completion(
- self, response: httpx.Response, key: ApiKey, max_cost_for_model: int
+ self,
+ response: httpx.Response,
+ key: ApiKey,
+ max_cost_for_model: int,
+ requested_model: str | None = None,
) -> StreamingResponse:
async def stream_with_cost(
max_cost_for_model: int,
@@ -1178,6 +1235,8 @@ class BaseUpstreamProvider:
stored_chunks.append(chunk)
try:
decoded_chunk = chunk.decode("utf-8", errors="ignore")
+ modified_lines = []
+ changed = False
for line in decoded_chunk.split("\n"):
if line.startswith("data: "):
try:
@@ -1187,6 +1246,20 @@ class BaseUpstreamProvider:
if msg and msg.get("model"):
last_model_seen = str(msg.get("model"))
+ if requested_model:
+ # Apply requested_model override
+ model_updated = False
+ if msg:
+ msg["model"] = requested_model
+ model_updated = True
+ if data.get("model"):
+ data["model"] = requested_model
+ model_updated = True
+
+ if model_updated:
+ line = "data: " + json.dumps(data)
+ changed = True
+
if usage := msg.get("usage"):
input_tokens += usage.get("input_tokens", 0)
output_tokens += usage.get(
@@ -1200,10 +1273,14 @@ class BaseUpstreamProvider:
)
except json.JSONDecodeError:
pass
- except Exception:
- pass
+ modified_lines.append(line)
- yield chunk
+ if changed:
+ yield "\n".join(modified_lines).encode("utf-8")
+ else:
+ yield chunk
+ except Exception:
+ yield chunk
usage_data = {
"input_tokens": input_tokens,
@@ -1225,6 +1302,11 @@ class BaseUpstreamProvider:
new_session,
max_cost_for_model,
)
+
+ self.inject_cost_metadata(
+ combined_data, cost_data, fresh_key
+ )
+
usage_finalized = True
yield f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode()
except Exception:
@@ -1264,11 +1346,22 @@ class BaseUpstreamProvider:
session: AsyncSession,
deducted_max_cost: int,
path: str,
+ requested_model: str | None = None,
) -> Response:
try:
content = await response.aread()
response_json = json.loads(content)
+ if requested_model:
+ if "model" in response_json:
+ response_json["model"] = requested_model
+ if (
+ "message" in response_json
+ and isinstance(response_json["message"], dict)
+ and "model" in response_json["message"]
+ ):
+ response_json["message"]["model"] = requested_model
+
if path.endswith("count_tokens") and "usage" not in response_json:
input_tokens = response_json.get("input_tokens", 0)
response_json["usage"] = {"input_tokens": input_tokens}
@@ -1276,7 +1369,8 @@ class BaseUpstreamProvider:
cost_data = await adjust_payment_for_tokens(
key, response_json, session, deducted_max_cost
)
- response_json["cost"] = cost_data
+
+ self.inject_cost_metadata(response_json, cost_data, key)
allowed_headers = {
"content-type",
@@ -1381,15 +1475,28 @@ class BaseUpstreamProvider:
stream=True,
)
- logger.info(
- "Received upstream response",
- extra={
- "status_code": response.status_code,
- "path": path,
- "key_hash": key.hashed_key[:8] + "...",
- "content_type": response.headers.get("content-type", "unknown"),
- },
- )
+ if response.status_code != 200:
+ logger.error(
+ "Received upstream response",
+ extra={
+ "reason_phrase": response.reason_phrase,
+ "status_code": response.status_code,
+ "path": path,
+ "key_hash": key.hashed_key[:8] + "...",
+ "content_type": response.headers.get("content-type", "unknown"),
+ },
+ )
+ else:
+ logger.info(
+ "Received upstream response",
+ extra={
+ "reason_phrase": response.reason_phrase,
+ "status_code": response.status_code,
+ "path": path,
+ "key_hash": key.hashed_key[:8] + "...",
+ "content_type": response.headers.get("content-type", "unknown"),
+ },
+ )
if response.status_code != 200:
if response.status_code >= 500:
@@ -1401,7 +1508,7 @@ class BaseUpstreamProvider:
)
try:
- mapped_error = await self.map_upstream_error_response(
+ mapped_error = await self.forward_upstream_error_response(
request, path, response
)
finally:
@@ -1430,7 +1537,10 @@ class BaseUpstreamProvider:
if is_streaming and response.status_code == 200:
result = await self.handle_streaming_messages_completion(
- response, key, max_cost_for_model
+ response,
+ key,
+ max_cost_for_model,
+ requested_model=original_model_id,
)
background_tasks = BackgroundTasks()
background_tasks.add_task(response.aclose)
@@ -1441,7 +1551,12 @@ class BaseUpstreamProvider:
if response.status_code == 200:
try:
return await self.handle_non_streaming_messages_completion(
- response, key, session, max_cost_for_model, path
+ response,
+ key,
+ session,
+ max_cost_for_model,
+ path,
+ requested_model=original_model_id,
)
finally:
await response.aclose()
@@ -1451,7 +1566,12 @@ class BaseUpstreamProvider:
if response.status_code == 200:
try:
return await self.handle_non_streaming_messages_completion(
- response, key, session, max_cost_for_model, path
+ response,
+ key,
+ session,
+ max_cost_for_model,
+ path,
+ requested_model=original_model_id,
)
finally:
await response.aclose()
@@ -1695,7 +1815,7 @@ class BaseUpstreamProvider:
)
try:
- mapped_error = await self.map_upstream_error_response(
+ mapped_error = await self.forward_upstream_error_response(
request, path, response
)
finally:
@@ -1867,7 +1987,7 @@ class BaseUpstreamProvider:
)
if response.status_code != 200:
try:
- mapped = await self.map_upstream_error_response(
+ mapped = await self.forward_upstream_error_response(
request, path, response
)
finally:
@@ -2521,14 +2641,25 @@ class BaseUpstreamProvider:
stream=True,
)
- logger.debug(
- "Received upstream response",
- extra={
- "status_code": response.status_code,
- "path": path,
- "response_headers": dict(response.headers),
- },
- )
+ if response.status_code != 200:
+ logger.error(
+ "Received upstream response",
+ extra={
+ "reason_phrase": response.reason_phrase,
+ "status_code": response.status_code,
+ "path": path,
+ "response_headers": dict(response.headers),
+ },
+ )
+ else:
+ logger.debug(
+ "Received upstream response",
+ extra={
+ "status_code": response.status_code,
+ "path": path,
+ "response_headers": dict(response.headers),
+ },
+ )
if response.status_code != 200:
logger.warning(
diff --git a/ui/components/api-endpoint-tester.tsx b/ui/components/api-endpoint-tester.tsx
index a77b4248..dbb91266 100644
--- a/ui/components/api-endpoint-tester.tsx
+++ b/ui/components/api-endpoint-tester.tsx
@@ -471,7 +471,11 @@ export function ApiEndpointTester({ models }: ApiEndpointTesterProps) {
testEndpointMutation.mutate(requestData);
};
- const enabledModels = models.filter((model) => model.isEnabled);
+ const enabledModels = Array.from(
+ new Map(
+ models.filter((model) => model.isEnabled).map((m) => [m.id, m])
+ ).values()
+ );
const credentials = selectedModel ? getModelCredentials(selectedModel) : null;
const endpointUrl = credentials
? buildEndpointUrl(
diff --git a/ui/components/model-tester.tsx b/ui/components/model-tester.tsx
index dee500a7..3ecd43b3 100644
--- a/ui/components/model-tester.tsx
+++ b/ui/components/model-tester.tsx
@@ -197,7 +197,11 @@ export function ModelTester({ models }: ModelTesterProps) {
testModelMutation.mutate(request);
};
- const enabledModels = models.filter((model) => model.isEnabled);
+ const enabledModels = Array.from(
+ new Map(
+ models.filter((model) => model.isEnabled).map((m) => [m.id, m])
+ ).values()
+ );
const credentials = selectedModel ? getModelCredentials(selectedModel) : null;
return (
diff --git a/ui/components/models-page.tsx b/ui/components/models-page.tsx
index f8505d7f..bff3a379 100644
--- a/ui/components/models-page.tsx
+++ b/ui/components/models-page.tsx
@@ -124,12 +124,14 @@ export function ModelsPage() {
>
Basic Testing
+ {/*
API Endpoints
+ */}