diff --git a/tests/integration/test_admin_forwarded_model_id.py b/tests/integration/test_admin_forwarded_model_id.py new file mode 100644 index 00000000..609dc18a --- /dev/null +++ b/tests/integration/test_admin_forwarded_model_id.py @@ -0,0 +1,143 @@ +"""Regression tests for admin model alias persistence (GitHub issue #639).""" + +import json +from datetime import datetime, timedelta, timezone +from unittest.mock import patch + +import pytest +from httpx import AsyncClient +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.admin import admin_sessions +from routstr.core.db import ModelRow, UpstreamProviderRow +from routstr.proxy import reinitialize_upstreams + +MODEL_ID = "deepseek/deepseek-chat" +ARCHITECTURE = { + "modality": "text", + "input_modalities": ["text"], + "output_modalities": ["text"], + "tokenizer": "unknown", + "instruct_type": None, +} +PRICING = { + "prompt": 1e-7, + "completion": 2e-7, + "request": 0.0, + "image": 0.0, + "web_search": 0.0, + "internal_reasoning": 0.0, +} + + +def _admin_headers() -> dict[str, str]: + token = "test-admin-forwarded-model-id-token" + admin_sessions[token] = int( + (datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp() + ) + return {"Authorization": f"Bearer {token}"} + + +def _model_payload( + *, forwarded_model_id: str | None, include_alias: bool +) -> dict[str, object]: + payload: dict[str, object] = { + "id": MODEL_ID, + "name": "DeepSeek Chat", + "description": "DeepSeek Chat test model", + "created": 0, + "context_length": 128000, + "architecture": ARCHITECTURE, + "pricing": PRICING, + "per_request_limits": None, + "top_provider": None, + "canonical_slug": None, + "alias_ids": [], + "enabled": True, + } + if include_alias: + payload["forwarded_model_id"] = forwarded_model_id + return payload + + +def _model_row(provider_id: int, *, forwarded_model_id: str | None) -> ModelRow: + return ModelRow( + id=MODEL_ID, + upstream_provider_id=provider_id, + name="DeepSeek Chat", + description="DeepSeek Chat test model", + created=0, + context_length=128000, + architecture=json.dumps(ARCHITECTURE), + pricing=json.dumps(PRICING), + enabled=True, + forwarded_model_id=forwarded_model_id, + ) + + +async def _create_provider( + session: AsyncSession, *, base_url: str +) -> UpstreamProviderRow: + provider = UpstreamProviderRow( + provider_type="generic", + base_url=base_url, + api_key="test-key", + provider_fee=1.0, + ) + session.add(provider) + await session.commit() + await session.refresh(provider) + assert provider.id is not None + return provider + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_create_without_forwarded_model_id_preserves_null( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """Saving an unaliased model must not invent a self-alias.""" + provider = await _create_provider( + integration_session, base_url="https://issue-639-create.example/v1" + ) + await reinitialize_upstreams() + + with patch("routstr.payment.models.sats_usd_price", return_value=1e-6): + response = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/models", + headers=_admin_headers(), + json=_model_payload(forwarded_model_id=None, include_alias=False), + ) + + assert response.status_code == 200 + row = await integration_session.get(ModelRow, (MODEL_ID, provider.id)) + assert row is not None + assert row.forwarded_model_id is None + + +@pytest.mark.integration +@pytest.mark.asyncio +async def test_admin_update_can_clear_forwarded_model_id( + integration_client: AsyncClient, + integration_session: AsyncSession, +) -> None: + """An explicit JSON null must remove an existing upstream alias.""" + provider = await _create_provider( + integration_session, base_url="https://issue-639-clear.example/v1" + ) + row = _model_row(provider.id, forwarded_model_id="upstream-deepseek-chat") + integration_session.add(row) + await integration_session.commit() + await reinitialize_upstreams() + + with patch("routstr.payment.models.sats_usd_price", return_value=1e-6): + response = await integration_client.post( + f"/admin/api/upstream-providers/{provider.id}/models", + headers=_admin_headers(), + json=_model_payload(forwarded_model_id=None, include_alias=True), + ) + + assert response.status_code == 200 + await integration_session.refresh(row) + assert row.forwarded_model_id is None diff --git a/tests/unit/test_algorithm.py b/tests/unit/test_algorithm.py index 21b418db..92452cc9 100644 --- a/tests/unit/test_algorithm.py +++ b/tests/unit/test_algorithm.py @@ -230,6 +230,45 @@ def test_create_model_mappings_applies_override_only_to_matching_provider( assert {p for _, p in provider_map["same-id"]} == {provider_a, provider_b} +def test_create_model_mappings_does_not_split_self_alias_from_base_identity( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A write-time self-alias must not advertise one shared model twice (#639).""" + model_id = "deepseek/deepseek-chat" + provider_a = create_test_provider( + "provider-a", + "https://provider-a.example/v1", + db_id=1, + models=[create_test_model(model_id)], + ) + provider_b = create_test_provider( + "provider-b", + "https://provider-b.example/v1", + db_id=2, + models=[create_test_model(model_id)], + ) + + admin_saved_model = create_test_model(model_id) + admin_saved_model.forwarded_model_id = model_id + override_row = SimpleNamespace(id=model_id, upstream_provider_id=1, enabled=True) + + def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def] + return admin_saved_model + + monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model) + + _, _, unique_models = create_model_mappings( + upstreams=[provider_a, provider_b], + overrides_by_key={(model_id, 1): (override_row, 1.0)}, + disabled_model_keys=set(), + ) + + advertised_ids = sorted( + model.forwarded_model_id or model.id for model in unique_models.values() + ) + assert advertised_ids == ["deepseek-chat"] + + def test_create_model_mappings_disables_only_matching_provider() -> None: """Disabled overrides are scoped to the provider row, not the shared model id.""" provider_a = create_test_provider(