From ea655b748b51b385bb3a3d453fa63b56fea33b23 Mon Sep 17 00:00:00 2001 From: Shroominic Date: Fri, 26 Dec 2025 12:33:42 +0100 Subject: [PATCH] optimize startup time --- routstr/proxy.py | 38 +++++++++++++++++-------------------- routstr/upstream/helpers.py | 19 +++++++++++++------ 2 files changed, 30 insertions(+), 27 deletions(-) diff --git a/routstr/proxy.py b/routstr/proxy.py index 0fa48f4f..f79fef1c 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -80,31 +80,27 @@ def get_unique_models() -> list[Model]: async def refresh_model_maps() -> None: """Refresh global model and provider maps using the cost-based algorithm.""" + from sqlalchemy.orm import selectinload + global _model_instances, _provider_map, _unique_models - # Gather database overrides and disabled models async with create_session() as session: - result = await session.exec(select(ModelRow).where(ModelRow.enabled)) - override_rows = result.all() - - provider_result = await session.exec(select(UpstreamProviderRow)) - providers_by_id = {p.id: p for p in provider_result.all()} - - overrides_by_id: dict[str, tuple[ModelRow, float]] = { - row.id: ( - row, - providers_by_id[row.upstream_provider_id].provider_fee - if row.upstream_provider_id in providers_by_id - else 1.01, - ) - for row in override_rows - if row.upstream_provider_id is not None - } - - disabled_result = await session.exec( - select(ModelRow.id).where(ModelRow.enabled == False) # noqa: E712 + # Fetch all providers with their models in a single logical operation + query = select(UpstreamProviderRow).options( + selectinload(UpstreamProviderRow.models) # type: ignore ) - disabled_model_ids = {row for row in disabled_result.all()} + result = await session.exec(query) + provider_rows = result.all() + + overrides_by_id: dict[str, tuple[ModelRow, float]] = {} + disabled_model_ids: set[str] = set() + + for provider in provider_rows: + for model in provider.models: + if model.enabled: + overrides_by_id[model.id] = (model, provider.provider_fee) + else: + disabled_model_ids.add(model.id) _model_instances, _provider_map, _unique_models = create_model_mappings( upstreams=_upstreams, diff --git a/routstr/upstream/helpers.py b/routstr/upstream/helpers.py index 36ad6d76..c2417d40 100644 --- a/routstr/upstream/helpers.py +++ b/routstr/upstream/helpers.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import os import re from typing import TYPE_CHECKING @@ -7,6 +8,8 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: from ..core.settings import Settings +from sqlmodel import select + from ..core import get_logger from ..core.db import AsyncSession, ModelRow, UpstreamProviderRow, create_session from ..payment.models import Model @@ -176,8 +179,6 @@ async def init_upstreams() -> list[BaseUpstreamProvider]: Seeds database with providers from settings if empty, then loads and instantiates provider instances from database records, and refreshes their models cache. """ - from sqlmodel import select - from ..core.settings import settings async with create_session() as session: @@ -193,16 +194,16 @@ async def init_upstreams() -> list[BaseUpstreamProvider]: result = await session.exec(select(UpstreamProviderRow)) existing_providers = result.all() - upstreams: list[BaseUpstreamProvider] = [] - for provider_row in existing_providers: + async def _init_single_provider( + provider_row: UpstreamProviderRow, + ) -> BaseUpstreamProvider | None: if not provider_row.enabled: logger.debug(f"Skipping disabled provider: {provider_row.base_url}") - continue + return None provider = _instantiate_provider(provider_row) if provider: await provider.refresh_models_cache() - upstreams.append(provider) logger.debug( f"Initialized {provider_row.provider_type} provider", extra={ @@ -210,6 +211,12 @@ async def init_upstreams() -> list[BaseUpstreamProvider]: "models_cached": len(provider.get_cached_models()), }, ) + return provider + return None + + tasks = [_init_single_provider(row) for row in existing_providers] + results = await asyncio.gather(*tasks) + upstreams = [p for p in results if p is not None] return upstreams