diff --git a/routstr/proxy.py b/routstr/proxy.py index f6977cf0..6a96e64c 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -104,6 +104,29 @@ def get_unique_models() -> list[Model]: return list(_unique_models.values()) +def _is_tinfoil_attestation_path(path: str) -> bool: + """Return True for Tinfoil attestation-bundle proxy paths.""" + return path in {"attestation", "tee/attestation"} + + +def _select_unauthenticated_get_upstreams( + path: str, upstreams: list[BaseUpstreamProvider] +) -> list[BaseUpstreamProvider]: + """Select upstream candidates for unauthenticated GET bypass paths. + + Tinfoil attestation endpoints are provider-specific. Trying every enabled + upstream can return an unrelated provider's 404 before Tinfoil is reached, + so route those paths only to Tinfoil providers. + """ + if _is_tinfoil_attestation_path(path): + return [ + upstream + for upstream in upstreams + if getattr(upstream, "provider_type", None) == "tinfoil" + ] + return upstreams + + async def refresh_model_maps() -> None: """Refresh global model and provider maps using the cost-based algorithm.""" from sqlalchemy.orm import selectinload @@ -211,20 +234,29 @@ async def proxy( model_id = request_body_dict.get("model", "unknown") # /tee/* and /attestation GET requests (e.g. Tinfoil attestation bundle) - # don't map to models — just forward to all enabled upstreams without - # model/cost/auth lookups. + # don't map to models — forward without model/cost/auth lookups. Tinfoil + # attestation paths are routed only to Tinfoil providers so an unrelated + # upstream's 404 cannot short-circuit before the attestation proxy is tried. if request.method == "GET" and ( path.startswith("tee/") or path.startswith("attestation") ): - all_upstreams = _upstreams + selected_upstreams = _select_unauthenticated_get_upstreams(path, _upstreams) + if not selected_upstreams: + return create_error_response( + "upstream_error", + "No upstream available for unauthenticated GET path", + 502, + request=request, + ) + last_error_response = None - for i, upstream in enumerate(all_upstreams): + for i, upstream in enumerate(selected_upstreams): try: headers = upstream.prepare_headers(dict(request.headers)) response = await upstream.forward_get_request(request, path, headers) - if response.status_code in [502, 429] and i < len(all_upstreams) - 1: + if response.status_code in [502, 429] and i < len(selected_upstreams) - 1: logger.warning( - "Upstream %s returned %s for tee GET %s, trying next", + "Upstream %s returned %s for unauthenticated GET %s, trying next", upstream.provider_type, response.status_code, path, @@ -233,12 +265,12 @@ async def proxy( return response except UpstreamError as e: logger.warning( - "Upstream %s failed for tee GET %s: %s", + "Upstream %s failed for unauthenticated GET %s: %s", upstream.provider_type, path, e, ) - if i == len(all_upstreams) - 1: + if i == len(selected_upstreams) - 1: last_error_response = create_error_response( "upstream_error", str(e), 502, request=request ) diff --git a/tests/unit/test_proxy_tinfoil_attestation_routing.py b/tests/unit/test_proxy_tinfoil_attestation_routing.py new file mode 100644 index 00000000..19c62c90 --- /dev/null +++ b/tests/unit/test_proxy_tinfoil_attestation_routing.py @@ -0,0 +1,94 @@ +from __future__ import annotations + +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import FastAPI +from fastapi.responses import Response +from httpx import ASGITransport, AsyncClient + +from routstr import proxy as proxy_module + + +@pytest.fixture +def proxy_app() -> FastAPI: + app = FastAPI() + app.include_router(proxy_module.proxy_router) + return app + + +@pytest.mark.asyncio +async def test_attestation_get_routes_directly_to_tinfoil_provider( + monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI +) -> None: + non_tinfoil = MagicMock() + non_tinfoil.provider_type = "openai" + non_tinfoil.prepare_headers = MagicMock(return_value={}) + non_tinfoil.forward_get_request = AsyncMock( + return_value=Response(status_code=404, content=b"wrong upstream") + ) + + tinfoil = MagicMock() + tinfoil.provider_type = "tinfoil" + tinfoil.prepare_headers = MagicMock(return_value={"accept": "application/json"}) + tinfoil.forward_get_request = AsyncMock( + return_value=Response(status_code=200, content=b'{"attestation":true}') + ) + + monkeypatch.setattr(proxy_module, "_upstreams", [non_tinfoil, tinfoil]) + + async with AsyncClient( + transport=ASGITransport(app=proxy_app), base_url="http://test" + ) as client: + response = await client.get("/attestation") + + assert response.status_code == 200 + assert response.content == b'{"attestation":true}' + non_tinfoil.forward_get_request.assert_not_called() + tinfoil.forward_get_request.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_tee_attestation_get_routes_directly_to_tinfoil_provider( + monkeypatch: pytest.MonkeyPatch, proxy_app: FastAPI +) -> None: + non_tinfoil = MagicMock() + non_tinfoil.provider_type = "openrouter" + non_tinfoil.prepare_headers = MagicMock(return_value={}) + non_tinfoil.forward_get_request = AsyncMock( + return_value=Response(status_code=404, content=b"wrong upstream") + ) + + tinfoil = MagicMock() + tinfoil.provider_type = "tinfoil" + tinfoil.prepare_headers = MagicMock(return_value={"accept": "application/json"}) + tinfoil.forward_get_request = AsyncMock( + return_value=Response(status_code=200, content=b'{"tee":true}') + ) + + monkeypatch.setattr(proxy_module, "_upstreams", [non_tinfoil, tinfoil]) + + async with AsyncClient( + transport=ASGITransport(app=proxy_app), base_url="http://test" + ) as client: + response = await client.get("/tee/attestation") + + assert response.status_code == 200 + assert response.content == b'{"tee":true}' + non_tinfoil.forward_get_request.assert_not_called() + tinfoil.forward_get_request.assert_awaited_once() + + +def test_attestation_upstream_selection_is_tinfoil_only() -> None: + non_tinfoil = MagicMock(provider_type="openai") + tinfoil = MagicMock(provider_type="tinfoil") + + assert proxy_module._select_unauthenticated_get_upstreams( + "attestation", [non_tinfoil, tinfoil] + ) == [tinfoil] + assert proxy_module._select_unauthenticated_get_upstreams( + "tee/attestation", [non_tinfoil, tinfoil] + ) == [tinfoil] + assert proxy_module._select_unauthenticated_get_upstreams( + "tee/other", [non_tinfoil, tinfoil] + ) == [non_tinfoil, tinfoil]