Team Ai
Apppublic

Yash030/claude-code-proxy

sourceHugging Faceupdated 5mo agoView on Hugging Face
2likes
services.py421 linesDownload Raw Back to api
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