From e903aa3a9f37210b36e6851654d6b72a94d9a926 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 2 Aug 2026 23:16:01 +0200 Subject: [PATCH] clean up --- routstr/upstream/model_paths.py | 34 ++++++++++++++++++++++----------- tests/unit/test_model_paths.py | 21 +++++++------------- 2 files changed, 30 insertions(+), 25 deletions(-) diff --git a/routstr/upstream/model_paths.py b/routstr/upstream/model_paths.py index 020f0fbb..95b6edf2 100644 --- a/routstr/upstream/model_paths.py +++ b/routstr/upstream/model_paths.py @@ -31,6 +31,8 @@ from ..core.db import ModelPathRow, ModelRow, UpstreamProviderRow, create_sessio from ..core.logging import get_logger if TYPE_CHECKING: + from sqlmodel.ext.asyncio.session import AsyncSession + from .base import BaseUpstreamProvider logger = get_logger(__name__) @@ -782,18 +784,28 @@ async def get_all_model_paths() -> dict: async def get_paths_for_model(model_id: str) -> dict: - """Return paths only for the exact model ID advertised by ``/v1/models``.""" - async with create_session() as session: - rows = ( - await session.exec( - select(ModelPathRow) - .where(col(ModelPathRow.model_id) == model_id) - .order_by( - col(ModelPathRow.path), - col(ModelPathRow.upstream_provider_id), + """Return paths for an advertised ID or its provider-prefixed alias.""" + + async def load_rows(session: AsyncSession, lookup_id: str) -> list[ModelPathRow]: + return list( + ( + await session.exec( + select(ModelPathRow) + .where(col(ModelPathRow.model_id) == lookup_id) + .order_by( + col(ModelPathRow.path), + col(ModelPathRow.upstream_provider_id), + ) ) - ) - ).all() + ).all() + ) + + async with create_session() as session: + rows = await load_rows(session, model_id) + if not rows: + unprefixed_id = public_model_id(model_id) + if unprefixed_id != model_id: + rows = await load_rows(session, unprefixed_id) seen: set[str] = set() paths: list[dict] = [] diff --git a/tests/unit/test_model_paths.py b/tests/unit/test_model_paths.py index 1bbe25e0..240a4199 100644 --- a/tests/unit/test_model_paths.py +++ b/tests/unit/test_model_paths.py @@ -462,9 +462,7 @@ async def test_disabling_model_on_one_provider_keeps_other_provider( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == { - _expected_path(1, "shared-model") - } + assert _paths_of(payload, "shared-model") == {_expected_path(1, "shared-model")} @pytest.mark.asyncio @@ -498,12 +496,8 @@ async def test_override_alias_not_applied_across_providers( await mp.refresh_model_paths([p1, p2]) payload = await mp.get_all_model_paths() - assert _paths_of(payload, "shared-model") == { - _expected_path(1, "shared-model") - } - assert _paths_of(payload, "private-alias") == { - _expected_path(2, "private-alias") - } + assert _paths_of(payload, "shared-model") == {_expected_path(1, "shared-model")} + assert _paths_of(payload, "private-alias") == {_expected_path(2, "private-alias")} @pytest.mark.asyncio @@ -1134,7 +1128,7 @@ async def test_get_paths_for_model_falls_back_to_provider_prefixed_id( @pytest.mark.asyncio -async def test_get_paths_for_model_requires_exact_advertised_id( +async def test_get_paths_for_model_accepts_provider_prefixed_alias( patched_session: AsyncEngine, ) -> None: p1 = _FakeProvider( @@ -1158,15 +1152,14 @@ async def test_get_paths_for_model_requires_exact_advertised_id( _path_entry(4, "deepseek-v4-pro"), _path_entry(7, "deepseek-v4-pro"), ] - assert prefixed_paths == [] + assert prefixed_paths == short_paths @pytest.mark.asyncio async def test_get_paths_for_model_multi_segment_id_matches_models_listing( patched_session: AsyncEngine, ) -> None: - """For three-segment ids the discovery id must be the same base id the - rest of the system exposes (first-slash rule), not the last segment.""" + """Three-segment upstream IDs resolve to the same first-slash public ID.""" provider = _FakeProvider( provider_type="generic", base_url="https://x/v1", @@ -1181,7 +1174,7 @@ async def test_get_paths_for_model_multi_segment_id_matches_models_listing( ] assert (await mp.get_paths_for_model("accounts/fireworks/models/glm-5"))[ "data" - ] == [] + ] == [_path_entry(1, "fireworks/models/glm-5")] # --------------------------------------------------------------------------- #