Spaces:
Sleeping
Sleeping
File size: 4,822 Bytes
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 | """Configured provider model validation."""
import asyncio
from collections import defaultdict
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
from free_claude_code.config.model_refs import (
ConfiguredChatModelRef,
configured_chat_model_refs,
)
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 .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__}"
class ConfiguredModelValidator:
"""Validate configured provider/model refs against upstream model lists."""
def __init__(
self,
settings: Settings,
provider_resolver: ProviderResolver,
model_cache: ProviderModelCache,
) -> None:
self._settings = settings
self._provider_resolver = provider_resolver
self._model_cache = model_cache
async def validate_configured_models(self) -> None:
"""Fail unless every configured chat model exists upstream."""
refs = configured_chat_model_refs(self._settings)
refs_by_provider: dict[str, list[ConfiguredChatModelRef]] = defaultdict(list)
for ref in refs:
refs_by_provider[ref.provider_id].append(ref)
failures: list[str] = []
tasks: dict[str, asyncio.Task[frozenset[ProviderModelInfo]]] = {}
for provider_id, provider_refs in refs_by_provider.items():
try:
provider = self._provider_resolver(provider_id)
except Exception as exc:
failures.extend(
self._format_provider_query_failures(provider_refs, exc)
)
continue
tasks[provider_id] = asyncio.create_task(provider.list_model_infos())
if tasks:
results = await asyncio.gather(*tasks.values(), return_exceptions=True)
for (provider_id, _task), result in zip(
tasks.items(), results, strict=True
):
provider_refs = refs_by_provider[provider_id]
if isinstance(result, BaseException):
if isinstance(result, asyncio.CancelledError):
raise result
failures.extend(
self._format_provider_query_failures(provider_refs, result)
)
continue
model_ids = frozenset(info.model_id for info in result)
self._model_cache.cache_model_infos(provider_id, result)
failures.extend(
self._format_missing_model_failure(ref)
for ref in provider_refs
if ref.model_id not in model_ids
)
if failures:
message = "Configured model validation failed:\n" + "\n".join(
f"- {failure}" for failure in failures
)
raise ApplicationUnavailableError(message)
logger.info(
"Configured provider models validated: models={} providers={}",
len(refs),
len(refs_by_provider),
)
def _format_provider_query_failures(
self,
refs: list[ConfiguredChatModelRef],
exc: BaseException,
) -> list[str]:
reason = provider_query_failure_reason(exc, self._settings)
return [self._format_model_validation_failure(ref, reason) for ref in refs]
def _format_missing_model_failure(self, ref: ConfiguredChatModelRef) -> str:
return self._format_model_validation_failure(ref, "missing model")
@staticmethod
def _format_model_validation_failure(
ref: ConfiguredChatModelRef, problem: str
) -> str:
return (
f"sources={','.join(ref.sources)} provider={ref.provider_id} "
f"model={ref.model_id} problem={problem}"
)
|