Yash030/claude-code-proxy
2
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 