CoolFace
Apppublic

Blablablab/audio-classification

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes
agent_runner_manager.py227 linesDownload Raw Back to potato
1"""2Agent Runner Session Manager3 4Singleton that manages active AgentRunner sessions.5Keyed by "{user_id}:{instance_id}" for per-user, per-instance isolation.6Includes TTL-based cleanup and max concurrent session limits.7"""8 9import atexit10import logging11import threading12import time13from typing import Dict, Optional14 15from potato.agent_runner import AgentConfig, AgentRunner, AgentState16 17logger = logging.getLogger(__name__)18 19# Default limits20DEFAULT_MAX_SESSIONS = 1021DEFAULT_SESSION_TTL = 3600  # 1 hour22 23 24class AgentRunnerManager:25    """26    Manages active AgentRunner sessions with lifecycle control.27 28    Thread-safe singleton. Sessions are keyed by "{user_id}:{instance_id}".29    """30 31    _instance = None32    _lock = threading.Lock()33 34    def __init__(35        self,36        max_sessions: int = DEFAULT_MAX_SESSIONS,37        session_ttl: int = DEFAULT_SESSION_TTL,38    ):39        self._sessions: Dict[str, AgentRunner] = {}40        self._session_created: Dict[str, float] = {}41        self._session_meta: Dict[str, Dict] = {}42        self._lock = threading.Lock()43        self.max_sessions = max_sessions44        self.session_ttl = session_ttl45 46        # Start cleanup thread47        self._cleanup_stop = threading.Event()48        self._cleanup_thread = threading.Thread(49            target=self._cleanup_loop, daemon=True, name="agent-cleanup"50        )51        self._cleanup_thread.start()52 53    @classmethod54    def get_instance(cls, **kwargs) -> "AgentRunnerManager":55        """Get or create the singleton instance."""56        if cls._instance is None:57            with cls._lock:58                if cls._instance is None:59                    cls._instance = cls(**kwargs)60        return cls._instance61 62    @classmethod63    def clear_instance(cls):64        """Clear the singleton (for testing)."""65        with cls._lock:66            if cls._instance is not None:67                cls._instance.shutdown()68                cls._instance = None69 70    def create_session(71        self,72        user_id: str,73        instance_id: str,74        config: AgentConfig,75        screenshot_dir: str,76    ) -> AgentRunner:77        """78        Create a new agent session.79 80        Args:81            user_id: Annotator user ID82            instance_id: Annotation instance ID83            config: Agent configuration84            screenshot_dir: Directory to store screenshots85 86        Returns:87            AgentRunner instance88 89        Raises:90            RuntimeError: If max sessions reached or session already exists91        """92        session_key = f"{user_id}:{instance_id}"93 94        with self._lock:95            # Clean up expired sessions first96            self._cleanup_expired_locked()97 98            # Check for existing active session99            if session_key in self._sessions:100                existing = self._sessions[session_key]101                if existing.state in (AgentState.RUNNING, AgentState.PAUSED, AgentState.TAKEOVER):102                    raise RuntimeError(103                        f"Active session already exists for {session_key}. "104                        f"Stop it first."105                    )106                # Old completed/error session — remove it107                del self._sessions[session_key]108                del self._session_created[session_key]109                if session_key in self._session_meta:110                    del self._session_meta[session_key]111 112            # Check capacity113            active_count = sum(114                1115                for s in self._sessions.values()116                if s.state in (AgentState.RUNNING, AgentState.PAUSED, AgentState.TAKEOVER)117            )118            if active_count >= self.max_sessions:119                raise RuntimeError(120                    f"Maximum concurrent sessions ({self.max_sessions}) reached"121                )122 123            import uuid124            session_id = str(uuid.uuid4())[:12]125            runner = AgentRunner(session_id, config, screenshot_dir)126 127            self._sessions[session_key] = runner128            self._session_created[session_key] = time.time()129            self._session_meta[session_key] = {130                "user_id": user_id,131                "instance_id": instance_id,132                "session_id": session_id,133            }134 135            logger.info(136                f"Created agent session {session_id} for {session_key}"137            )138            return runner139 140    def get_session(self, session_id: str) -> Optional[AgentRunner]:141        """Get a session by its session_id."""142        with self._lock:143            for runner in self._sessions.values():144                if runner.session_id == session_id:145                    return runner146        return None147 148    def get_session_by_key(self, user_id: str, instance_id: str) -> Optional[AgentRunner]:149        """Get a session by user_id and instance_id."""150        session_key = f"{user_id}:{instance_id}"151        with self._lock:152            return self._sessions.get(session_key)153 154    def remove_session(self, session_id: str):155        """Remove a session by session_id."""156        with self._lock:157            key_to_remove = None158            for key, runner in self._sessions.items():159                if runner.session_id == session_id:160                    key_to_remove = key161                    break162            if key_to_remove:163                runner = self._sessions.pop(key_to_remove)164                self._session_created.pop(key_to_remove, None)165                self._session_meta.pop(key_to_remove, None)166                runner.stop()167                logger.info(f"Removed agent session {session_id}")168 169    def list_sessions(self) -> list:170        """List all active sessions."""171        with self._lock:172            result = []173            for key, runner in self._sessions.items():174                meta = self._session_meta.get(key, {})175                result.append({176                    "session_id": runner.session_id,177                    "user_id": meta.get("user_id"),178                    "instance_id": meta.get("instance_id"),179                    "state": runner.state.value,180                    "step_count": runner.step_count,181                    "created": self._session_created.get(key),182                })183            return result184 185    def _cleanup_expired_locked(self):186        """Remove expired sessions. Must be called with self._lock held."""187        now = time.time()188        expired_keys = []189        for key, created_at in self._session_created.items():190            if now - created_at > self.session_ttl:191                runner = self._sessions.get(key)192                if runner and runner.state in (AgentState.COMPLETED, AgentState.ERROR, AgentState.IDLE):193                    expired_keys.append(key)194                elif runner and now - created_at > self.session_ttl * 2:195                    # Force-stop sessions that have been running too long196                    runner.stop()197                    expired_keys.append(key)198 199        for key in expired_keys:200            self._sessions.pop(key, None)201            self._session_created.pop(key, None)202            self._session_meta.pop(key, None)203            logger.info(f"Cleaned up expired session: {key}")204 205    def _cleanup_loop(self):206        """Background cleanup thread."""207        while not self._cleanup_stop.is_set():208            self._cleanup_stop.wait(60)  # Check every 60 seconds209            if self._cleanup_stop.is_set():210                break211            with self._lock:212                self._cleanup_expired_locked()213 214    def shutdown(self):215        """Stop all sessions and cleanup thread."""216        self._cleanup_stop.set()217        with self._lock:218            for key, runner in self._sessions.items():219                try:220                    runner.stop()221                except Exception as e:222                    logger.warning(f"Error stopping session {key}: {e}")223            self._sessions.clear()224            self._session_created.clear()225            self._session_meta.clear()226        logger.info("AgentRunnerManager shut down")227