From a713f819616c2deab7bd8fb9d5772c38e635858d Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 30 Sep 2026 02:15:57 +0200 Subject: [PATCH 1/2] perf: cache provider catalogs and load admin models page progressively --- routstr/core/admin.py | 92 ++++++++++++--- tests/conftest.py | 15 +++ tests/unit/test_admin_remote_models_cache.py | 112 +++++++++++++++++++ ui/app/model/loading.tsx | 23 ++++ ui/components/model-provider-section.tsx | 19 +++- ui/components/model-selector.tsx | 46 +++++--- ui/components/models-page.tsx | 30 +++-- ui/lib/api/services/admin.ts | 91 +++++++++------ ui/lib/hooks/use-models-with-providers.ts | 50 +++++++++ ui/lib/hooks/use-progressive-list.ts | 46 ++++++++ 10 files changed, 446 insertions(+), 78 deletions(-) create mode 100644 tests/unit/test_admin_remote_models_cache.py create mode 100644 ui/app/model/loading.tsx create mode 100644 ui/lib/hooks/use-models-with-providers.ts create mode 100644 ui/lib/hooks/use-progressive-list.ts diff --git a/routstr/core/admin.py b/routstr/core/admin.py index 365e2f68..f3862edc 100644 --- a/routstr/core/admin.py +++ b/routstr/core/admin.py @@ -1,6 +1,8 @@ +import asyncio import json import re import secrets +import time from datetime import datetime, timezone from pathlib import Path @@ -54,6 +56,9 @@ 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) @@ -1183,6 +1188,7 @@ 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} @@ -1196,13 +1202,78 @@ 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) -> dict[str, object]: - from ..upstream.helpers import _instantiate_provider - +async def get_provider_models( + provider_id: str, + include_remote: bool = Query(True), + refresh_remote: bool = Query(False), +) -> dict[str, object]: async with create_session() as session: provider = await _get_upstream_provider_by_ref(session, provider_id) provider_pk = _provider_pk(provider) @@ -1214,16 +1285,11 @@ async def get_provider_models(provider_id: str) -> dict[str, object]: apply_fees=False, ) - 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}" - ) + upstream_models: list = [] + if include_remote: + upstream_models = await _get_remote_models( + provider, provider_pk, force_refresh=refresh_remote + ) db_model_ids = {model.id for model in db_models} filtered_remote_models = [ diff --git a/tests/conftest.py b/tests/conftest.py index d1bfa919..74c3d058 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -31,3 +31,18 @@ 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 new file mode 100644 index 00000000..6a22ce00 --- /dev/null +++ b/tests/unit/test_admin_remote_models_cache.py @@ -0,0 +1,112 @@ +"""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/app/model/loading.tsx b/ui/app/model/loading.tsx new file mode 100644 index 00000000..390be4d8 --- /dev/null +++ b/ui/app/model/loading.tsx @@ -0,0 +1,23 @@ +import { AppPageShell } from '@/components/app-page-shell'; +import { PageHeader } from '@/components/page-header'; +import { Skeleton } from '@/components/ui/skeleton'; + +/** + * Route-level fallback so clicking "Models" lands on the page immediately + * instead of holding the previous route until this one's chunk is parsed. + */ +export default function ModelPageLoading() { + return ( + +
+ + + + +
+
+ ); +} diff --git a/ui/components/model-provider-section.tsx b/ui/components/model-provider-section.tsx index 9f441678..3acd132e 100644 --- a/ui/components/model-provider-section.tsx +++ b/ui/components/model-provider-section.tsx @@ -1,5 +1,6 @@ import { useMemo } from 'react'; import type { Model } from '@/lib/api/schemas/models'; +import { useProgressiveList } from '@/lib/hooks/use-progressive-list'; import type { AdminModelGroup } from '@/lib/api/services/admin'; import type { DisplayUnit } from '@/lib/types/units'; import { ModelItemCard } from '@/components/model-item-card'; @@ -24,6 +25,7 @@ import { Edit3, Globe, Key, + Loader2, MoreVertical, RefreshCw, } from 'lucide-react'; @@ -103,10 +105,21 @@ export function ModelProviderSection({ }); }, [provider, providerModels]); + const { visibleItems: visibleProviderModels, hiddenCount } = + useProgressiveList(keyedProviderModels); + + const pendingRowsNotice = + hiddenCount > 0 ? ( +
+ + Rendering {hiddenCount} more model{hiddenCount === 1 ? '' : 's'}… +
+ ) : null; + if (filterProvider) { return (
- {keyedProviderModels.map(({ model, renderKey }) => ( + {visibleProviderModels.map(({ model, renderKey }) => ( onDeleteModel(model.id)} /> ))} + {pendingRowsNotice}
); } @@ -217,7 +231,7 @@ export function ModelProviderSection({
- {keyedProviderModels.map(({ model, renderKey }) => ( + {visibleProviderModels.map(({ model, renderKey }) => ( onDeleteModel(model.id)} /> ))} + {pendingRowsNotice}
diff --git a/ui/components/model-selector.tsx b/ui/components/model-selector.tsx index 8636fb36..32552ef3 100644 --- a/ui/components/model-selector.tsx +++ b/ui/components/model-selector.tsx @@ -1,7 +1,7 @@ 'use client'; import React, { useState, useMemo } from 'react'; -import { useQuery, useMutation, useQueryClient } from '@tanstack/react-query'; +import { useMutation, useQueryClient } from '@tanstack/react-query'; import { type Model, type GroupSettings } from '@/lib/api/schemas/models'; import { AdminService, @@ -13,6 +13,7 @@ import { AddProviderModelDialog } from '@/components/add-provider-model-dialog'; import { EditGroupForm } from '@/components/edit-group-form'; import { ModelProviderSection } from '@/components/model-provider-section'; import { useDisplayCurrency } from '@/lib/hooks/use-display-currency'; +import { useModelsWithProviders } from '@/lib/hooks/use-models-with-providers'; import { Button } from '@/components/ui/button'; import { Checkbox } from '@/components/ui/checkbox'; import { Skeleton } from '@/components/ui/skeleton'; @@ -27,7 +28,7 @@ import { AlertDialogHeader, AlertDialogTitle, } from '@/components/ui/alert-dialog'; -import { Trash2, Ban, CheckCircle, Plus } from 'lucide-react'; +import { Trash2, Ban, CheckCircle, Loader2, Plus } from 'lucide-react'; import { toast } from 'sonner'; import { sortModels, @@ -131,19 +132,15 @@ export function ModelSelector({ const queryClient = useQueryClient(); - // Fetch models and groups + // Shared with the page shell, so mounting this panel costs no extra fetch. const { - data: modelsData, + models, + groups, isLoading: isLoadingModels, + isFetchingRemote, error: modelsError, refetch: refetchModels, - } = useQuery({ - queryKey: ['models-with-providers'], - queryFn: () => AdminService.getModelsWithProviders(), - refetchOnWindowFocus: false, - }); - - const { models = [], groups = [] } = modelsData || {}; + } = useModelsWithProviders(); const allOverrideModels = useMemo( () => models.filter(isOverrideModel), [models] @@ -868,11 +865,19 @@ export function ModelSelector({ {Object.keys(groupedModels).length === 0 ? ( -
-

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

-
+ isFetchingRemote ? ( +
+ + +
+ ) : ( +
+

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

+
+ ) ) : null} {/* Provider Groups or Filtered Models */} @@ -922,6 +927,15 @@ 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 && ( import('@/components/model-tester').then((m) => m.ModelTester), + { loading: () => , ssr: false } +); + +const ApiEndpointTester = dynamic( + () => + import('@/components/api-endpoint-tester').then((m) => m.ApiEndpointTester), + { loading: () => , ssr: false } +); + export function ModelsPage() { const [filteredModels, setFilteredModels] = useState( undefined @@ -31,16 +42,11 @@ export function ModelsPage() { useState('all'); const { - data: modelsData, + models, + groups, isLoading: isLoadingModels, error: modelsError, - } = useQuery({ - queryKey: ['admin-models-with-providers'], - queryFn: () => AdminService.getModelsWithProviders(), - refetchOnWindowFocus: false, - }); - - const { models = [], groups = [] } = modelsData || {}; + } = useModelsWithProviders(); const groupedModels = useMemo( () => groupAndSortModelsByProvider(models), diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index e510a571..78c8f01c 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -317,9 +317,14 @@ export class AdminService { ); } - static async getProviderModels(providerId: number): Promise { + static async getProviderModels( + providerId: number, + options: { includeRemote?: boolean } = {} + ): Promise { + const query = + options.includeRemote === false ? '?include_remote=false' : ''; const data = await apiClient.get( - `/admin/api/upstream-providers/${providerId}/models` + `/admin/api/upstream-providers/${providerId}/models${query}` ); // Convert pricing for all models in the list so the UI receives "per 1M tokens" values @@ -428,7 +433,9 @@ export class AdminService { ); } - static async getModelsWithProviders(): Promise<{ + static async getModelsWithProviders( + options: { includeRemote?: boolean } = {} + ): Promise<{ models: AdminModelAsModel[]; groups: AdminModelGroup[]; }> { @@ -446,14 +453,50 @@ export class AdminService { const allModels: AdminModelAsModel[] = []; const seenModelIds = new Set(); - for (const provider of providers) { - try { - const providerModels = await this.getProviderModels(provider.id); + // One provider's catalog never depends on another's, and each miss costs an + // upstream round trip, so the whole fan-out happens in a single wave. + const providerResults = await Promise.all( + providers.map(async (provider) => { + try { + return { + provider, + models: await this.getProviderModels(provider.id, options), + }; + } catch (error) { + console.error( + `Failed to fetch models for provider ${provider.id}:`, + error + ); + return null; + } + }) + ); - providerModels.db_models.forEach((dbModel) => { - seenModelIds.add(dbModel.id); + for (const result of providerResults) { + if (!result) { + continue; + } + const { provider, models: providerModels } = result; + providerModels.db_models.forEach((dbModel) => { + seenModelIds.add(dbModel.id); + const modelWithProvider = { + ...dbModel, + upstream_provider_id: provider.id, + }; + allModels.push({ + ...this.transformAdminModelToModel( + modelWithProvider, + provider.provider_type + ), + has_own_api_key: false, + api_key_type: 'group', + }); + }); + + providerModels.remote_models.forEach((remoteModel) => { + if (!seenModelIds.has(remoteModel.id)) { const modelWithProvider = { - ...dbModel, + ...remoteModel, upstream_provider_id: provider.id, }; allModels.push({ @@ -462,33 +505,11 @@ export class AdminService { provider.provider_type ), has_own_api_key: false, - api_key_type: 'group', + api_key_type: 'remote', + soft_deleted: false, }); - }); - - providerModels.remote_models.forEach((remoteModel) => { - if (!seenModelIds.has(remoteModel.id)) { - const modelWithProvider = { - ...remoteModel, - upstream_provider_id: provider.id, - }; - allModels.push({ - ...this.transformAdminModelToModel( - modelWithProvider, - provider.provider_type - ), - has_own_api_key: false, - api_key_type: 'remote', - soft_deleted: false, - }); - } - }); - } catch (error) { - console.error( - `Failed to fetch models for provider ${provider.id}:`, - error - ); - } + } + }); } return { models: allModels, groups }; diff --git a/ui/lib/hooks/use-models-with-providers.ts b/ui/lib/hooks/use-models-with-providers.ts new file mode 100644 index 00000000..90c2f1f5 --- /dev/null +++ b/ui/lib/hooks/use-models-with-providers.ts @@ -0,0 +1,50 @@ +'use client'; + +import { useQuery, useQueryClient } 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. + */ +export function useModelsWithProviders() { + const queryClient = useQueryClient(); + + const localQuery = useQuery({ + queryKey: localModelsQueryKey, + queryFn: () => + AdminService.getModelsWithProviders({ includeRemote: false }), + refetchOnWindowFocus: false, + staleTime: 30_000, + }); + + const fullQuery = 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, + }); + }, + }; +} diff --git a/ui/lib/hooks/use-progressive-list.ts b/ui/lib/hooks/use-progressive-list.ts new file mode 100644 index 00000000..8b55ffc4 --- /dev/null +++ b/ui/lib/hooks/use-progressive-list.ts @@ -0,0 +1,46 @@ +'use client'; + +import { useEffect, useState } from 'react'; + +/** + * Reveal a long list in frame-sized batches. + * + * A provider catalog can hold thousands of rows, and mounting them in one + * commit blocks the main thread long enough that the page looks frozen right + * after navigation. Each batch yields back to the browser, so the first rows + * paint immediately and the rest fill in without freezing input. + */ +export function useProgressiveList( + items: T[], + initialCount = 40, + step = 80 +): { visibleItems: T[]; hiddenCount: number } { + const [count, setCount] = useState(initialCount); + const [trackedItems, setTrackedItems] = useState(items); + + // Reset during render, not in an effect: an effect would first commit the new + // list at the old (possibly full) count, which is the freeze this avoids. + if (trackedItems !== items) { + setTrackedItems(items); + setCount(initialCount); + } + + useEffect(() => { + if (count >= items.length) { + return; + } + + const frame = requestAnimationFrame(() => { + setCount((current) => Math.min(items.length, current + step)); + }); + + return () => cancelAnimationFrame(frame); + }, [count, items.length, step]); + + const visibleCount = Math.min(count, items.length); + + return { + visibleItems: items.slice(0, visibleCount), + hiddenCount: items.length - visibleCount, + }; +} From 1c61dc1bd0af3cabcb5e93bf669ea65944f0be3a Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 30 Sep 2026 02:21:37 +0200 Subject: [PATCH 2/2] 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, }; }