From 1c61dc1bd0af3cabcb5e93bf669ea65944f0be3a Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 30 Sep 2026 02:21:37 +0200 Subject: [PATCH] refactor: drop provider catalog cache, fetch on page mount with skeleton --- routstr/core/admin.py | 92 +++------------ tests/conftest.py | 15 --- tests/unit/test_admin_remote_models_cache.py | 112 ------------------- ui/components/model-selector.tsx | 30 +---- ui/lib/api/services/admin.ts | 15 +-- ui/lib/hooks/use-models-with-providers.ts | 42 ++----- 6 files changed, 32 insertions(+), 274 deletions(-) delete mode 100644 tests/unit/test_admin_remote_models_cache.py diff --git a/routstr/core/admin.py b/routstr/core/admin.py index f3862edc..365e2f68 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1,8 +1,6 @@ -import asyncio import json import re import secrets -import time from datetime import datetime, timezone from pathlib import Path @@ -56,9 +54,6 @@ async def _refresh_provider_model_paths(upstream_provider_id: int) -> None: """Queue discovery sync without blocking the committed admin mutation.""" from ..upstream.model_paths import schedule_model_paths_refresh_for_provider - # Every provider/model mutation funnels through here, so it is also the one - # place that can keep the cached admin catalog from serving a stale listing. - invalidate_remote_models_cache(upstream_provider_id) await schedule_model_paths_refresh_for_provider(upstream_provider_id) @@ -1188,7 +1183,6 @@ async def delete_upstream_provider(provider_id: str) -> dict[str, object]: await session.delete(provider) await session.commit() - invalidate_remote_models_cache(deleted_id) await reinitialize_upstreams() await refresh_model_maps() return {"ok": True, "deleted_id": deleted_id} @@ -1202,78 +1196,13 @@ async def get_provider_types() -> list[dict[str, object]]: return [cls.get_provider_metadata() for cls in upstream_provider_classes] -# The admin catalog view is opened repeatedly and by several panels at once, -# while every miss costs a live upstream round trip. Keep the raw listing for a -# short window and let concurrent readers share one in-flight fetch. -_REMOTE_MODELS_TTL_SECONDS = 120.0 -_REMOTE_MODELS_FETCH_TIMEOUT_SECONDS = 20.0 -_remote_models_cache: dict[int, tuple[float, list]] = {} -_remote_models_locks: dict[int, asyncio.Lock] = {} -# Bumped on every invalidation so a fetch that started against the old provider -# config cannot write its result back after the cache was cleared. -_remote_models_generation = 0 - - -def invalidate_remote_models_cache(provider_pk: int | None = None) -> None: - global _remote_models_generation - _remote_models_generation += 1 - if provider_pk is None: - _remote_models_cache.clear() - _remote_models_locks.clear() - else: - _remote_models_cache.pop(provider_pk, None) - - -async def _get_remote_models( - provider: UpstreamProviderRow, provider_pk: int, force_refresh: bool = False -) -> list: - from ..upstream.helpers import _instantiate_provider - - now = time.monotonic() - cached = _remote_models_cache.get(provider_pk) - if not force_refresh and cached and now - cached[0] < _REMOTE_MODELS_TTL_SECONDS: - return cached[1] - - lock = _remote_models_locks.setdefault(provider_pk, asyncio.Lock()) - async with lock: - cached = _remote_models_cache.get(provider_pk) - now = time.monotonic() - if ( - not force_refresh - and cached - and now - cached[0] < _REMOTE_MODELS_TTL_SECONDS - ): - return cached[1] - - upstream_instance = _instantiate_provider(provider) - if not upstream_instance: - return [] - - generation = _remote_models_generation - try: - models = await asyncio.wait_for( - upstream_instance.fetch_models(), - timeout=_REMOTE_MODELS_FETCH_TIMEOUT_SECONDS, - ) - except Exception as e: - logger.error(f"Failed to fetch models from {provider.provider_type}: {e}") - # A stale listing beats an empty one for an operator view. - return cached[1] if cached else [] - - if generation == _remote_models_generation: - _remote_models_cache[provider_pk] = (time.monotonic(), models) - return models - - @admin_router.get( "/api/upstream-providers/{provider_id}/models", dependencies=[Depends(require_admin_api)], ) -async def get_provider_models( - provider_id: str, - include_remote: bool = Query(True), - refresh_remote: bool = Query(False), -) -> dict[str, object]: +async def get_provider_models(provider_id: str) -> dict[str, object]: + from ..upstream.helpers import _instantiate_provider + async with create_session() as session: provider = await _get_upstream_provider_by_ref(session, provider_id) provider_pk = _provider_pk(provider) @@ -1285,11 +1214,16 @@ async def get_provider_models( apply_fees=False, ) - upstream_models: list = [] - if include_remote: - upstream_models = await _get_remote_models( - provider, provider_pk, force_refresh=refresh_remote - ) + upstream_models = [] + upstream_instance = _instantiate_provider(provider) + if upstream_instance: + try: + raw_models = await upstream_instance.fetch_models() + upstream_models = raw_models + except Exception as e: + logger.error( + f"Failed to fetch models from {provider.provider_type}: {e}" + ) db_model_ids = {model.id for model in db_models} filtered_remote_models = [ diff --git a/tests/conftest.py b/tests/conftest.py index 74c3d058..d1bfa919 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -31,18 +31,3 @@ def _isolate_redemption_negative_cache() -> Iterator[None]: redemption_negative_cache.clear() yield redemption_negative_cache.clear() - - -@pytest.fixture(autouse=True) -def _isolate_admin_remote_models_cache() -> Iterator[None]: - """Clear the admin catalog cache between tests. - - Provider primary keys restart at 1 for every fresh test database, so a - cached listing from an earlier test would otherwise answer for a different - provider that happens to reuse the same key. - """ - from routstr.core.admin import invalidate_remote_models_cache - - invalidate_remote_models_cache() - yield - invalidate_remote_models_cache() diff --git a/tests/unit/test_admin_remote_models_cache.py b/tests/unit/test_admin_remote_models_cache.py deleted file mode 100644 index 6a22ce00..00000000 --- a/tests/unit/test_admin_remote_models_cache.py +++ /dev/null @@ -1,112 +0,0 @@ -"""Cache behavior of the admin provider catalog listing.""" - -import asyncio -from types import SimpleNamespace -from typing import Any -from unittest.mock import AsyncMock, patch - -import pytest - -from routstr.core import admin -from routstr.core.admin import _get_remote_models, invalidate_remote_models_cache - -PROVIDER = SimpleNamespace(provider_type="generic") - - -def _upstream(fetch: Any) -> Any: - return patch( - "routstr.upstream.helpers._instantiate_provider", - return_value=SimpleNamespace(fetch_models=fetch), - ) - - -@pytest.mark.asyncio -async def test_second_read_is_served_from_cache() -> None: - fetch = AsyncMock(return_value=["a"]) - with _upstream(fetch): - assert await _get_remote_models(PROVIDER, 1) == ["a"] # type: ignore[arg-type] - assert await _get_remote_models(PROVIDER, 1) == ["a"] # type: ignore[arg-type] - assert fetch.await_count == 1 - - -@pytest.mark.asyncio -async def test_force_refresh_and_invalidation_refetch() -> None: - fetch = AsyncMock(side_effect=[["a"], ["b"], ["c"]]) - with _upstream(fetch): - await _get_remote_models(PROVIDER, 1) # type: ignore[arg-type] - assert await _get_remote_models(PROVIDER, 1, force_refresh=True) == ["b"] # type: ignore[arg-type] - invalidate_remote_models_cache(1) - assert await _get_remote_models(PROVIDER, 1) == ["c"] # type: ignore[arg-type] - - -@pytest.mark.asyncio -async def test_expired_entry_is_refetched() -> None: - fetch = AsyncMock(side_effect=[["a"], ["b"]]) - with _upstream(fetch): - await _get_remote_models(PROVIDER, 1) # type: ignore[arg-type] - stamp, models = admin._remote_models_cache[1] - admin._remote_models_cache[1] = ( - stamp - admin._REMOTE_MODELS_TTL_SECONDS - 1, - models, - ) - assert await _get_remote_models(PROVIDER, 1) == ["b"] # type: ignore[arg-type] - - -@pytest.mark.asyncio -async def test_failed_refresh_falls_back_to_stale_listing() -> None: - fetch = AsyncMock(side_effect=[["a"], RuntimeError("upstream down")]) - with _upstream(fetch): - await _get_remote_models(PROVIDER, 1) # type: ignore[arg-type] - assert await _get_remote_models(PROVIDER, 1, force_refresh=True) == ["a"] # type: ignore[arg-type] - - -@pytest.mark.asyncio -async def test_failed_first_fetch_returns_empty_and_caches_nothing() -> None: - fetch = AsyncMock(side_effect=RuntimeError("upstream down")) - with _upstream(fetch): - assert await _get_remote_models(PROVIDER, 1) == [] # type: ignore[arg-type] - assert 1 not in admin._remote_models_cache - - -@pytest.mark.asyncio -async def test_concurrent_readers_share_one_fetch() -> None: - release = asyncio.Event() - - async def slow_fetch() -> list[str]: - await release.wait() - return ["a"] - - fetch = AsyncMock(side_effect=slow_fetch) - with _upstream(fetch): - readers = [ - asyncio.create_task(_get_remote_models(PROVIDER, 1)) # type: ignore[arg-type] - for _ in range(5) - ] - await asyncio.sleep(0) - release.set() - results = await asyncio.gather(*readers) - assert results == [["a"]] * 5 - assert fetch.await_count == 1 - - -@pytest.mark.asyncio -async def test_invalidation_during_fetch_discards_in_flight_result() -> None: - release = asyncio.Event() - - async def slow_fetch() -> list[str]: - await release.wait() - return ["old"] - - with _upstream(AsyncMock(side_effect=slow_fetch)): - reader = asyncio.create_task(_get_remote_models(PROVIDER, 1)) # type: ignore[arg-type] - await asyncio.sleep(0) - invalidate_remote_models_cache(1) - release.set() - assert await reader == ["old"] - assert 1 not in admin._remote_models_cache - - -@pytest.mark.asyncio -async def test_uninstantiable_provider_returns_empty() -> None: - with patch("routstr.upstream.helpers._instantiate_provider", return_value=None): - assert await _get_remote_models(PROVIDER, 1) == [] # type: ignore[arg-type] diff --git a/ui/components/model-selector.tsx b/ui/components/model-selector.tsx index 32552ef3..54a345b8 100644 --- a/ui/components/model-selector.tsx +++ b/ui/components/model-selector.tsx @@ -28,7 +28,7 @@ import { AlertDialogHeader, AlertDialogTitle, } from '@/components/ui/alert-dialog'; -import { Trash2, Ban, CheckCircle, Loader2, Plus } from 'lucide-react'; +import { Trash2, Ban, CheckCircle, Plus } from 'lucide-react'; import { toast } from 'sonner'; import { sortModels, @@ -137,7 +137,6 @@ export function ModelSelector({ models, groups, isLoading: isLoadingModels, - isFetchingRemote, error: modelsError, refetch: refetchModels, } = useModelsWithProviders(); @@ -865,19 +864,11 @@ export function ModelSelector({ {Object.keys(groupedModels).length === 0 ? ( - isFetchingRemote ? ( -
- - -
- ) : ( -
-

- Try broadening your search or switch to a different provider - scope. -

-
- ) +
+

