From a713f819616c2deab7bd8fb9d5772c38e635858d Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Wed, 30 Sep 2026 02:15:57 +0200 Subject: [PATCH] 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, + }; +}