Blablablab/audio-classification
0
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 