+ Try broadening your search or switch to a different provider scope. +

+
) : null} {/* Provider Groups or Filtered Models */} @@ -927,15 +918,6 @@ export function ModelSelector({ ); })} - {/* The stored rows render first; provider catalogs arrive after their - upstream calls return, so the list says more is still on the way. */} - {isFetchingRemote && Object.keys(groupedModels).length > 0 ? ( -
- - Loading provider catalogs… -
- ) : null} - {/* Forms and Dialogs */} {modelDialogState.providerId && ( { - const query = - options.includeRemote === false ? '?include_remote=false' : ''; + static async getProviderModels(providerId: number): Promise { const data = await apiClient.get( - `/admin/api/upstream-providers/${providerId}/models${query}` + `/admin/api/upstream-providers/${providerId}/models` ); // Convert pricing for all models in the list so the UI receives "per 1M tokens" values @@ -433,9 +428,7 @@ export class AdminService { ); } - static async getModelsWithProviders( - options: { includeRemote?: boolean } = {} - ): Promise<{ + static async getModelsWithProviders(): Promise<{ models: AdminModelAsModel[]; groups: AdminModelGroup[]; }> { @@ -460,7 +453,7 @@ export class AdminService { try { return { provider, - models: await this.getProviderModels(provider.id, options), + models: await this.getProviderModels(provider.id), }; } catch (error) { console.error( diff --git a/ui/lib/hooks/use-models-with-providers.ts b/ui/lib/hooks/use-models-with-providers.ts index 90c2f1f5..6aeec5cf 100644 --- a/ui/lib/hooks/use-models-with-providers.ts +++ b/ui/lib/hooks/use-models-with-providers.ts @@ -1,50 +1,26 @@ 'use client'; -import { useQuery, useQueryClient } from '@tanstack/react-query'; +import { useQuery } from '@tanstack/react-query'; import { AdminService } from '@/lib/api/services/admin'; export const modelsWithProvidersQueryKey = ['models-with-providers'] as const; -const localModelsQueryKey = ['models-with-providers', 'local'] as const; /** - * Shared catalog read for every models view. - * - * The database rows come back without touching an upstream, so they render - * first; the listing that needs live provider calls replaces them once it - * lands. Both queries live under one key prefix, so a single - * `invalidateQueries(['models-with-providers'])` still refreshes the pair, and - * every consumer of this hook shares one request instead of fanning out again. + * Shared catalog read for every models view, so the page shell and the + * selector panel share one request instead of each fanning out to providers. */ export function useModelsWithProviders() { - const queryClient = useQueryClient(); - - const localQuery = useQuery({ - queryKey: localModelsQueryKey, - queryFn: () => - AdminService.getModelsWithProviders({ includeRemote: false }), - refetchOnWindowFocus: false, - staleTime: 30_000, - }); - - const fullQuery = useQuery({ + const query = useQuery({ queryKey: modelsWithProvidersQueryKey, queryFn: () => AdminService.getModelsWithProviders(), refetchOnWindowFocus: false, - staleTime: 60_000, }); - const data = fullQuery.data ?? localQuery.data; - return { - models: data?.models ?? [], - groups: data?.groups ?? [], - isLoading: !data && (localQuery.isLoading || fullQuery.isLoading), - isFetchingRemote: fullQuery.isFetching, - error: data ? null : (fullQuery.error ?? localQuery.error), - refetch: async () => { - await queryClient.invalidateQueries({ - queryKey: modelsWithProvidersQueryKey, - }); - }, + models: query.data?.models ?? [], + groups: query.data?.groups ?? [], + isLoading: query.isLoading, + error: query.error, + refetch: query.refetch, }; }