Spaces:
Sleeping
Sleeping
File size: 6,624 Bytes
2415446 a1bab2d 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 a1bab2d 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 2415446 61bb677 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 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 | """Model routing for Claude-compatible requests."""
from __future__ import annotations
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 (
PROVIDER_CATALOG,
SUPPORTED_PROVIDER_IDS,
)
from free_claude_code.config.reasoning import ReasoningPreference
from free_claude_code.config.settings import 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 .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
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)
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,
)
@staticmethod
def _validate_provider_id(provider_id: str) -> None:
if provider_id not in PROVIDER_CATALOG:
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)
|