Team Ai
Apppublic

Yash030/claude-code-proxy

sourceHugging Faceupdated 5mo agoView on Hugging Face
2likes
session_tracker.py272 linesDownload Raw Back to core
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