Team Ai
Apppublic

Yash030/claude-code-proxy

sourceHugging Faceupdated 5mo agoView on Hugging Face
2likes
runtime.py401 linesDownload Raw Back to api
1"""Application runtime composition and lifecycle ownership."""2 3from __future__ import annotations4 5import asyncio6import os7from contextlib import suppress8from dataclasses import dataclass, field9from typing import TYPE_CHECKING, Any10 11from fastapi import FastAPI12from loguru import logger13 14from config.settings import Settings, get_settings15from providers.exceptions import ServiceUnavailableError16from providers.registry import ProviderRegistry17 18if TYPE_CHECKING:19    from cli.manager import CLISessionManager20    from messaging.handler import ClaudeMessageHandler21    from messaging.platforms.base import MessagingPlatform22    from messaging.session import SessionStore23 24_SHUTDOWN_TIMEOUT_S = 5.025 26 27async def best_effort(28    name: str,29    awaitable: Any,30    timeout_s: float = _SHUTDOWN_TIMEOUT_S,31    *,32    log_verbose_errors: bool = False,33) -> None:34    """Run a shutdown step with timeout; never raise to callers."""35    try:36        await asyncio.wait_for(awaitable, timeout=timeout_s)37    except TimeoutError:38        logger.warning("Shutdown step timed out: {} ({}s)", name, timeout_s)39    except Exception as e:40        if log_verbose_errors:41            logger.warning(42                "Shutdown step failed: {}: {}: {}",43                name,44                type(e).__name__,45                e,46            )47        else:48            logger.warning(49                "Shutdown step failed: {}: exc_type={}",50                name,51                type(e).__name__,52            )53 54 55def warn_if_process_auth_token(settings: Settings) -> None:56    """Warn when server auth was implicitly inherited from the shell."""57    if settings.uses_process_anthropic_auth_token():58        logger.warning(59            "ANTHROPIC_AUTH_TOKEN is set in the process environment but not in "60            "a configured .env file. The proxy will require that token. Add "61            "ANTHROPIC_AUTH_TOKEN= to .env to disable proxy auth, or set the "62            "same token in .env to make server auth explicit."63        )64 65 66def log_startup_failure(settings: Settings, exc: Exception) -> None:67    """Log startup failures without traceback noise unless verbose diagnostics are enabled."""68    message = startup_failure_message(settings, exc)69    logger.error("Startup failed:\n{}", message)70 71 72def startup_failure_message(settings: Settings, exc: Exception) -> str:73    """Return a concise startup failure message for logs and ASGI lifespan failure."""74    if isinstance(exc, ServiceUnavailableError):75        return exc.message.strip() or "Server startup failed."76 77    if settings.log_api_error_tracebacks:78        return f"{type(exc).__name__}: {exc}"79 80    return f"Server startup failed: exc_type={type(exc).__name__}"81 82 83def _should_continue_after_model_validation_failure(exc: Exception) -> bool:84    """Return whether a model-validation failure should be downgraded to a warning.85 86    Provider discovery can fail transiently or due to local environment issues87    (for example, a missing runtime dependency in the provider's process path).88    We keep startup alive in those cases so the configured proxy can still serve89    requests and advertise the models that are already known from settings.90    """91    if not isinstance(exc, ServiceUnavailableError):92        return False93 94    message = (exc.message or str(exc)).lower()95    return "problem=query failure:" in message96 97 98@dataclass(slots=True)99class AppRuntime:100    """Own optional messaging, CLI, session, and provider runtime resources."""101 102    app: FastAPI103    settings: Settings104    _provider_registry: ProviderRegistry | None = field(default=None, init=False)105    messaging_platform: MessagingPlatform | None = None106    message_handler: ClaudeMessageHandler | None = None107    cli_manager: CLISessionManager | None = None108    _session_cleanup_task: asyncio.Task | None = field(default=None, init=False)109 110    @classmethod111    def for_app(112        cls,113        app: FastAPI,114        settings: Settings | None = None,115    ) -> AppRuntime:116        return cls(app=app, settings=settings or get_settings())117 118    async def startup(self) -> None:119        logger.info("Starting Claude Code Proxy...")120        self._provider_registry = ProviderRegistry()121        self.app.state.provider_registry = self._provider_registry122        try:123            warn_if_process_auth_token(self.settings)124            try:125                # Use a reasonable timeout for startup validation to prevent blocking healthy checks.126                await asyncio.wait_for(127                    self._provider_registry.validate_configured_models(self.settings),128                    timeout=15.0,129                )130            except Exception as exc:131                logger.warning(132                    "Startup model validation skipped or timed out: continuing in lazy mode. "133                    "Reason: {}",134                    str(exc) or type(exc).__name__,135                )136            self._provider_registry.start_model_list_refresh(self.settings)137            # Pre-warm provider connections on startup for faster first request138            await self._warmup_providers()139            await self._start_messaging_if_configured()140            # Start background session cleanup141            from core.session_tracker import SessionTracker142 143            self._session_cleanup_task = asyncio.create_task(144                SessionTracker.get_instance().start_cleanup_loop()145            )146            self._publish_state()147        except Exception as exc:148            log_startup_failure(self.settings, exc)149            await best_effort(150                "provider_registry.cleanup",151                self._provider_registry.cleanup(),152                log_verbose_errors=self.settings.log_api_error_tracebacks,153            )154            raise155 156    async def shutdown(self) -> None:157        verbose = self.settings.log_api_error_tracebacks158        # Cancel session cleanup task159        if self._session_cleanup_task is not None:160            self._session_cleanup_task.cancel()161            with suppress(asyncio.CancelledError, asyncio.TimeoutError):162                await asyncio.wait_for(self._session_cleanup_task, timeout=2.0)163        if self.message_handler is not None:164            try:165                self.message_handler.session_store.flush_pending_save()166            except Exception as e:167                if verbose:168                    logger.warning("Session store flush on shutdown: {}", e)169                else:170                    logger.warning(171                        "Session store flush on shutdown: exc_type={}",172                        type(e).__name__,173                    )174 175        logger.info("Shutdown requested, cleaning up...")176        if self.messaging_platform:177            await best_effort(178                "messaging_platform.stop",179                self.messaging_platform.stop(),180                log_verbose_errors=verbose,181            )182        if self.cli_manager:183            await best_effort(184                "cli_manager.stop_all",185                self.cli_manager.stop_all(),186                log_verbose_errors=verbose,187            )188        if self._provider_registry is not None:189            await best_effort(190                "provider_registry.cleanup",191                self._provider_registry.cleanup(),192                log_verbose_errors=verbose,193            )194        await self._shutdown_limiter()195        logger.info("Server shut down cleanly")196 197    async def _start_messaging_if_configured(self) -> None:198        try:199            from messaging.platforms.factory import (200                MessagingPlatformOptions,201                create_messaging_platform,202            )203 204            self.messaging_platform = create_messaging_platform(205                self.settings.messaging_platform,206                MessagingPlatformOptions(207                    telegram_bot_token=self.settings.telegram_bot_token,208                    allowed_telegram_user_id=self.settings.allowed_telegram_user_id,209                    discord_bot_token=self.settings.discord_bot_token,210                    allowed_discord_channels=self.settings.allowed_discord_channels,211                    voice_note_enabled=self.settings.voice_note_enabled,212                    whisper_model=self.settings.whisper_model,213                    whisper_device=self.settings.whisper_device,214                    hf_token=self.settings.hf_token,215                    nvidia_nim_api_key=self.settings.nvidia_nim_api_key_qwen,216                    messaging_rate_limit=self.settings.messaging_rate_limit,217                    messaging_rate_window=self.settings.messaging_rate_window,218                    log_raw_messaging_content=self.settings.log_raw_messaging_content,219                    log_api_error_tracebacks=self.settings.log_api_error_tracebacks,220                ),221            )222 223            if self.messaging_platform:224                await self._start_message_handler()225 226        except ImportError as e:227            if self.settings.log_api_error_tracebacks:228                logger.warning("Messaging module import error: {}", e)229            else:230                logger.warning(231                    "Messaging module import error: exc_type={}",232                    type(e).__name__,233                )234        except Exception as e:235            if self.settings.log_api_error_tracebacks:236                logger.error("Failed to start messaging platform: {}", e)237                import traceback238 239                logger.error(traceback.format_exc())240            else:241                logger.error(242                    "Failed to start messaging platform: exc_type={}",243                    type(e).__name__,244                )245 246    async def _start_message_handler(self) -> None:247        from cli.manager import CLISessionManager248        from messaging.handler import ClaudeMessageHandler249        from messaging.session import SessionStore250 251        workspace = (252            os.path.abspath(self.settings.allowed_dir)253            if self.settings.allowed_dir254            else os.getcwd()255        )256        os.makedirs(workspace, exist_ok=True)257 258        data_path = os.path.abspath(self.settings.claude_workspace)259        os.makedirs(data_path, exist_ok=True)260 261        api_url = f"http://{self.settings.host}:{self.settings.port}/v1"262        allowed_dirs = [workspace] if self.settings.allowed_dir else []263        plans_dir_abs = os.path.abspath(264            os.path.join(self.settings.claude_workspace, "plans")265        )266        plans_directory = os.path.relpath(plans_dir_abs, workspace)267        self.cli_manager = CLISessionManager(268            workspace_path=workspace,269            api_url=api_url,270            allowed_dirs=allowed_dirs,271            plans_directory=plans_directory,272            claude_bin=self.settings.claude_cli_bin,273            log_raw_cli_diagnostics=self.settings.log_raw_cli_diagnostics,274            log_messaging_error_details=self.settings.log_messaging_error_details,275        )276 277        session_store = SessionStore(278            storage_path=os.path.join(data_path, "sessions.json"),279            message_log_cap=self.settings.max_message_log_entries_per_chat,280        )281        platform = self.messaging_platform282        assert platform is not None283        self.message_handler = ClaudeMessageHandler(284            platform=platform,285            cli_manager=self.cli_manager,286            session_store=session_store,287            debug_platform_edits=self.settings.debug_platform_edits,288            debug_subagent_stack=self.settings.debug_subagent_stack,289            log_raw_messaging_content=self.settings.log_raw_messaging_content,290            log_raw_cli_diagnostics=self.settings.log_raw_cli_diagnostics,291            log_messaging_error_details=self.settings.log_messaging_error_details,292        )293        self._restore_tree_state(session_store)294 295        platform.on_message(self.message_handler.handle_message)296        await platform.start()297        logger.info(f"{platform.name} platform started with message handler")298 299    async def _warmup_providers(self) -> None:300        """Pre-establish HTTP connections to providers for faster first request."""301        logger.info("Warming up provider connections...")302        try:303            from api.dependencies import resolve_provider304 305            # Get all configured provider types306            provider_types = ["zen", "nvidia_nim"]307            warmup_tasks = []308            for provider_type in provider_types:309                try:310                    provider = resolve_provider(311                        provider_type, app=self.app, settings=self.settings312                    )313                    # Trigger lazy initialization by accessing client314                    if hasattr(provider, "_client"):315                        warmup_tasks.append(316                            self._warmup_provider(provider, provider_type)317                        )318                except Exception:319                    pass  # Skip if provider not configured320 321            if warmup_tasks:322                # Give connections a small window to establish323                await asyncio.wait_for(324                    asyncio.gather(*warmup_tasks, return_exceptions=True), timeout=5.0325                )326                logger.info("Provider warmup complete")327        except Exception as e:328            logger.warning("Provider warmup skipped: {}", e)329 330    async def _warmup_provider(self, provider, provider_type: str) -> None:331        """Force connection pool pre-warming on startup."""332        try:333            import httpx334 335            if hasattr(provider, "_http_client"):336                # Touch the connection pool to establish TCP+TLS connections337                http = provider._http_client338                await asyncio.wait_for(339                    http.get("/", timeout=httpx.Timeout(3.0, connect=2.0)),340                    timeout=4.0,341                )342                logger.debug("Provider {} HTTP pool warmed up", provider_type)343        except Exception:344            pass  # Warmup failures are non-fatal345 346    def _restore_tree_state(self, session_store: SessionStore) -> None:347        saved_trees = session_store.get_all_trees()348        if not saved_trees:349            return350        if self.message_handler is None:351            return352 353        logger.info(f"Restoring {len(saved_trees)} conversation trees...")354        from messaging.trees.queue_manager import TreeQueueManager355 356        self.message_handler.replace_tree_queue(357            TreeQueueManager.from_dict(358                {359                    "trees": saved_trees,360                    "node_to_tree": session_store.get_node_mapping(),361                },362                queue_update_callback=self.message_handler.update_queue_positions,363                node_started_callback=self.message_handler.mark_node_processing,364            )365        )366        if self.message_handler.tree_queue.cleanup_stale_nodes() > 0:367            tree_data = self.message_handler.tree_queue.to_dict()368            session_store.sync_from_tree_data(369                tree_data["trees"], tree_data["node_to_tree"]370            )371 372    def _publish_state(self) -> None:373        self.app.state.messaging_platform = self.messaging_platform374        self.app.state.message_handler = self.message_handler375        self.app.state.cli_manager = self.cli_manager376 377    async def _shutdown_limiter(self) -> None:378        verbose = self.settings.log_api_error_tracebacks379        try:380            from messaging.limiter import MessagingRateLimiter381        except Exception as e:382            if verbose:383                logger.debug(384                    "Rate limiter shutdown skipped (import failed): {}: {}",385                    type(e).__name__,386                    e,387                )388            else:389                logger.debug(390                    "Rate limiter shutdown skipped (import failed): exc_type={}",391                    type(e).__name__,392                )393            return394 395        await best_effort(396            "MessagingRateLimiter.shutdown_instance",397            MessagingRateLimiter.shutdown_instance(),398            timeout_s=2.0,399            log_verbose_errors=verbose,400        )401