Yash030/claude-code-proxy
2
1"""Model routing for Claude-compatible requests."""2 3from __future__ import annotations4 5from dataclasses import dataclass6 7from loguru import logger8 9from config.provider_ids import SUPPORTED_PROVIDER_IDS10from config.settings import Settings11from core.model_capabilities import find_best_model_for_task12from core.session_tracker import SessionTracker13from core.task_detector import TaskDetector14from providers.rate_limit import GlobalRateLimiter15 16from .gateway_model_ids import decode_gateway_model_id17from .models.anthropic import MessagesRequest, TokenCountRequest18 19# Default NIM models to include in auto routing (in order of preference)20DEFAULT_NIM_AUTO_MODELS = [21 "nvidia_nim/qwen/qwen3-coder-480b-a35b-instruct",22 "nvidia_nim/z-ai/glm4.7",23 "nvidia_nim/stepfun-ai/step-3.5-flash",24 "nvidia_nim/mistralai/mistral-large-3-675b-instruct-2512",25 "nvidia_nim/abacusai/dracarys-llama-3.1-70b-instruct",26 "nvidia_nim/bytedance/seed-oss-36b-instruct",27 "nvidia_nim/mistralai/mistral-nemotron",28]29 30 31@dataclass(frozen=True, slots=True)32class ResolvedModel:33 original_model: str34 provider_id: str35 provider_model: str36 provider_model_ref: str37 thinking_enabled: bool38 39 40@dataclass(frozen=True, slots=True)41class RoutedMessagesRequest:42 request: MessagesRequest43 resolved: ResolvedModel44 45 46@dataclass(frozen=True, slots=True)47class RoutedTokenCountRequest:48 request: TokenCountRequest49 resolved: ResolvedModel50 51 52class ModelRouter:53 """Resolve incoming Claude model names to configured provider/model pairs."""54 55 def __init__(self, settings: Settings):56 self._settings = settings57 58 def _is_auto(self, model_name: str) -> bool:59 """Return whether the model name refers to the virtual 'auto' model."""60 name_lower = model_name.lower()61 return name_lower == "auto" or name_lower == "anthropic/auto"62 63 def _normalize_candidate_ref(self, raw_ref: str) -> str | None:64 """Normalize auto candidate refs to ``provider/model`` when possible."""65 candidate = (raw_ref or "").strip()66 if not candidate:67 return None68 69 provider_id, separator, remainder = candidate.partition("/")70 if separator and provider_id in SUPPORTED_PROVIDER_IDS and remainder:71 return f"{provider_id}/{remainder}"72 73 # Treat bare model ids and vendor/model ids as NVIDIA NIM models.74 return f"nvidia_nim/{candidate}"75 76 def resolve(self, claude_model_name: str) -> ResolvedModel:77 # Special virtual model 'auto' maps to the configured default MODEL and78 # enables provider-side fallbacks. Resolve it to the configured model79 # while preserving the original requested name.80 if self._is_auto(claude_model_name):81 # If the user configured an explicit AUTO_MODEL_ORDER, try each82 # provider/model pair in order and pick the first provider that is83 # plausibly configured. Fall back to the single configured MODEL.84 order_csv = (self._settings.auto_model_order or "").strip()85 if order_csv:86 for cand in [c.strip() for c in order_csv.split(",") if c.strip()]:87 if "/" not in cand:88 # assume vendor-prefixed entries; skip malformed89 continue90 provider_id = Settings.parse_provider_type(cand)91 provider_model = Settings.parse_model_name(cand)92 if self._settings.provider_is_configured(provider_id):93 thinking_enabled = self._settings.resolve_thinking(94 claude_model_name95 )96 return ResolvedModel(97 original_model=claude_model_name,98 provider_id=provider_id,99 provider_model=provider_model,100 provider_model_ref=cand,101 thinking_enabled=thinking_enabled,102 )103 # No explicit order matched or none configured — fall back to default MODEL104 provider_model_ref = self._settings.model105 provider_id = Settings.parse_provider_type(provider_model_ref)106 provider_model = Settings.parse_model_name(provider_model_ref)107 thinking_enabled = self._settings.resolve_thinking(claude_model_name)108 return ResolvedModel(109 original_model=claude_model_name,110 provider_id=provider_id,111 provider_model=provider_model,112 provider_model_ref=provider_model_ref,113 thinking_enabled=thinking_enabled,114 )115 116 (117 direct_provider_id,118 direct_provider_model,119 force_thinking_enabled,120 ) = self._direct_provider_model(claude_model_name)121 if direct_provider_id is not None and direct_provider_model is not None:122 thinking_enabled = (123 force_thinking_enabled124 if force_thinking_enabled is not None125 else self._settings.resolve_thinking(direct_provider_model)126 )127 logger.debug(128 "MODEL DIRECT: '{}' -> provider='{}' model='{}' thinking={}",129 claude_model_name,130 direct_provider_id,131 direct_provider_model,132 thinking_enabled,133 )134 return ResolvedModel(135 original_model=claude_model_name,136 provider_id=direct_provider_id,137 provider_model=direct_provider_model,138 provider_model_ref=claude_model_name,139 thinking_enabled=thinking_enabled,140 )141 142 provider_model_ref = self._settings.resolve_model(claude_model_name)143 thinking_enabled = self._settings.resolve_thinking(claude_model_name)144 provider_id = Settings.parse_provider_type(provider_model_ref)145 provider_model = Settings.parse_model_name(provider_model_ref)146 if provider_model != claude_model_name:147 logger.debug(148 "MODEL MAPPING: '{}' -> '{}'", claude_model_name, provider_model149 )150 return ResolvedModel(151 original_model=claude_model_name,152 provider_id=provider_id,153 provider_model=provider_model,154 provider_model_ref=provider_model_ref,155 thinking_enabled=thinking_enabled,156 )157 158 def resolve_candidates(self, claude_model_name: str) -> list[ResolvedModel]:159 """Resolve a model name to a prioritized list of candidates.160 161 Used by the 'auto' routing logic to implement provider-side failover.162 Considers session load for fair resource sharing across multiple clients.163 164 Priority order:165 1. AUTO_MODEL_ORDER (if configured)166 2. MODEL (primary)167 3. NVIDIA NIM fallback models (if configured, or DEFAULT_NIM_AUTO_MODELS)168 4. MODEL_OPUS, MODEL_SONNET, MODEL_HAIKU169 """170 if not self._is_auto(claude_model_name):171 return [self.resolve(claude_model_name)]172 173 healthy_candidates: list[ResolvedModel] = []174 blocked_candidates: list[ResolvedModel] = []175 seen: set[str] = set()176 session_tracker = SessionTracker.get_instance()177 178 def add_candidate(ref: str | None, source: str) -> None:179 normalized_ref = self._normalize_candidate_ref(ref or "")180 if normalized_ref is None or normalized_ref in seen:181 return182 provider_id = Settings.parse_provider_type(normalized_ref)183 provider_model = Settings.parse_model_name(normalized_ref)184 if self._settings.provider_is_configured(provider_id):185 seen.add(normalized_ref)186 resolved = ResolvedModel(187 original_model=claude_model_name,188 provider_id=provider_id,189 provider_model=provider_model,190 provider_model_ref=normalized_ref,191 thinking_enabled=self._settings.resolve_thinking(claude_model_name),192 )193 194 limiter = GlobalRateLimiter.get_scoped_instance(provider_id)195 is_blocked = limiter.is_blocked()196 197 # For Zen provider, never consider it blocked (no rate limits)198 if provider_id == "zen":199 is_blocked = False200 201 # Check model health (recent failures)202 is_healthy = limiter.is_healthy(normalized_ref)203 204 if is_blocked or not is_healthy:205 reason = "BLOCKED" if is_blocked else "UNHEALTHY"206 logger.debug(207 "Routing: candidate '{}' (from {}) is {} (health={})",208 normalized_ref,209 source,210 reason,211 is_healthy,212 )213 blocked_candidates.append(resolved)214 else:215 # Smart ordering: Zen (no rate limits) gets priority, then by load216 logger.debug(217 "Routing: added candidate '{}' (from {})",218 normalized_ref,219 source,220 )221 healthy_candidates.append(resolved)222 223 else:224 logger.debug(225 "Routing: candidate '{}' (from {}) is NOT CONFIGURED",226 normalized_ref,227 source,228 )229 230 # 1. AUTO_MODEL_ORDER (user-configured priority)231 order_csv = (self._settings.auto_model_order or "").strip()232 if order_csv:233 for cand in [c.strip() for c in order_csv.split(",") if c.strip()]:234 add_candidate(cand, "AUTO_MODEL_PRIORITY")235 236 # 2. Primary MODEL237 add_candidate(self._settings.model, "MODEL")238 239 # 3. NVIDIA Fallbacks - use configured or defaults240 nim_csv = (self._settings.nvidia_nim_fallback_models or "").strip()241 if nim_csv:242 for cand in [c.strip() for c in nim_csv.split(",") if c.strip()]:243 add_candidate(cand, "NVIDIA_NIM_FALLBACK_MODELS")244 else:245 # Use default NIM models when no explicit fallback configured246 for cand in DEFAULT_NIM_AUTO_MODELS:247 add_candidate(cand, "DEFAULT_NIM_AUTO_MODELS")248 249 # 4. Model-specific overrides250 add_candidate(self._settings.model_opus, "MODEL_OPUS")251 add_candidate(self._settings.model_sonnet, "MODEL_SONNET")252 add_candidate(self._settings.model_haiku, "MODEL_HAIKU")253 254 # Smart ordering: Zen goes first (no rate limits), then sort by load255 def provider_priority(c: ResolvedModel) -> tuple:256 # Priority: zen > others, then by active request count257 is_zen = 0 if c.provider_id == "zen" else 1258 active = session_tracker._provider_active.get(c.provider_id, 0)259 return (is_zen, active)260 261 healthy_candidates.sort(key=provider_priority)262 263 all_candidates = healthy_candidates + blocked_candidates264 logger.info(265 "Routing: resolved '{}' to {} candidates: {}",266 claude_model_name,267 len(all_candidates),268 ", ".join(c.provider_model_ref for c in all_candidates),269 )270 return all_candidates271 272 def _direct_provider_model(273 self, model_name: str274 ) -> tuple[str | None, str | None, bool | None]:275 decoded = decode_gateway_model_id(model_name)276 if decoded is not None:277 if decoded.provider_id not in SUPPORTED_PROVIDER_IDS:278 return None, None, None279 return (280 decoded.provider_id,281 decoded.provider_model,282 decoded.force_thinking_enabled,283 )284 285 provider_id, separator, provider_model = model_name.partition("/")286 if not separator:287 return None, None, None288 if provider_id not in SUPPORTED_PROVIDER_IDS:289 return None, None, None290 if not provider_model:291 return None, None, None292 return provider_id, provider_model, None293 294 def resolve_messages_request(295 self, request: MessagesRequest296 ) -> RoutedMessagesRequest:297 """Return an internal routed request context."""298 resolved = self.resolve(request.model)299 routed = request.model_copy(deep=True)300 routed.model = resolved.provider_model301 return RoutedMessagesRequest(request=routed, resolved=resolved)302 303 def resolve_token_count_request(304 self, request: TokenCountRequest305 ) -> RoutedTokenCountRequest:306 """Return an internal token-count request context."""307 resolved = self.resolve(request.model)308 routed = request.model_copy(309 update={"model": resolved.provider_model}, deep=True310 )311 return RoutedTokenCountRequest(request=routed, resolved=resolved)312 313 def resolve_with_task_awareness(314 self,315 claude_model_name: str,316 messages: list,317 ) -> ResolvedModel:318 """Resolve model with task-based capability matching.319 320 For 'auto' model, detects task requirements and routes to best-capable model.321 """322 if not self._is_auto(claude_model_name):323 return self.resolve(claude_model_name)324 325 # Detect what capabilities are needed326 detector = TaskDetector()327 requirements = detector.detect_requirements(messages)328 329 logger.info(330 "Task-aware routing: detected requirements={} confidence={:.2f}",331 requirements.required_capabilities,332 requirements.confidence,333 )334 335 # Get available candidates336 candidates = self.resolve_candidates(claude_model_name)337 338 if not candidates:339 # Fallback to default340 return self.resolve(claude_model_name)341 342 # If confidence is low or only general text needed, use load-based selection343 if requirements.confidence < 0.7 or (344 not requirements.requires_vision345 and not requirements.requires_coding346 and not requirements.requires_reasoning347 ):348 logger.debug(349 "Task-aware routing: low confidence, using load-based selection"350 )351 return candidates[0]352 353 # Find best model matching required capabilities354 required_caps = set()355 if requirements.requires_coding:356 required_caps.add("coding")357 if requirements.requires_reasoning:358 required_caps.add("reasoning")359 if requirements.requires_vision:360 required_caps.add("vision")361 362 if required_caps:363 model_refs = [c.provider_model_ref for c in candidates]364 best = find_best_model_for_task(required_caps, model_refs)365 if best:366 # Find the matching candidate367 for cand in candidates:368 if cand.provider_model_ref == best.model_ref:369 logger.info(370 "Task-aware routing: selected {} for capabilities={}",371 best.model_ref,372 required_caps,373 )374 return cand375 376 # Default to first candidate (load-balanced)377 return candidates[0]378 379 def get_routing_hint(self, messages: list) -> str:380 """Get a hint about what kind of model would be best."""381 detector = TaskDetector()382 requirements = detector.detect_requirements(messages)383 return detector.get_priority_hint(requirements)384 