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),
        )