Yash030/claude-code-proxy
2
1"""Session-aware request tracking for fair resource sharing across Claude Code instances."""2 3from __future__ import annotations4 5import asyncio6import time7from collections import defaultdict8from dataclasses import dataclass9from typing import ClassVar, cast10 11from loguru import logger12 13 14@dataclass(slots=True)15class SessionState:16 """State for a single session across all providers."""17 18 requests_in_window: int = 019 last_request_time: float = 0.020 total_requests: int = 021 22 23@dataclass(frozen=True, slots=True)24class ProviderLoad:25 """Load information for a single provider."""26 27 provider_id: str28 active_requests: int29 session_count: int30 requests_per_minute: float31 is_healthy: bool # Not rate limited32 33 34@dataclass(frozen=True, slots=True)35class SessionLoad:36 """Load information for a session across all providers."""37 38 session_id: str39 total_requests: int40 providers: dict[str, int] # provider_id -> request count41 42 43class SessionTracker:44 """45 Track request rates per session and per provider for fair resource sharing.46 47 This enables multiple Claude Code instances to share the proxy efficiently48 without one session starving others.49 """50 51 _instance: ClassVar[SessionTracker | None] = None52 53 def __init__(54 self,55 *,56 max_sessions: int = 50,57 window_seconds: float = 60.0,58 per_session_rate_limit: int = 30,59 retention_seconds: float | None = None,60 ):61 if hasattr(self, "_initialized"):62 return63 64 self._sessions: dict[str, SessionState] = {}65 self._session_requests: dict[str, dict[str, int]] = defaultdict(66 lambda: defaultdict(int)67 )68 self._provider_active: dict[str, int] = defaultdict(int)69 self._max_sessions = max_sessions70 self._window_seconds = window_seconds71 self._per_session_rate_limit = per_session_rate_limit72 self._retention_seconds = (73 retention_seconds if retention_seconds is not None else window_seconds * 274 )75 self._lock = asyncio.Lock()76 self._initialized = True77 78 logger.info(79 "SessionTracker initialized (max_sessions={}, window={}s, per_session_limit={}/min, retention={}s)",80 max_sessions,81 window_seconds,82 per_session_rate_limit,83 self._retention_seconds,84 )85 86 @classmethod87 def get_instance(cls, **kwargs) -> SessionTracker:88 """Get or create the singleton instance."""89 if cls._instance is None:90 cls._instance = cls(**kwargs)91 return cls._instance92 93 @classmethod94 def reset_instance(cls) -> None:95 """Reset singleton (for testing)."""96 cls._instance = None97 98 def _cleanup_old_sessions(self) -> None:99 """Remove sessions with no recent activity (must be called with lock held)."""100 now = time.monotonic()101 cutoff = now - self._retention_seconds102 to_remove = [103 sid104 for sid, state in self._sessions.items()105 if state.last_request_time < cutoff106 ]107 for sid in to_remove:108 del self._sessions[sid]109 if sid in self._session_requests:110 del self._session_requests[sid]111 if to_remove:112 logger.debug(113 "SessionTracker: cleaned up {} stale sessions ({} remaining)",114 len(to_remove),115 len(self._sessions),116 )117 118 async def start_cleanup_loop(self, interval: float = 60.0) -> None:119 """Background task: periodically clean up stale sessions."""120 while True:121 await asyncio.sleep(interval)122 async with self._lock:123 self._cleanup_old_sessions()124 125 def _evict_lru_session(self) -> None:126 """Evict least recently used session when at capacity."""127 if not self._sessions:128 return129 lru_sid = min(self._sessions.items(), key=lambda x: x[1].last_request_time)[0]130 del self._sessions[lru_sid]131 if lru_sid in self._session_requests:132 del self._session_requests[lru_sid]133 logger.warning("SessionTracker: Evicted LRU session '{}'", lru_sid)134 135 async def track_request(self, session_id: str, provider_id: str) -> None:136 """Record a request for a session to a provider (async-safe)."""137 self.track_request_sync(session_id, provider_id)138 139 def track_request_sync(self, session_id: str, provider_id: str) -> None:140 """Record a request for a session to a provider (sync version for hot path)."""141 # Hot path - no cleanup on every call, just update state142 # Cleanup runs periodically in background, not on every request143 if session_id not in self._sessions:144 if len(self._sessions) >= self._max_sessions:145 self._evict_lru_session()146 self._sessions[session_id] = SessionState()147 148 state = self._sessions[session_id]149 state.requests_in_window += 1150 state.last_request_time = time.monotonic()151 state.total_requests += 1152 153 self._session_requests[session_id][provider_id] += 1154 self._provider_active[provider_id] += 1155 156 async def track_request_async(self, session_id: str, provider_id: str) -> None:157 """Async version with lock for when called from async contexts that need guarantees."""158 async with self._lock:159 self.track_request_sync(session_id, provider_id)160 161 async def release_request(self, session_id: str, provider_id: str) -> None:162 """Release a request slot when streaming completes."""163 async with self._lock:164 self._provider_active[provider_id] = max(165 0, self._provider_active[provider_id] - 1166 )167 168 def get_provider_load(169 self, provider_id: str, blocked: bool = False170 ) -> ProviderLoad:171 """Get current load information for a provider."""172 session_count = sum(173 1174 for sid in self._sessions175 if self._session_requests[sid].get(provider_id, 0) > 0176 )177 total_requests = sum(178 self._session_requests[sid].get(provider_id, 0) for sid in self._sessions179 )180 181 return ProviderLoad(182 provider_id=provider_id,183 active_requests=self._provider_active.get(provider_id, 0),184 session_count=session_count,185 requests_per_minute=total_requests,186 is_healthy=not blocked,187 )188 189 def get_all_provider_loads(190 self, blocked_providers: set[str] | None = None191 ) -> dict[str, ProviderLoad]:192 """Get load information for all providers."""193 blocked = blocked_providers or set()194 all_providers = set(self._provider_active.keys())195 196 # Add providers from sessions even if not currently active197 for sid in self._session_requests:198 for provider_id in self._session_requests[sid]:199 all_providers.add(provider_id)200 201 return {202 pid: self.get_provider_load(pid, pid in blocked) for pid in all_providers203 }204 205 def get_session_load(self, session_id: str) -> SessionLoad | None:206 """Get load information for a specific session."""207 if session_id not in self._sessions:208 return None209 210 state = self._sessions[session_id]211 provider_counts = dict(self._session_requests[session_id])212 213 return SessionLoad(214 session_id=session_id,215 total_requests=state.total_requests,216 providers=provider_counts,217 )218 219 def get_all_session_loads(self) -> dict[str, SessionLoad]:220 """Get load information for all active sessions."""221 return {222 sid: cast(SessionLoad, self.get_session_load(sid))223 for sid in self._sessions224 if self.get_session_load(sid) is not None225 }226 227 async def check_session_allowed(self, session_id: str) -> tuple[bool, str]:228 """229 Check if a session is within its rate limit.230 231 Returns (allowed, reason) tuple.232 """233 async with self._lock:234 if session_id not in self._sessions:235 return True, "new session"236 237 state = self._sessions[session_id]238 if state.requests_in_window > self._per_session_rate_limit:239 return (240 False,241 f"rate limit exceeded ({state.requests_in_window}/{self._per_session_rate_limit}/min)",242 )243 244 return True, "ok"245 246 def get_healthy_provider_priority(247 self,248 candidates: list[str],249 blocked_providers: set[str] | None = None,250 ) -> list[str]:251 """252 Return candidates sorted by health/load priority.253 254 Healthy providers with lower load come first.255 """256 blocked = blocked_providers or set()257 return sorted(258 candidates,259 key=lambda pid: (260 pid in blocked, # Blocked providers go last261 self._provider_active.get(pid, 0), # Lower load first262 ),263 )264 265 def stats(self) -> dict:266 """Return current statistics."""267 return {268 "active_sessions": len(self._sessions),269 "total_providers": len(self._provider_active),270 "provider_active": dict(self._provider_active),271 }272 