perf: cache provider catalogs and load admin models page progressively

This commit is contained in:
9qeklajc
2026-09-30 02:15:57 +02:00
parent 0dae9fe521
commit a713f81961
10 changed files with 446 additions and 78 deletions
+79 -13
View File
@@ -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 = [
+15
View File
@@ -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()
@@ -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]
+23
View File
@@ -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 (
<AppPageShell contentClassName='mx-auto w-full max-w-5xl'>
<div className='space-y-3 sm:space-y-4'>
<PageHeader
title='Model Management'
description='Manage provider model catalogs and validate endpoints from one place.'
/>
<Skeleton className='h-10 w-full' />
<Skeleton className='h-16 w-full' />
<Skeleton className='h-[420px] w-full' />
</div>
</AppPageShell>
);
}
+17 -2
View File
@@ -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 ? (
<div className='text-muted-foreground flex items-center justify-center gap-2 p-3 text-xs sm:text-sm'>
<Loader2 className='h-3.5 w-3.5 animate-spin' />
Rendering {hiddenCount} more model{hiddenCount === 1 ? '' : 's'}…
</div>
) : null;
if (filterProvider) {
return (
<div className='bg-card/35 border-border/70 md:divide-border/75 overflow-hidden rounded-lg border md:divide-y'>
{keyedProviderModels.map(({ model, renderKey }) => (
{visibleProviderModels.map(({ model, renderKey }) => (
<ModelItemCard
key={renderKey}
model={model}
@@ -125,6 +138,7 @@ export function ModelProviderSection({
onDelete={() => onDeleteModel(model.id)}
/>
))}
{pendingRowsNotice}
</div>
);
}
@@ -217,7 +231,7 @@ export function ModelProviderSection({
<CardContent className='px-3 pt-0 pb-3 sm:px-6 sm:pb-6'>
<div className='bg-card/35 border-border/70 md:divide-border/75 overflow-hidden rounded-lg border md:divide-y'>
{keyedProviderModels.map(({ model, renderKey }) => (
{visibleProviderModels.map(({ model, renderKey }) => (
<ModelItemCard
key={renderKey}
model={model}
@@ -236,6 +250,7 @@ export function ModelProviderSection({
onDelete={() => onDeleteModel(model.id)}
/>
))}
{pendingRowsNotice}
</div>
</CardContent>
</Card>
+30 -16
View File
@@ -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({
</div>
{Object.keys(groupedModels).length === 0 ? (
<div className='border-border/40 rounded-lg border border-dashed p-4 text-center sm:p-5'>
<p className='text-muted-foreground text-sm'>
Try broadening your search or switch to a different provider scope.
</p>
</div>
isFetchingRemote ? (
<div className='grid gap-3 sm:gap-4'>
<Skeleton className='h-[200px]' />
<Skeleton className='h-[200px]' />
</div>
) : (
<div className='border-border/40 rounded-lg border border-dashed p-4 text-center sm:p-5'>
<p className='text-muted-foreground text-sm'>
Try broadening your search or switch to a different provider
scope.
</p>
</div>
)
) : 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 ? (
<div className='border-border/40 text-muted-foreground flex items-center justify-center gap-2 rounded-lg border border-dashed p-3 text-sm sm:p-4'>
<Loader2 className='h-4 w-4 animate-spin' />
Loading provider catalogs…
</div>
) : null}
{/* Forms and Dialogs */}
{modelDialogState.providerId && (
<AddProviderModelDialog
+18 -12
View File
@@ -1,16 +1,14 @@
'use client';
import { useMemo, useState } from 'react';
import { useQuery } from '@tanstack/react-query';
import dynamic from 'next/dynamic';
import { AlertCircle } from 'lucide-react';
import type { Model } from '@/lib/api/schemas/models';
import { AdminService } from '@/lib/api/services/admin';
import { useModelsWithProviders } from '@/lib/hooks/use-models-with-providers';
import { groupAndSortModelsByProvider } from '@/lib/utils/model-sort';
import { AppPageShell } from '@/components/app-page-shell';
import { PageHeader } from '@/components/page-header';
import { ModelSelector } from '@/components/model-selector';
import { ModelTester } from '@/components/model-tester';
import { ApiEndpointTester } from '@/components/api-endpoint-tester';
import { ModelSearchFilter } from '@/components/model-search-filter';
import { Alert, AlertDescription } from '@/components/ui/alert';
import {
@@ -23,6 +21,19 @@ import {
import { Skeleton } from '@/components/ui/skeleton';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
// The testing tabs are never the landing view, so keeping them out of this
// route's chunk is what lets the navigation itself resolve quickly.
const ModelTester = dynamic(
() => import('@/components/model-tester').then((m) => m.ModelTester),
{ loading: () => <Skeleton className='h-[420px] w-full' />, ssr: false }
);
const ApiEndpointTester = dynamic(
() =>
import('@/components/api-endpoint-tester').then((m) => m.ApiEndpointTester),
{ loading: () => <Skeleton className='h-[420px] w-full' />, ssr: false }
);
export function ModelsPage() {
const [filteredModels, setFilteredModels] = useState<Model[] | undefined>(
undefined
@@ -31,16 +42,11 @@ export function ModelsPage() {
useState<string>('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),
+56 -35
View File
@@ -317,9 +317,14 @@ export class AdminService {
);
}
static async getProviderModels(providerId: number): Promise<ProviderModels> {
static async getProviderModels(
providerId: number,
options: { includeRemote?: boolean } = {}
): Promise<ProviderModels> {
const query =
options.includeRemote === false ? '?include_remote=false' : '';
const data = await apiClient.get<ProviderModels>(
`/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<string>();
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 };
+50
View File
@@ -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,
});
},
};
}
+46
View File
@@ -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<T>(
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,
};
}