Jack1808/Claude_Code
0
1"""Dependency injection for FastAPI."""2 3from urllib.parse import unquote_plus4 5from fastapi import Depends, HTTPException, Request6from loguru import logger7 8from config.settings import Settings9from config.settings import get_settings as _get_settings10from providers.base import BaseProvider, ProviderConfig11from providers.common import get_user_facing_error_message12from providers.exceptions import AuthenticationError13from providers.llamacpp import LlamaCppProvider14from providers.lmstudio import LMStudioProvider15from providers.nvidia_nim import NVIDIA_NIM_BASE_URL, NvidiaNimProvider16from providers.open_router import OPENROUTER_BASE_URL, OpenRouterProvider17 18# Provider registry: keyed by provider type string, lazily populated19_providers: dict[str, BaseProvider] = {}20 21 22def get_settings() -> Settings:23 """Get application settings via dependency injection."""24 return _get_settings()25 26 27def _create_provider_for_type(provider_type: str, settings: Settings) -> BaseProvider:28 """Construct and return a new provider instance for the given provider type."""29 if provider_type == "nvidia_nim":30 if not settings.nvidia_nim_api_key or not settings.nvidia_nim_api_key.strip():31 raise AuthenticationError(32 "NVIDIA_NIM_API_KEY is not set. Add it to your .env file. "33 "Get a key at https://build.nvidia.com/settings/api-keys"34 )35 config = ProviderConfig(36 api_key=settings.nvidia_nim_api_key,37 base_url=NVIDIA_NIM_BASE_URL,38 rate_limit=settings.provider_rate_limit,39 rate_window=settings.provider_rate_window,40 max_concurrency=settings.provider_max_concurrency,41 http_read_timeout=settings.http_read_timeout,42 http_write_timeout=settings.http_write_timeout,43 http_connect_timeout=settings.http_connect_timeout,44 )45 return NvidiaNimProvider(config, nim_settings=settings.nim)46 if provider_type == "open_router":47 if not settings.open_router_api_key or not settings.open_router_api_key.strip():48 raise AuthenticationError(49 "OPENROUTER_API_KEY is not set. Add it to your .env file. "50 "Get a key at https://openrouter.ai/keys"51 )52 config = ProviderConfig(53 api_key=settings.open_router_api_key,54 base_url=OPENROUTER_BASE_URL,55 rate_limit=settings.provider_rate_limit,56 rate_window=settings.provider_rate_window,57 max_concurrency=settings.provider_max_concurrency,58 http_read_timeout=settings.http_read_timeout,59 http_write_timeout=settings.http_write_timeout,60 http_connect_timeout=settings.http_connect_timeout,61 )62 return OpenRouterProvider(config)63 if provider_type == "lmstudio":64 config = ProviderConfig(65 api_key="lm-studio",66 base_url=settings.lm_studio_base_url,67 rate_limit=settings.provider_rate_limit,68 rate_window=settings.provider_rate_window,69 max_concurrency=settings.provider_max_concurrency,70 http_read_timeout=settings.http_read_timeout,71 http_write_timeout=settings.http_write_timeout,72 http_connect_timeout=settings.http_connect_timeout,73 )74 return LMStudioProvider(config)75 if provider_type == "llamacpp":76 config = ProviderConfig(77 api_key="llamacpp",78 base_url=settings.llamacpp_base_url,79 rate_limit=settings.provider_rate_limit,80 rate_window=settings.provider_rate_window,81 max_concurrency=settings.provider_max_concurrency,82 http_read_timeout=settings.http_read_timeout,83 http_write_timeout=settings.http_write_timeout,84 http_connect_timeout=settings.http_connect_timeout,85 )86 return LlamaCppProvider(config)87 logger.error(88 "Unknown provider_type: '{}'. Supported: 'nvidia_nim', 'open_router', 'lmstudio', 'llamacpp'",89 provider_type,90 )91 raise ValueError(92 f"Unknown provider_type: '{provider_type}'. "93 f"Supported: 'nvidia_nim', 'open_router', 'lmstudio', 'llamacpp'"94 )95 96 97def get_provider_for_type(provider_type: str) -> BaseProvider:98 """Get or create a provider for the given provider type.99 100 Providers are cached in the registry and reused across requests.101 """102 if provider_type not in _providers:103 try:104 _providers[provider_type] = _create_provider_for_type(105 provider_type, get_settings()106 )107 except AuthenticationError as e:108 raise HTTPException(109 status_code=503, detail=get_user_facing_error_message(e)110 ) from e111 logger.info("Provider initialized: {}", provider_type)112 return _providers[provider_type]113 114 115def validate_request_api_key(request: Request, settings: Settings) -> None:116 """Validate a request against configured server API key.117 118 Checks `x-api-key` header, `Authorization: Bearer ...`, or query parameter `psw`119 against `Settings.anthropic_auth_token`. If `ANTHROPIC_AUTH_TOKEN` is empty, this is a no-op.120 121 Supports Hugging Face Spaces private deployments via query parameter authentication:122 - Append `?psw=your-token` to the base URL123 - Or `?psw:your-token` (URL-encoded colon becomes %3A)124 """125 anthropic_auth_token = getattr(settings, "anthropic_auth_token", None)126 if not anthropic_auth_token:127 # No API key configured -> allow128 return129 130 # Keep Space health/app shell reachable even when API auth is enabled.131 if _is_public_probe_request(request):132 return133 134 # Allow Hugging Face private Space signed browser requests for UI pages.135 # This keeps API routes protected while avoiding 401 on Space shell probes.136 if _is_hf_signed_page_request(request):137 return138 139 token = None140 141 # Check headers first (preferred)142 header = (143 request.headers.get("x-api-key")144 or request.headers.get("authorization")145 or request.headers.get("anthropic-auth-token")146 )147 if header:148 # Support both raw key in X-API-Key and Bearer token in Authorization149 token = header150 if header.lower().startswith("bearer "):151 token = header.split(" ", 1)[1]152 # Strip anything after the first colon to handle tokens with appended model names153 if token and ":" in token:154 token = token.split(":", 1)[0]155 else:156 token = _extract_query_token(request)157 158 if not token:159 raise HTTPException(status_code=401, detail="Missing API key")160 161 if token != anthropic_auth_token:162 raise HTTPException(status_code=401, detail="Invalid API key")163 164 165def _extract_query_token(request: Request) -> str | None:166 """Extract auth token from query string for private proxy deployments."""167 query_params = request.query_params168 if "psw" in query_params:169 token = query_params["psw"]170 if token and ":" in token:171 return token.split(":", 1)[0]172 return token or None173 174 raw_query_bytes = request.scope.get("query_string", b"")175 raw_query = raw_query_bytes.decode("utf-8", errors="ignore")176 if not raw_query:177 return None178 179 for part in raw_query.split("&"):180 if part.startswith("psw:"):181 token = unquote_plus(part[len("psw:") :])182 if token and ":" in token:183 return token.split(":", 1)[0]184 return token or None185 if part.startswith("psw%3A") or part.startswith("psw%3a"):186 token = unquote_plus(part[len("psw%3A") :])187 if token and ":" in token:188 return token.split(":", 1)[0]189 return token or None190 191 return None192 193 194def _is_hf_signed_page_request(request: Request) -> bool:195 """Return True for Hugging Face signed browser requests to non-API pages."""196 if request.method not in {"GET", "HEAD"}:197 return False198 199 if request.url.path.startswith("/v1/"):200 return False201 202 if "__sign" not in request.query_params:203 return False204 205 accept = request.headers.get("accept", "").lower()206 return "text/html" in accept or "*/*" in accept207 208 209def _is_public_probe_request(request: Request) -> bool:210 """Return True for read-only public endpoints used by app shells/health checks."""211 if request.method not in {"GET", "HEAD"}:212 return False213 return request.url.path in {"/", "/health"}214 215 216def require_api_key(217 request: Request, settings: Settings = Depends(get_settings)218) -> None:219 """FastAPI dependency wrapper for API key validation."""220 validate_request_api_key(request, settings)221 222 223def get_provider() -> BaseProvider:224 """Get or create the default provider (based on MODEL env var).225 226 Backward-compatible convenience for health/root endpoints and tests.227 """228 return get_provider_for_type(get_settings().provider_type)229 230 231async def cleanup_provider():232 """Cleanup all provider resources."""233 global _providers234 for provider in _providers.values():235 await provider.cleanup()236 _providers = {}237 logger.debug("Provider cleanup completed")238 