Yash030/claude-code-proxy
2
1"""Application services for the Claude-compatible API."""2 3from __future__ import annotations4 5import traceback6import uuid7from collections.abc import AsyncIterator, Callable8from typing import Any9 10from fastapi import HTTPException, Request11from fastapi.responses import StreamingResponse12from loguru import logger13 14from config.settings import Settings, get_settings15from core.anthropic import get_token_count, get_user_facing_error_message16from core.anthropic.sse import ANTHROPIC_SSE_RESPONSE_HEADERS, format_sse_event17from core.session_tracker import SessionTracker18from providers.base import BaseProvider19from providers.exceptions import (20 InvalidRequestError,21 OverloadedError,22 ProviderError,23 RateLimitError,24)25 26from .model_router import ModelRouter, ResolvedModel27from .models.anthropic import MessagesRequest, TokenCountRequest28from .models.responses import TokenCountResponse29from .optimization_handlers import try_optimizations30from .web_tools.egress import WebFetchEgressPolicy31from .web_tools.request import (32 is_web_server_tool_request,33 openai_chat_upstream_server_tool_error,34)35 36TokenCounter = Callable[[list[Any], str | list[Any] | None, list[Any] | None], int]37 38ProviderGetter = Callable[[str], BaseProvider]39 40# Providers that use ``/chat/completions`` + Anthropic-to-OpenAI conversion (not native Messages).41_OPENAI_CHAT_UPSTREAM_IDS = frozenset({"nvidia_nim", "groq", "cerebras", "silicon"})42 43 44def anthropic_sse_streaming_response(45 body: AsyncIterator[str],46) -> StreamingResponse:47 """Return a :class:`StreamingResponse` for Anthropic-style SSE streams."""48 return StreamingResponse(49 body,50 media_type="text/event-stream",51 headers=ANTHROPIC_SSE_RESPONSE_HEADERS,52 )53 54 55def _http_status_for_unexpected_service_exception(_exc: BaseException) -> int:56 """HTTP status for uncaught non-provider failures (stable client contract)."""57 return 50058 59 60def _log_unexpected_service_exception(61 settings: Settings,62 exc: BaseException,63 *,64 context: str,65 request_id: str | None = None,66) -> None:67 """Log service-layer failures without echoing exception text unless opted in."""68 if settings.log_api_error_tracebacks:69 if request_id is not None:70 logger.error("{} request_id={}: {}", context, request_id, exc)71 else:72 logger.error("{}: {}", context, exc)73 logger.error(traceback.format_exc())74 return75 if request_id is not None:76 logger.error(77 "{} request_id={} exc_type={}",78 context,79 request_id,80 type(exc).__name__,81 )82 else:83 logger.error("{} exc_type={}", context, type(exc).__name__)84 85 86def _require_non_empty_messages(messages: list[Any]) -> None:87 if not messages:88 raise InvalidRequestError("messages cannot be empty")89 90 91def _get_client_ip(request: Request) -> str | None:92 """Extract client IP from gateway headers or return None for direct connections."""93 # Check for proxy/gateway headers94 forwarded = request.headers.get("X-Forwarded-For")95 if forwarded:96 return forwarded.split(",")[0].strip()97 real_ip = request.headers.get("X-Real-IP")98 if real_ip:99 return real_ip100 client_ip = request.headers.get("X-Client-IP")101 if client_ip:102 return client_ip103 via = request.headers.get("Via")104 if via:105 return request.client.host # Gateway/proxy IP106 return None # Direct connection107 108 109def _get_session_id(request: Request) -> str:110 """Get session ID from X-Session-ID header or fall back to gateway IP.111 112 Claude Code sends X-Session-ID when started with --session-id <uuid>.113 """114 session = request.headers.get("X-Session-ID")115 if session:116 return session117 ip = _get_client_ip(request)118 return f"gateway_{ip}" if ip else "direct"119 120 121class ClaudeProxyService:122 """Coordinate request optimization, model routing, and providers."""123 124 def __init__(125 self,126 settings: Settings,127 provider_getter: ProviderGetter,128 model_router: ModelRouter | None = None,129 token_counter: TokenCounter = get_token_count,130 ):131 self._settings = settings132 self._provider_getter = provider_getter133 self._model_router = model_router or ModelRouter(settings)134 self._token_counter = token_counter135 settings_local = get_settings()136 self._session_tracker = SessionTracker.get_instance(137 retention_seconds=settings_local.session_retention_minutes * 60138 )139 140 def create_message(self, request: Request, request_data: MessagesRequest) -> object:141 """Create a message response or streaming response with optional failover."""142 try:143 _require_non_empty_messages(request_data.messages)144 145 candidates = self._model_router.resolve_candidates(request_data.model)146 if not candidates:147 raise InvalidRequestError(148 f"No configured models available for '{request_data.model}'"149 )150 151 # Debug log what we're routing to152 from loguru import logger153 154 logger.info(155 "REQUEST_MODEL_ROUTING: requested={} resolved_provider={} resolved_model={}",156 request_data.model,157 candidates[0].provider_id,158 candidates[0].provider_model,159 )160 161 # For 'auto' requests with multiple candidates, we wrap the stream in a failover loop.162 if len(candidates) > 1:163 return anthropic_sse_streaming_response(164 self._stream_with_fallbacks(request, candidates, request_data)165 )166 167 # Standard path for single-model requests168 return self._create_single_message(request, candidates[0], request_data)169 170 except ProviderError:171 raise172 except Exception as e:173 _log_unexpected_service_exception(174 self._settings, e, context="CREATE_MESSAGE_ERROR"175 )176 raise HTTPException(177 status_code=_http_status_for_unexpected_service_exception(e),178 detail=get_user_facing_error_message(e),179 ) from e180 181 def _create_single_message(182 self, request: Request, resolved: ResolvedModel, request_data: MessagesRequest183 ) -> object:184 """Create a single message response from a resolved model."""185 routed_request = request_data.model_copy(deep=True)186 routed_request.model = resolved.provider_model187 188 if resolved.provider_id in _OPENAI_CHAT_UPSTREAM_IDS:189 tool_err = openai_chat_upstream_server_tool_error(190 routed_request,191 web_tools_enabled=self._settings.enable_web_server_tools,192 )193 if tool_err is not None:194 raise InvalidRequestError(tool_err)195 196 if self._settings.enable_web_server_tools and is_web_server_tool_request(197 routed_request198 ):199 from .web_tools.streaming import stream_web_server_tool_response200 201 input_tokens = self._token_counter(202 routed_request.messages, routed_request.system, routed_request.tools203 )204 logger.info("Optimization: Handling Anthropic web server tool")205 egress = WebFetchEgressPolicy(206 allow_private_network_targets=self._settings.web_fetch_allow_private_networks,207 allowed_schemes=self._settings.web_fetch_allowed_scheme_set(),208 )209 return anthropic_sse_streaming_response(210 stream_web_server_tool_response(211 routed_request,212 input_tokens=input_tokens,213 web_fetch_egress=egress,214 verbose_client_errors=self._settings.log_api_error_tracebacks,215 ),216 )217 218 optimized = try_optimizations(routed_request, self._settings)219 if optimized is not None:220 return optimized221 222 provider = self._provider_getter(resolved.provider_id)223 provider.preflight_stream(224 routed_request,225 thinking_enabled=resolved.thinking_enabled,226 )227 228 session_id = _get_session_id(request)229 self._session_tracker.track_request_sync(session_id, resolved.provider_id)230 231 request_id = f"req_{uuid.uuid4().hex[:12]}"232 logger.info(233 "API_REQUEST: request_id={} model={} messages={}",234 request_id,235 routed_request.model,236 len(routed_request.messages),237 )238 239 input_tokens = self._token_counter(240 routed_request.messages, routed_request.system, routed_request.tools241 )242 return anthropic_sse_streaming_response(243 provider.stream_response(244 routed_request,245 input_tokens=input_tokens,246 request_id=request_id,247 thinking_enabled=resolved.thinking_enabled,248 ),249 )250 251 async def _stream_with_fallbacks(252 self,253 request: Request,254 candidates: list[ResolvedModel],255 request_data: MessagesRequest,256 ) -> AsyncIterator[str]:257 """Iterate through candidates until one succeeds or all fail."""258 last_exc: Exception | None = None259 260 for i, resolved in enumerate(candidates):261 try:262 # Pre-check: skip candidates that are currently rate limited or unhealthy263 from providers.rate_limit import GlobalRateLimiter264 265 limiter = GlobalRateLimiter.get_scoped_instance(resolved.provider_id)266 if limiter.is_blocked() and resolved.provider_id != "zen":267 # Silently skip — no failure penalty for temporary rate limit268 logger.debug(269 "Skipping blocked provider '{}' (no penalty)",270 resolved.provider_id,271 )272 continue273 274 # Check model health (recent failures)275 if not limiter.is_healthy(resolved.provider_model_ref):276 logger.warning(277 "Provider '{}' has recent failures, skipping to next candidate...",278 resolved.provider_model_ref,279 )280 last_exc = Exception("Recent failures")281 continue282 283 provider = self._provider_getter(resolved.provider_id)284 routed_request = request_data.model_copy(deep=True)285 routed_request.model = resolved.provider_model286 287 provider.preflight_stream(288 routed_request,289 thinking_enabled=resolved.thinking_enabled,290 )291 292 session_id = _get_session_id(request)293 self._session_tracker.track_request_sync(294 session_id, resolved.provider_id295 )296 297 request_id = f"req_{uuid.uuid4().hex[:12]}"298 logger.info(299 "API_REQUEST (auto fallback {}/{}): request_id={} provider={} model={}",300 i + 1,301 len(candidates),302 request_id,303 resolved.provider_id,304 resolved.provider_model,305 )306 307 input_tokens = self._token_counter(308 routed_request.messages, routed_request.system, routed_request.tools309 )310 311 # Attempt to stream from this provider.312 async for event in provider.stream_response(313 routed_request,314 input_tokens=input_tokens,315 request_id=request_id,316 thinking_enabled=resolved.thinking_enabled,317 ):318 yield event319 # CRITICAL: If we have yielded even one event, we have committed to this provider.320 # We must not fallback to another candidate mid-stream.321 return # Success, exit the fallback loop.322 323 except (RateLimitError, OverloadedError) as e:324 logger.warning(325 "Provider '{}' is rate limited or overloaded ({}). Trying next candidate...",326 resolved.provider_id,327 e.status_code,328 )329 limiter.record_failure(resolved.provider_model_ref)330 last_exc = e331 continue332 except TimeoutError as e:333 # Timeout = slow model, try next candidate for faster response334 logger.warning(335 "Provider '{}' timed out ({}). Trying next candidate...",336 resolved.provider_id,337 type(e).__name__,338 )339 limiter.record_failure(resolved.provider_model_ref)340 last_exc = e341 continue342 except Exception as e:343 # Check if it's a transient error that should trigger fallback344 error_str = str(e).lower()345 is_transient = any(346 kw in error_str347 for kw in [348 "timeout",349 "connection",350 "refused",351 "reset",352 "unavailable",353 "service",354 ]355 )356 if is_transient:357 logger.warning(358 "Provider '{}' failed with transient error ({}): {}. Trying next candidate...",359 resolved.provider_id,360 type(e).__name__,361 e,362 )363 limiter.record_failure(resolved.provider_model_ref)364 last_exc = e365 continue366 367 logger.error(368 "Provider '{}' failed with unexpected error: {}. Trying next candidate...",369 resolved.provider_id,370 e,371 )372 last_exc = e373 continue374 375 err_msg = str(last_exc) if last_exc else "No candidates succeeded"376 yield format_sse_event(377 "error",378 {379 "type": "error",380 "error": {381 "type": "api_error",382 "message": f"All fallback candidates failed: {err_msg}",383 },384 },385 )386 if last_exc:387 raise last_exc388 raise InvalidRequestError("No candidates succeeded")389 390 def count_tokens(self, request_data: TokenCountRequest) -> TokenCountResponse:391 """Count tokens for a request after applying configured model routing."""392 request_id = f"req_{uuid.uuid4().hex[:12]}"393 with logger.contextualize(request_id=request_id):394 try:395 _require_non_empty_messages(request_data.messages)396 routed = self._model_router.resolve_token_count_request(request_data)397 tokens = self._token_counter(398 routed.request.messages, routed.request.system, routed.request.tools399 )400 logger.info(401 "COUNT_TOKENS: request_id={} model={} messages={} input_tokens={}",402 request_id,403 routed.request.model,404 len(routed.request.messages),405 tokens,406 )407 return TokenCountResponse(input_tokens=tokens)408 except ProviderError:409 raise410 except Exception as e:411 _log_unexpected_service_exception(412 self._settings,413 e,414 context="COUNT_TOKENS_ERROR",415 request_id=request_id,416 )417 raise HTTPException(418 status_code=_http_status_for_unexpected_service_exception(e),419 detail=get_user_facing_error_message(e),420 ) from e421 