diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 957684af..8de51de7 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1733,8 +1733,9 @@ async def certify_upstream_provider( provider.provider_fee, sats_to_usd, ) - # The timeout applies per upstream call. The run makes up to five - # calls, so the request can stay open for up to five times it. + # The timeout applies per upstream call. The run makes up to six + # calls (models, two short probes after a max_completion_tokens retry, + # three cache probes), so the request can stay open for six times it. live_rows = await run_live_checks( provider.base_url, provider.api_key, diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 731c6de5..0120e26e 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -689,7 +689,8 @@ def _model_test_target( upstream.build_request_url(path, model_obj), upstream.prepare_headers({"content-type": "application/json"}), dict(upstream.prepare_params(path, None)), - upstream.transform_model_name(model_id), + # The proxy forwards ``model.id``, not the row's client alias. + upstream.transform_model_name(model_obj.id), ) diff --git a/routstr/upstream/certification.py b/routstr/upstream/certification.py index cb4d83e4..fc345f33 100644 --- a/routstr/upstream/certification.py +++ b/routstr/upstream/certification.py @@ -247,17 +247,14 @@ def shape_body( ) -> Any: """The JSON body the proxy would forward, model-name transforms included. - ``prepare_request_body`` rewrites ``model`` from ``model.id``; the probe - keeps the id it chose (``forwarded_model_id`` first) and only applies the - provider's own name transform to it. + ``prepare_request_body`` sets ``model`` from ``model.id``, exactly as + ``forward_request`` does, so an alias row's ``forwarded_model_id`` never + reaches the upstream here either. """ if upstream is None or model is None: return body shaped = upstream.prepare_request_body(json.dumps(body).encode(), model) - data = json.loads(shaped) if shaped else dict(body) - if isinstance(data, dict) and isinstance(body.get("model"), str): - data["model"] = upstream.transform_model_name(body["model"]) - return data + return json.loads(shaped) if shaped else body async def probe_upstream( @@ -876,7 +873,7 @@ async def run_live_checks( probe = await probe_upstream( base_url, api_key, - model.forwarded_model_id or model.id, + model.id, endpoint_tag=endpoint_tag, client=client, timeout=timeout, diff --git a/routstr/upstream/certification_cache.py b/routstr/upstream/certification_cache.py index a7459683..fef1d8a4 100644 --- a/routstr/upstream/certification_cache.py +++ b/routstr/upstream/certification_cache.py @@ -604,7 +604,7 @@ async def run_cache_checks( probe = await probe_cache( base_url, api_key, - model.forwarded_model_id or model.id, + model.id, endpoint_tag=endpoint_tag, client=client, timeout=timeout, diff --git a/tests/integration/test_certify_alias_paths.py b/tests/integration/test_certify_alias_paths.py index 6786a8a6..98b35bfd 100644 --- a/tests/integration/test_certify_alias_paths.py +++ b/tests/integration/test_certify_alias_paths.py @@ -106,7 +106,9 @@ async def test_certify_forwarded_alias_listed_path_succeeds( assert response.status_code == 200, response.text assert chat.call_count == 1 body: dict[str, Any] = json.loads(chat.calls[0].request.content) - assert body["model"] == forwarded + # The path is keyed by the exposed id, but the upstream gets what the + # proxy sends for this row: transform_model_name(model.id). + assert body["model"] == "local-alias" assert body["provider"] == {"order": ["endpoint"], "allow_fallbacks": False} mismatch = await integration_client.post( f"/admin/api/upstream-providers/{provider.id}/certify", diff --git a/tests/integration/test_certify_matches_proxy_model.py b/tests/integration/test_certify_matches_proxy_model.py new file mode 100644 index 00000000..d36357bc --- /dev/null +++ b/tests/integration/test_certify_matches_proxy_model.py @@ -0,0 +1,90 @@ +"""Certification must send the upstream the model id the proxy sends. + +An admin alias row has ``id`` and ``forwarded_model_id`` that differ. The +proxy forwards ``transform_model_name(model.id)`` (``prepare_request_body``), +so a probe that sends the forwarded id certifies a request no client can make. +""" + +import json +import time +from typing import Any +from unittest.mock import patch + +import pytest +import respx +from httpx import AsyncClient, Response +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.db import ApiKey +from routstr.proxy import reinitialize_upstreams + +from .test_certify_endpoint import _admin_headers, _make_provider, _model_row + + +@pytest.mark.integration +@pytest.mark.asyncio +@respx.mock +async def test_certify_and_proxy_send_the_same_model( + integration_session: AsyncSession, + integration_client: AsyncClient, +) -> None: + base_url = "https://certify-upstream.example/v1" + respx.get(f"{base_url}/models").mock(return_value=Response(200, json={"data": []})) + chat = respx.post(f"{base_url}/chat/completions").mock( + return_value=Response( + 200, + json={ + "id": "x", + "object": "chat.completion", + "model": "m", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 1}, + }, + ) + ) + provider = await _make_provider(integration_session) + model = _model_row(provider.id, model_id="row-id") # type: ignore[arg-type] + model.forwarded_model_id = "client-alias" + integration_session.add(model) + integration_session.add( + ApiKey( + hashed_key="certify-contract", balance=10**9, created_at=int(time.time()) + ) + ) + await integration_session.commit() + + with ( + patch("routstr.payment.models.sats_usd_price", return_value=0.0005), + patch("routstr.payment.cost_calculation.sats_usd_price", return_value=0.0005), + patch("routstr.payment.price.SATS_USD_PRICE", 0.0005), + ): + await reinitialize_upstreams() + proxied = await integration_client.post( + "/v1/chat/completions", + headers={"Authorization": "Bearer sk-certify-contract"}, + json={ + "model": "client-alias", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 1, + }, + ) + assert proxied.status_code == 200, proxied.text + certified = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/certify", + headers=_admin_headers(), + json={"model_id": "row-id", "check_cache": False}, + ) + assert certified.status_code == 200, certified.text + + bodies: list[dict[str, Any]] = [ + json.loads(call.request.content) for call in chat.calls + ] + assert len(bodies) == 2 + proxy_model, certify_model = bodies[0]["model"], bodies[1]["model"] + assert certify_model == proxy_model