From 1b6188b1308968d460afa31f0ae99483010fd55b Mon Sep 17 00:00:00 2001 From: 9qeklajc <9qeklajc> Date: Sun, 2 Nov 2025 23:42:18 +0100 Subject: [PATCH] better model update --- routstr/core/main.py | 7 ++ routstr/payment/models.py | 110 ++++++++++++++++++++++++++++++++ ui/components/EditModelForm.tsx | 2 +- ui/components/ModelSelector.tsx | 28 ++++---- ui/lib/api/services/admin.ts | 2 +- 5 files changed, 131 insertions(+), 18 deletions(-) diff --git a/routstr/core/main.py b/routstr/core/main.py index 5a4df7e5..22bcb5ce 100644 --- a/routstr/core/main.py +++ b/routstr/core/main.py @@ -14,6 +14,7 @@ from ..balance import balance_router, deprecated_wallet_router from ..discovery import providers_cache_refresher, providers_router from ..nip91 import announce_provider from ..payment.models import ( + cleanup_enabled_models_periodically, models_router, update_sats_pricing, ) @@ -48,6 +49,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: nip91_task = None providers_task = None models_refresh_task = None + models_cleanup_task = None model_maps_refresh_task = None try: @@ -87,6 +89,7 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: models_refresh_task = asyncio.create_task( refresh_upstreams_models_periodically(get_upstreams()) ) + models_cleanup_task = asyncio.create_task(cleanup_enabled_models_periodically()) model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically()) payout_task = asyncio.create_task(periodic_payout()) nip91_task = asyncio.create_task(announce_provider()) @@ -115,6 +118,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: providers_task.cancel() if models_refresh_task is not None: models_refresh_task.cancel() + if models_cleanup_task is not None: + models_cleanup_task.cancel() if model_maps_refresh_task is not None: model_maps_refresh_task.cancel() @@ -132,6 +137,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: tasks_to_wait.append(providers_task) if models_refresh_task is not None: tasks_to_wait.append(models_refresh_task) + if models_cleanup_task is not None: + tasks_to_wait.append(models_cleanup_task) if model_maps_refresh_task is not None: tasks_to_wait.append(model_maps_refresh_task) diff --git a/routstr/payment/models.py b/routstr/payment/models.py index 90d56581..61955035 100644 --- a/routstr/payment/models.py +++ b/routstr/payment/models.py @@ -526,6 +526,116 @@ async def update_sats_pricing() -> None: logger.error(f"Error updating sats pricing: {e}") +async def cleanup_enabled_models_periodically() -> None: + """Background task to clean up enabled models that match upstream pricing. + + When model is enabled (enabled=True), remove it from DB if it matches upstream pricing. + Keep it in DB only if pricing differs from upstream or if it's disabled. + """ + interval = getattr( + settings, "models_cleanup_interval_seconds", 300 + ) # 5 minutes default + if not interval or interval <= 0: + return + + while True: + try: + await _cleanup_enabled_models_once() + except asyncio.CancelledError: + break + except Exception as e: + logger.error( + "Error during enabled models cleanup", + extra={"error": str(e), "error_type": type(e).__name__}, + ) + + try: + jitter = max(0.0, float(interval) * 0.1) + await asyncio.sleep(interval + random.uniform(0, jitter)) + except asyncio.CancelledError: + break + + +async def _cleanup_enabled_models_once() -> None: + """Clean up enabled models that match upstream pricing.""" + from ..proxy import get_upstreams + + async with create_session() as session: + # Get all enabled models from DB + result = await session.exec( + select(ModelRow).where( + ModelRow.enabled, # Only enabled models + ) + ) + db_models = result.all() + + if not db_models: + return + + upstreams = get_upstreams() + models_to_remove = [] + + for db_model in db_models: + # Find corresponding upstream model + print(db_model.id) + upstream_model = None + for upstream in upstreams: + upstream_model = upstream.get_cached_model_by_id(db_model.id) + if upstream_model: + break + + if not upstream_model: + continue + + # Compare pricing to see if they match + db_pricing = json.loads(db_model.pricing) + upstream_pricing = upstream_model.pricing.dict() + + # Check if pricing matches (with small tolerance for float comparison) + pricing_matches = _pricing_matches(db_pricing, upstream_pricing) + + if pricing_matches: + models_to_remove.append(db_model) + logger.info( + f"Removing enabled model {db_model.id} - matches upstream pricing", + extra={"model_id": db_model.id}, + ) + + # Remove models that match upstream pricing + for model in models_to_remove: + await session.delete(model) + + if models_to_remove: + await session.commit() + logger.info( + f"Cleaned up {len(models_to_remove)} enabled models that match upstream pricing" + ) + + +def _pricing_matches( + db_pricing: dict, upstream_pricing: dict, tolerance: float = 0.1 +) -> bool: + """Check if pricing dictionaries match within tolerance.""" + keys_to_compare = [ + "prompt", + "completion", + "request", + "image", + "web_search", + "internal_reasoning", + ] + + for key in keys_to_compare: + db_val = float(db_pricing.get(key, 0.0)) * 1000000 + upstream_val = float(upstream_pricing.get(key, 0.0)) * 1000000 + print(db_val - upstream_val) + + if abs(db_val - upstream_val) > tolerance: + return False + + return True + + async def refresh_models_periodically() -> None: """Background task: periodically fetch OpenRouter models and insert new ones. diff --git a/ui/components/EditModelForm.tsx b/ui/components/EditModelForm.tsx index af4c53f5..0ac46cb4 100644 --- a/ui/components/EditModelForm.tsx +++ b/ui/components/EditModelForm.tsx @@ -121,7 +121,7 @@ export function EditModelForm({ const adminModel = await AdminService.getProviderModel( providerId, - model.full_name + model.id ); setAdminModelData(adminModel as AdminModelData); diff --git a/ui/components/ModelSelector.tsx b/ui/components/ModelSelector.tsx index 63c417e1..5b03e350 100644 --- a/ui/components/ModelSelector.tsx +++ b/ui/components/ModelSelector.tsx @@ -139,7 +139,7 @@ export function ModelSelector({ if (providerId === 'unknown') { continue; } - const modelFullNames = providerModels.map((m) => m.full_name); + const modelFullNames = providerModels.map((m) => m.id); const result = await AdminService.deleteModels( modelFullNames, providerId @@ -191,22 +191,18 @@ export function ModelSelector({ try { const existingModel = await AdminService.getProviderModel( providerIdNum, - model.full_name - ); - await AdminService.updateProviderModel( - providerIdNum, - model.full_name, - { - ...existingModel, - enabled: false, - } + model.id ); + await AdminService.updateProviderModel(providerIdNum, model.id, { + ...existingModel, + enabled: false, + }); totalDisabled++; } catch (fetchError: unknown) { const error = fetchError as { message?: string; status?: number }; if (error.message?.includes('404') || error.status === 404) { const newOverride = { - id: model.full_name, + id: model.id, name: model.name, description: model.description || '', created: Math.floor(Date.now() / 1000), @@ -340,10 +336,10 @@ export function ModelSelector({ const providerId = parseInt(model.provider_id); const existingModel = await AdminService.getProviderModel( providerId, - model.full_name + model.id ); - await AdminService.updateProviderModel(providerId, model.full_name, { + await AdminService.updateProviderModel(providerId, model.id, { ...existingModel, enabled: true, }); @@ -389,7 +385,7 @@ export function ModelSelector({ try { const existingModel = await AdminService.getProviderModel( providerIdNum, - model.full_name + model.id ); await AdminService.updateProviderModel( providerIdNum, @@ -747,7 +743,7 @@ export function ModelSelector({ try { const existingModel = await AdminService.getProviderModel( providerId, - model.full_name + model.id ); await AdminService.updateProviderModel(providerId, model.full_name, { @@ -759,7 +755,7 @@ export function ModelSelector({ const error = fetchError as { message?: string; status?: number }; if (error.message?.includes('404') || error.status === 404) { const newOverride = { - id: model.full_name, + id: model.id, name: model.name, description: model.description || '', created: Math.floor(Date.now() / 1000), diff --git a/ui/lib/api/services/admin.ts b/ui/lib/api/services/admin.ts index 0a03327f..e08f3e17 100644 --- a/ui/lib/api/services/admin.ts +++ b/ui/lib/api/services/admin.ts @@ -497,7 +497,7 @@ export class AdminService { internal_reasoning: 0, }; - const modelId = (data.id as string) || (data.full_name as string); + const modelId = data.id as string; const payload = { model_id: modelId,