File size: 10,205 Bytes
b2005c6 2f1cc5a 0d0f0f9 2f1cc5a af75bef 2f1cc5a dc04eee e8f2acc 7b3c6de 0602b26 dc04eee e8f2acc dc04eee b2005c6 dc04eee b2005c6 71665ff dc04eee 71665ff b2005c6 e8f2acc b2005c6 1fdea9f dc04eee 1fdea9f 4714e0f dc04eee 1fdea9f 4714e0f dc04eee 4714e0f dc04eee 4714e0f dc04eee 4714e0f 562782c e8f2acc dc04eee 562782c 9a1b3fc 562782c b2005c6 dc04eee b2005c6 e8f2acc af75bef 2f1cc5a 1fdea9f 4714e0f dc04eee 1fdea9f dc04eee 1fdea9f dc04eee 1fdea9f 4714e0f dc04eee 4714e0f dc04eee 4714e0f dc04eee 4714e0f 562782c dc04eee 562782c dc04eee 562782c dc04eee 562782c 71665ff b2005c6 dc04eee b2005c6 71665ff b2005c6 71665ff | 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 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 | """Model routing for Claude-compatible requests."""
from dataclasses import dataclass
from loguru import logger
from free_claude_code.application.errors import UnknownProviderError
from free_claude_code.config.model_refs import parse_model_name, parse_provider_type
from free_claude_code.config.provider_catalog import (
OPENAI_COMPATIBLE_INSTANCE_RE,
PROVIDER_CATALOG,
SUPPORTED_PROVIDER_IDS,
)
from free_claude_code.config.reasoning import ReasoningPreference
from free_claude_code.config.settings import OpenAICompatibleInstance, Settings
from free_claude_code.core.anthropic import MessagesRequest, TokenCountRequest
from free_claude_code.core.gateway_model_ids import decode_gateway_model_id
from free_claude_code.core.reasoning import ReasoningPolicy
from free_claude_code.providers.runtime.config import (
mask_api_key,
openai_compatible_model_candidates,
split_api_key_pool,
)
from .reasoning import resolve_reasoning_policy
_ROUTE_SETTINGS = (
("fable", "model_fable", "reasoning_fable"),
("opus", "model_opus", "reasoning_opus"),
("haiku", "model_haiku", "reasoning_haiku"),
("sonnet", "model_sonnet", "reasoning_sonnet"),
)
@dataclass(frozen=True, slots=True)
class ResolvedModel:
original_model: str
provider_id: str
provider_model: str
provider_model_ref: str
reasoning_preference: ReasoningPreference
@dataclass(frozen=True, slots=True)
class RoutedMessagesRequest:
request: MessagesRequest
resolved: ResolvedModel
reasoning: ReasoningPolicy
@dataclass(frozen=True, slots=True)
class RoutedTokenCountRequest:
request: TokenCountRequest
resolved: ResolvedModel
class ModelRouter:
"""Resolve incoming Claude model names to configured provider/model pairs."""
def __init__(self, settings: Settings):
self._settings = settings
# Per-model rotation counters for bare model ids served by multiple
# OpenAI-compatible endpoints. Best-effort rotation across requests;
# exact fairness is not required for request distribution.
self._rotations: dict[str, int] = {}
def resolve(self, claude_model_name: str) -> ResolvedModel:
(
direct_provider_id,
direct_provider_model,
force_reasoning_off,
) = self._direct_provider_model(claude_model_name)
if direct_provider_id is not None and direct_provider_model is not None:
reasoning_preference = (
ReasoningPreference.OFF
if force_reasoning_off
else self._settings.reasoning_policy
)
logger.debug(
"MODEL DIRECT: '{}' -> provider='{}' model='{}' reasoning={}",
claude_model_name,
direct_provider_id,
direct_provider_model,
reasoning_preference.value,
)
return ResolvedModel(
original_model=claude_model_name,
provider_id=direct_provider_id,
provider_model=direct_provider_model,
provider_model_ref=claude_model_name,
reasoning_preference=reasoning_preference,
)
provider_model_ref = self._resolve_model_ref(claude_model_name)
# A bare model id (no provider prefix) is valid when an OpenAI-
# compatible endpoint advertises it; the provider prefix is optional.
if "/" not in provider_model_ref:
resolved = self._resolve_bare_endpoint_model(
provider_model_ref, claude_model_name
)
if resolved is not None:
return resolved
reasoning_preference = self._resolve_reasoning_preference(claude_model_name)
provider_id = parse_provider_type(provider_model_ref)
self._validate_provider_id(provider_id)
provider_model = parse_model_name(provider_model_ref)
if provider_model != claude_model_name:
logger.debug(
"MODEL MAPPING: '{}' -> '{}'", claude_model_name, provider_model
)
return ResolvedModel(
original_model=claude_model_name,
provider_id=provider_id,
provider_model=provider_model,
provider_model_ref=provider_model_ref,
reasoning_preference=reasoning_preference,
)
def _resolve_bare_endpoint_model(
self, model_id: str, original_model: str
) -> ResolvedModel | None:
"""Resolve a bare model id to an OpenAI-compatible endpoint.
Returns ``None`` when no configured endpoint advertises the id so the
caller falls back to the configured model refs.
"""
candidates = openai_compatible_model_candidates(self._settings, model_id)
if not candidates:
return None
provider_id, _instance = self._select_rotation_candidate(model_id, candidates)
self._log_rotation(model_id, candidates, provider_id)
return ResolvedModel(
original_model=original_model,
provider_id=provider_id,
provider_model=model_id,
provider_model_ref=f"{provider_id}/{model_id}",
reasoning_preference=self._resolve_reasoning_preference(original_model),
)
def _select_rotation_candidate(
self,
model_id: str,
candidates: tuple[tuple[str, OpenAICompatibleInstance], ...],
) -> tuple[str, OpenAICompatibleInstance]:
"""Pick the endpoint for a request, rotating across duplicates."""
if len(candidates) == 1:
return candidates[0]
counter = self._rotations.get(model_id, 0)
self._rotations[model_id] = counter + 1
return candidates[counter % len(candidates)]
@staticmethod
def _log_rotation(
model_id: str,
candidates: tuple[tuple[str, OpenAICompatibleInstance], ...],
chosen_provider_id: str,
) -> None:
"""Log provider/key distribution when one model id spans endpoints."""
entries = []
for provider_id, instance in candidates:
key_count = len(split_api_key_pool(instance.api_keys))
entries.append(
f"{provider_id} (key {mask_api_key(instance.api_keys)}, "
f"{key_count} key{'s' if key_count != 1 else ''})"
)
if len(candidates) > 1:
logger.info(
"MODEL MULTI-PROVIDER: '{}' served by {}; rotating to {} "
"(provider + key pool round-robin)",
model_id,
", ".join(entries),
chosen_provider_id,
)
else:
logger.debug("MODEL ENDPOINT: '{}' served by {}", model_id, entries[0])
def _validate_provider_id(self, provider_id: str) -> None:
if provider_id in PROVIDER_CATALOG:
return
# Numbered OpenAI-compatible endpoint instances are dynamic (configured
# via settings.openai_compatible_instances), so they are not catalog
# entries. Whether the specific number exists is validated by the
# provider factory, which reports the configured instance ids.
if OPENAI_COMPATIBLE_INSTANCE_RE.fullmatch(provider_id):
return
raise UnknownProviderError.for_provider(provider_id, PROVIDER_CATALOG)
def _direct_provider_model(
self, model_name: str
) -> tuple[str | None, str | None, bool]:
decoded = decode_gateway_model_id(model_name)
if decoded is not None:
if decoded.provider_id not in SUPPORTED_PROVIDER_IDS:
return None, None, False
return (
decoded.provider_id,
decoded.provider_model,
decoded.force_reasoning_off,
)
provider_id, separator, provider_model = model_name.partition("/")
if not separator:
return None, None, False
if provider_id not in SUPPORTED_PROVIDER_IDS:
return None, None, False
if not provider_model:
return None, None, False
return provider_id, provider_model, False
def _resolve_model_ref(self, claude_model_name: str) -> str:
"""Resolve a Claude model name to the configured provider/model ref."""
route = self._matched_route(claude_model_name)
if route is not None:
model = getattr(self._settings, route[1])
if isinstance(model, str):
return model
return self._settings.model
def _resolve_reasoning_preference(
self, claude_model_name: str
) -> ReasoningPreference:
"""Resolve a route override without inspecting the provider model."""
route = self._matched_route(claude_model_name)
if route is not None:
preference = getattr(self._settings, route[2])
if preference is not ReasoningPreference.INHERIT:
return preference
return self._settings.reasoning_policy
@staticmethod
def _matched_route(model_name: str) -> tuple[str, str, str] | None:
normalized = model_name.lower()
return next(
(route for route in _ROUTE_SETTINGS if route[0] in normalized),
None,
)
def resolve_messages_request(
self, request: MessagesRequest
) -> RoutedMessagesRequest:
"""Return an internal routed request context."""
resolved = self.resolve(request.model)
routed = request.model_copy(deep=True)
routed.model = resolved.provider_model
return RoutedMessagesRequest(
request=routed,
resolved=resolved,
reasoning=resolve_reasoning_policy(
routed,
resolved.reasoning_preference,
),
)
def resolve_token_count_request(
self, request: TokenCountRequest
) -> RoutedTokenCountRequest:
"""Return an internal token-count request context."""
resolved = self.resolve(request.model)
routed = request.model_copy(
update={"model": resolved.provider_model}, deep=True
)
return RoutedTokenCountRequest(request=routed, resolved=resolved)
|