CoolFace
Apppublic

CHKIM79/scalable-ai-agent-system

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
execution_engine.py482 linesDownload Raw Back to execution
1"""2Execution Engine for Async Task Processing3Handles task control, retries, latency optimization, and concurrent execution scaling4"""5import asyncio6import logging7import time8import uuid9from typing import Dict, List, Any, Optional, Callable, Union10from dataclasses import dataclass, field11from enum import Enum12from datetime import datetime, timedelta13import json14 15 16class TaskStatus(Enum):17    PENDING = "pending"18    RUNNING = "running"19    COMPLETED = "completed"20    FAILED = "failed"21    CANCELLED = "cancelled"22    RETRYING = "retrying"23 24 25class TaskPriority(Enum):26    LOW = 127    NORMAL = 228    HIGH = 329    CRITICAL = 430 31 32@dataclass33class ExecutionTask:34    id: str = field(default_factory=lambda: str(uuid.uuid4()))35    name: str = ""36    function: Optional[Callable] = None37    args: tuple = field(default_factory=tuple)38    kwargs: Dict[str, Any] = field(default_factory=dict)39    priority: TaskPriority = TaskPriority.NORMAL40    status: TaskStatus = TaskStatus.PENDING41    created_at: datetime = field(default_factory=datetime.now)42    started_at: Optional[datetime] = None43    completed_at: Optional[datetime] = None44    result: Any = None45    error: Optional[str] = None46    retry_count: int = 047    max_retries: int = 348    timeout: float = 30.049    dependencies: List[str] = field(default_factory=list)50    metadata: Dict[str, Any] = field(default_factory=dict)51 52 53@dataclass54class ExecutionMetrics:55    total_tasks: int = 056    completed_tasks: int = 057    failed_tasks: int = 058    average_execution_time: float = 0.059    current_load: int = 060    max_concurrent_tasks: int = 1061    queue_size: int = 062 63 64class TaskQueue:65    """Priority-based task queue with dependency management"""66    67    def __init__(self):68        self.queue: asyncio.PriorityQueue = asyncio.PriorityQueue()69        self.tasks: Dict[str, ExecutionTask] = {}70        self.completed_tasks: Dict[str, ExecutionTask] = {}71    72    async def add_task(self, task: ExecutionTask):73        """Add task to queue"""74        self.tasks[task.id] = task75        # Use negative priority for correct ordering (higher priority first)76        await self.queue.put((-task.priority.value, task.created_at, task))77    78    async def get_next_task(self) -> Optional[ExecutionTask]:79        """Get next available task considering dependencies"""80        if self.queue.empty():81            return None82        83        # Get all available tasks84        available_tasks = []85        temp_tasks = []86        87        while not self.queue.empty():88            priority, created_at, task = await self.queue.get()89            temp_tasks.append((priority, created_at, task))90            91            # Check if dependencies are satisfied92            if self._dependencies_satisfied(task):93                available_tasks.append(task)94                # Remove from temp_tasks since we're using this one95                temp_tasks.pop()96                break97        98        # Put back unused tasks99        for priority, created_at, task in temp_tasks:100            await self.queue.put((priority, created_at, task))101        102        return available_tasks[0] if available_tasks else None103    104    def _dependencies_satisfied(self, task: ExecutionTask) -> bool:105        """Check if task dependencies are satisfied"""106        for dep_id in task.dependencies:107            if dep_id not in self.completed_tasks:108                return False109            if self.completed_tasks[dep_id].status != TaskStatus.COMPLETED:110                return False111        return True112    113    def mark_completed(self, task: ExecutionTask):114        """Mark task as completed"""115        self.completed_tasks[task.id] = task116        if task.id in self.tasks:117            del self.tasks[task.id]118    119    def get_queue_size(self) -> int:120        """Get current queue size"""121        return self.queue.qsize()122 123 124class RetryManager:125    """Manages task retry logic with exponential backoff"""126    127    def __init__(self):128        self.retry_delays = {}129    130    def should_retry(self, task: ExecutionTask) -> bool:131        """Determine if task should be retried"""132        return task.retry_count < task.max_retries133    134    def get_retry_delay(self, task: ExecutionTask) -> float:135        """Calculate retry delay with exponential backoff"""136        base_delay = 1.0137        max_delay = 60.0138        139        delay = min(base_delay * (2 ** task.retry_count), max_delay)140        return delay141    142    async def schedule_retry(self, task: ExecutionTask, queue: TaskQueue):143        """Schedule task for retry"""144        if self.should_retry(task):145            task.retry_count += 1146            task.status = TaskStatus.RETRYING147            148            delay = self.get_retry_delay(task)149            await asyncio.sleep(delay)150            151            task.status = TaskStatus.PENDING152            await queue.add_task(task)153            return True154        return False155 156 157class LoadBalancer:158    """Manages concurrent execution load"""159    160    def __init__(self, max_concurrent_tasks: int = 10):161        self.max_concurrent_tasks = max_concurrent_tasks162        self.current_tasks: Dict[str, ExecutionTask] = {}163        self.semaphore = asyncio.Semaphore(max_concurrent_tasks)164    165    async def acquire(self, task: ExecutionTask) -> bool:166        """Acquire execution slot"""167        acquired = await self.semaphore.acquire()168        if acquired:169            self.current_tasks[task.id] = task170        return acquired171    172    def release(self, task: ExecutionTask):173        """Release execution slot"""174        if task.id in self.current_tasks:175            del self.current_tasks[task.id]176        self.semaphore.release()177    178    def get_current_load(self) -> int:179        """Get current load"""180        return len(self.current_tasks)181    182    def adjust_capacity(self, new_capacity: int):183        """Dynamically adjust capacity"""184        if new_capacity > self.max_concurrent_tasks:185            # Increase capacity186            for _ in range(new_capacity - self.max_concurrent_tasks):187                self.semaphore.release()188        elif new_capacity < self.max_concurrent_tasks:189            # Decrease capacity (acquire extra permits)190            for _ in range(self.max_concurrent_tasks - new_capacity):191                asyncio.create_task(self.semaphore.acquire())192        193        self.max_concurrent_tasks = new_capacity194 195 196class ExecutionEngine:197    """Main execution engine for async task processing"""198    199    def __init__(self, max_concurrent_tasks: int = 10):200        self.task_queue = TaskQueue()201        self.retry_manager = RetryManager()202        self.load_balancer = LoadBalancer(max_concurrent_tasks)203        self.metrics = ExecutionMetrics(max_concurrent_tasks=max_concurrent_tasks)204        205        self.running = False206        self.worker_tasks: List[asyncio.Task] = []207        self.execution_times: List[float] = []208        209        self.logger = logging.getLogger(__name__)210    211    async def initialize(self):212        """Initialize execution engine"""213        self.running = True214        215        # Start worker tasks216        for i in range(3):  # Start with 3 workers217            worker = asyncio.create_task(self._worker(f"worker-{i}"))218            self.worker_tasks.append(worker)219        220        # Start metrics updater221        asyncio.create_task(self._update_metrics())222        223        self.logger.info("Execution engine initialized")224    225    async def submit_task_dict(self, task_data: Dict[str, Any]) -> Dict[str, Any]:226        """Submit task from dictionary data (for testing)"""227        async def dummy_task():228            await asyncio.sleep(0.1)229            return {"result": "Task completed", "data": task_data}230        231        try:232            task_id = await self.submit_task(233                function=dummy_task,234                name=task_data.get('id', 'test_task')235            )236            return {"success": True, "task_id": task_id}237        except Exception as e:238            return {"success": False, "error": str(e)}239    240    async def submit_task(self, 241                         function: Callable,242                         args: tuple = (),243                         kwargs: Dict[str, Any] = None,244                         name: str = "",245                         priority: TaskPriority = TaskPriority.NORMAL,246                         timeout: float = 30.0,247                         max_retries: int = 3,248                         dependencies: List[str] = None) -> str:249        """Submit task for execution"""250        251        task = ExecutionTask(252            name=name or function.__name__,253            function=function,254            args=args,255            kwargs=kwargs or {},256            priority=priority,257            timeout=timeout,258            max_retries=max_retries,259            dependencies=dependencies or []260        )261        262        await self.task_queue.add_task(task)263        self.metrics.total_tasks += 1264        265        self.logger.info(f"Task submitted: {task.name} ({task.id})")266        return task.id267    268    async def get_task_result(self, task_id: str, timeout: float = None) -> Any:269        """Wait for task completion and return result"""270        start_time = time.time()271        272        while True:273            # Check if task is completed274            if task_id in self.task_queue.completed_tasks:275                task = self.task_queue.completed_tasks[task_id]276                if task.status == TaskStatus.COMPLETED:277                    return task.result278                elif task.status == TaskStatus.FAILED:279                    raise Exception(f"Task failed: {task.error}")280            281            # Check timeout282            if timeout and (time.time() - start_time) > timeout:283                raise asyncio.TimeoutError(f"Task {task_id} did not complete within {timeout}s")284            285            await asyncio.sleep(0.1)286    287    async def cancel_task(self, task_id: str) -> bool:288        """Cancel a pending or running task"""289        # Check if task is in queue290        if task_id in self.task_queue.tasks:291            task = self.task_queue.tasks[task_id]292            task.status = TaskStatus.CANCELLED293            return True294        295        # Check if task is running296        if task_id in self.load_balancer.current_tasks:297            task = self.load_balancer.current_tasks[task_id]298            task.status = TaskStatus.CANCELLED299            return True300        301        return False302    303    async def _worker(self, worker_name: str):304        """Worker coroutine that processes tasks"""305        self.logger.info(f"Worker {worker_name} started")306        307        while self.running:308            try:309                # Get next task310                task = await self.task_queue.get_next_task()311                if not task:312                    await asyncio.sleep(0.1)313                    continue314                315                # Acquire execution slot316                await self.load_balancer.acquire(task)317                318                try:319                    await self._execute_task(task)320                finally:321                    self.load_balancer.release(task)322                    323            except Exception as e:324                self.logger.error(f"Worker {worker_name} error: {e}")325                await asyncio.sleep(1)326        327        self.logger.info(f"Worker {worker_name} stopped")328    329    async def _execute_task(self, task: ExecutionTask):330        """Execute a single task"""331        task.status = TaskStatus.RUNNING332        task.started_at = datetime.now()333        334        self.logger.info(f"Executing task: {task.name} ({task.id})")335        336        try:337            # Execute with timeout338            if asyncio.iscoroutinefunction(task.function):339                result = await asyncio.wait_for(340                    task.function(*task.args, **task.kwargs),341                    timeout=task.timeout342                )343            else:344                # Run sync function in thread pool345                result = await asyncio.get_event_loop().run_in_executor(346                    None, 347                    lambda: task.function(*task.args, **task.kwargs)348                )349            350            # Task completed successfully351            task.result = result352            task.status = TaskStatus.COMPLETED353            task.completed_at = datetime.now()354            355            # Update metrics356            execution_time = (task.completed_at - task.started_at).total_seconds()357            self.execution_times.append(execution_time)358            self.metrics.completed_tasks += 1359            360            self.task_queue.mark_completed(task)361            362            self.logger.info(f"Task completed: {task.name} ({task.id}) in {execution_time:.2f}s")363            364        except asyncio.TimeoutError:365            task.error = f"Task timed out after {task.timeout}s"366            task.status = TaskStatus.FAILED367            await self._handle_task_failure(task)368            369        except Exception as e:370            task.error = str(e)371            task.status = TaskStatus.FAILED372            await self._handle_task_failure(task)373    374    async def _handle_task_failure(self, task: ExecutionTask):375        """Handle task failure and retry logic"""376        self.logger.error(f"Task failed: {task.name} ({task.id}) - {task.error}")377        378        # Try to retry379        if await self.retry_manager.schedule_retry(task, self.task_queue):380            self.logger.info(f"Task scheduled for retry: {task.name} (attempt {task.retry_count})")381        else:382            # No more retries383            task.completed_at = datetime.now()384            self.metrics.failed_tasks += 1385            self.task_queue.mark_completed(task)386            self.logger.error(f"Task permanently failed: {task.name} ({task.id})")387    388    async def _update_metrics(self):389        """Periodically update execution metrics"""390        while self.running:391            try:392                # Update current load393                self.metrics.current_load = self.load_balancer.get_current_load()394                self.metrics.queue_size = self.task_queue.get_queue_size()395                396                # Update average execution time397                if self.execution_times:398                    self.metrics.average_execution_time = sum(self.execution_times) / len(self.execution_times)399                    400                    # Keep only recent execution times (last 100)401                    if len(self.execution_times) > 100:402                        self.execution_times = self.execution_times[-100:]403                404                # Auto-scale based on load405                await self._auto_scale()406                407                await asyncio.sleep(5)  # Update every 5 seconds408                409            except Exception as e:410                self.logger.error(f"Metrics update error: {e}")411                await asyncio.sleep(5)412    413    async def _auto_scale(self):414        """Auto-scale execution capacity based on load"""415        current_load = self.metrics.current_load416        queue_size = self.metrics.queue_size417        max_capacity = self.metrics.max_concurrent_tasks418        419        # Scale up if queue is building up420        if queue_size > 10 and current_load >= max_capacity * 0.8:421            new_capacity = min(max_capacity + 2, 20)  # Max 20 concurrent tasks422            self.load_balancer.adjust_capacity(new_capacity)423            self.metrics.max_concurrent_tasks = new_capacity424            self.logger.info(f"Scaled up to {new_capacity} concurrent tasks")425        426        # Scale down if load is consistently low427        elif queue_size == 0 and current_load < max_capacity * 0.3 and max_capacity > 5:428            new_capacity = max(max_capacity - 1, 5)  # Min 5 concurrent tasks429            self.load_balancer.adjust_capacity(new_capacity)430            self.metrics.max_concurrent_tasks = new_capacity431            self.logger.info(f"Scaled down to {new_capacity} concurrent tasks")432    433    def get_metrics(self) -> Dict[str, Any]:434        """Get current execution metrics"""435        return {436            "total_tasks": self.metrics.total_tasks,437            "completed_tasks": self.metrics.completed_tasks,438            "failed_tasks": self.metrics.failed_tasks,439            "success_rate": self.metrics.completed_tasks / max(self.metrics.total_tasks, 1),440            "average_execution_time": self.metrics.average_execution_time,441            "current_load": self.metrics.current_load,442            "max_concurrent_tasks": self.metrics.max_concurrent_tasks,443            "queue_size": self.metrics.queue_size,444            "active_workers": len(self.worker_tasks)445        }446    447    def get_task_status(self, task_id: str) -> Optional[Dict[str, Any]]:448        """Get status of a specific task"""449        # Check active tasks450        if task_id in self.task_queue.tasks:451            task = self.task_queue.tasks[task_id]452        elif task_id in self.task_queue.completed_tasks:453            task = self.task_queue.completed_tasks[task_id]454        else:455            return None456        457        return {458            "id": task.id,459            "name": task.name,460            "status": task.status.value,461            "priority": task.priority.value,462            "created_at": task.created_at.isoformat(),463            "started_at": task.started_at.isoformat() if task.started_at else None,464            "completed_at": task.completed_at.isoformat() if task.completed_at else None,465            "retry_count": task.retry_count,466            "error": task.error467        }468    469    async def shutdown(self):470        """Gracefully shutdown execution engine"""471        self.logger.info("Shutting down execution engine...")472        self.running = False473        474        # Cancel all worker tasks475        for worker in self.worker_tasks:476            worker.cancel()477        478        # Wait for workers to finish479        await asyncio.gather(*self.worker_tasks, return_exceptions=True)480        481        self.logger.info("Execution engine shutdown complete")482