Spaces:
Sleeping
Sleeping
File size: 6,433 Bytes
2415446 c817fe8 2415446 c817fe8 2415446 c817fe8 2415446 c817fe8 2415446 0a54372 2415446 c817fe8 2415446 a1bab2d 2415446 a1bab2d 2415446 0a54372 2415446 a1bab2d 2415446 a1bab2d 2415446 a1bab2d 2415446 a1bab2d 2415446 a1bab2d 2415446 c817fe8 2415446 a1bab2d 2415446 a1bab2d 2415446 c817fe8 2415446 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 | """Provider model-list discovery and background refresh."""
import asyncio
from collections.abc import Callable
import httpx
from loguru import logger
from free_claude_code.application.errors import ApplicationUnavailableError
from free_claude_code.application.model_metadata import (
ProviderModelInfo,
ProviderModelRefreshResult,
)
from free_claude_code.config.model_refs import configured_chat_model_refs
from free_claude_code.config.provider_catalog import PROVIDER_CATALOG
from free_claude_code.config.settings import Settings
from free_claude_code.core.failures import ExecutionFailure
from free_claude_code.providers.base import BaseProvider
from free_claude_code.providers.model_listing import ModelListResponseError
from .config import has_provider_configuration
from .model_cache import ProviderModelCache
ProviderResolver = Callable[[str], BaseProvider]
def _provider_query_failure_reason(exc: BaseException, settings: Settings) -> str:
"""Return a concise model-list query failure reason for user-facing logs."""
if isinstance(exc, ModelListResponseError):
return f"malformed model-list response: {exc.message}"
if isinstance(exc, httpx.HTTPStatusError):
return f"query failure: HTTP {exc.response.status_code}"
if isinstance(exc, ApplicationUnavailableError):
return f"query failure: {exc.message}"
if isinstance(exc, ExecutionFailure) and settings.log_api_error_tracebacks:
return f"query failure: {exc.message}"
return f"query failure: {type(exc).__name__}"
def referenced_provider_ids(settings: Settings) -> tuple[str, ...]:
"""Return unique provider ids referenced by configured chat models."""
return tuple(
dict.fromkeys(ref.provider_id for ref in configured_chat_model_refs(settings))
)
def model_cache_provider_ids_for_settings(
settings: Settings,
connected_provider_ids: tuple[str, ...] = (),
) -> tuple[str, ...]:
"""Return providers whose model metadata is valid for these settings."""
configured = tuple(
provider_id
for provider_id, descriptor in PROVIDER_CATALOG.items()
if has_provider_configuration(descriptor, settings)
)
available = set(configured) | set(connected_provider_ids)
return tuple(
provider_id for provider_id in PROVIDER_CATALOG if provider_id in available
)
def model_list_provider_ids_for_settings(
settings: Settings,
connected_provider_ids: tuple[str, ...] = (),
) -> tuple[str, ...]:
"""Return providers worth discovering for this process configuration."""
referenced_ids = referenced_provider_ids(settings)
return tuple(
provider_id
for provider_id in model_cache_provider_ids_for_settings(
settings, connected_provider_ids
)
if not PROVIDER_CATALOG[provider_id].local or provider_id in referenced_ids
)
class ProviderModelDiscovery:
"""Refresh provider model-list metadata for one provider runtime."""
def __init__(
self,
settings: Settings,
provider_resolver: ProviderResolver,
model_cache: ProviderModelCache,
connected_provider_ids: tuple[str, ...] = (),
) -> None:
self._settings = settings
self._provider_resolver = provider_resolver
self._model_cache = model_cache
self._connected_provider_ids = connected_provider_ids
async def warm_referenced_model_cache(self) -> ProviderModelRefreshResult:
"""Synchronously cache model metadata for routed providers."""
return await self._refresh_model_infos(referenced_provider_ids(self._settings))
async def refresh_model_list_cache(
self, *, only_missing: bool = False
) -> ProviderModelRefreshResult:
"""Best-effort refresh of model lists for usable providers."""
provider_ids = model_list_provider_ids_for_settings(
self._settings, self._connected_provider_ids
)
if only_missing:
provider_ids = tuple(
provider_id
for provider_id in provider_ids
if not self._model_cache.has_provider(provider_id)
)
return await self._refresh_model_infos(provider_ids)
async def refresh_provider(self, provider_id: str) -> ProviderModelRefreshResult:
"""Refresh exactly one dynamically changed provider."""
return await self._refresh_model_infos((provider_id,))
async def _refresh_model_infos(
self, provider_ids: tuple[str, ...]
) -> ProviderModelRefreshResult:
failed_provider_ids: list[str] = []
tasks: dict[str, asyncio.Task[frozenset[ProviderModelInfo]]] = {}
for provider_id in provider_ids:
try:
provider = self._provider_resolver(provider_id)
except Exception as exc:
self._log_discovery_failure(provider_id, exc)
failed_provider_ids.append(provider_id)
continue
tasks[provider_id] = asyncio.create_task(provider.list_model_infos())
refreshed_provider_ids: list[str] = []
if tasks:
results = await asyncio.gather(*tasks.values(), return_exceptions=True)
for (provider_id, _task), result in zip(
tasks.items(), results, strict=True
):
if isinstance(result, BaseException):
if isinstance(result, asyncio.CancelledError):
raise result
self._log_discovery_failure(provider_id, result)
failed_provider_ids.append(provider_id)
continue
self._model_cache.cache_model_infos(provider_id, result)
refreshed_provider_ids.append(provider_id)
logger.info(
"Provider model discovery cached: provider={} models={}",
provider_id,
len(result),
)
return ProviderModelRefreshResult(
refreshed_provider_ids=tuple(refreshed_provider_ids),
failed_provider_ids=tuple(failed_provider_ids),
)
def _log_discovery_failure(self, provider_id: str, exc: BaseException) -> None:
logger.warning(
"Provider model discovery skipped: provider={} reason={}",
provider_id,
_provider_query_failure_reason(exc, self._settings),
)
|