Coverage for agentos/queue/task_queue.py: 43%
149 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
1"""
2AgentOS v0.40 Task Queue — 异步任务调度与重试。
3支持:内存队列(开发)/ Redis队列(生产)、优先级、重试、死信队列。
4"""
6from __future__ import annotations
8import asyncio
9import heapq
10import time
11import uuid
12from collections.abc import Callable
13from dataclasses import dataclass, field
14from enum import Enum, StrEnum
15from typing import Any
18class TaskState(StrEnum):
19 """任务状态枚举。"""
21 PENDING = "pending"
22 RUNNING = "running"
23 SUCCESS = "success"
24 FAILED = "failed"
25 RETRYING = "retrying"
26 CANCELLED = "cancelled"
27 DEAD = "dead" # 死信
30class TaskPriority(int, Enum):
31 """任务优先级枚举。"""
33 LOW = 0
34 NORMAL = 50
35 HIGH = 100
36 CRITICAL = 200
39@dataclass(order=True)
40class QueuedTask:
41 """带优先级的任务节点(priority取负以实现最大堆)。"""
43 priority: int # -priority for max-heap
44 created_at: float = field(compare=False)
45 task: ScheduledTask = field(compare=False)
48@dataclass
49class ScheduledTask:
50 """调度任务。"""
52 id: str = field(default_factory=lambda: uuid.uuid4().hex[:12])
53 name: str = ""
54 payload: dict = field(default_factory=dict)
55 priority: TaskPriority = TaskPriority.NORMAL
56 state: TaskState = TaskState.PENDING
57 max_retries: int = 3
58 retry_delay: float = 1.0 # 秒
59 timeout: float = 60.0
60 callback: Callable | None = field(default=None, repr=False)
61 result: Any = None
62 error: str = ""
63 retry_count: int = 0
64 created_at: float = field(default_factory=time.time)
65 started_at: float = 0
66 completed_at: float = 0
67 _tags: dict = field(default_factory=dict)
70class MemoryQueue:
71 """基于堆内存的任务队列 — 开发环境默认。"""
73 def __init__(self, max_size: int = 10000):
74 self._heap: list[QueuedTask] = []
75 self._pending: dict[str, ScheduledTask] = {}
76 self._dead: list[ScheduledTask] = []
77 self.max_size = max_size
78 self._lock = asyncio.Lock()
80 async def enqueue(self, task: ScheduledTask) -> str:
81 async with self._lock:
82 if len(self._pending) >= self.max_size:
83 raise RuntimeError(f"Queue full ({self.max_size})")
84 self._pending[task.id] = task
85 heapq.heappush(
86 self._heap,
87 QueuedTask(priority=-task.priority.value, created_at=task.created_at, task=task),
88 )
89 return task.id
91 async def dequeue(self) -> ScheduledTask | None:
92 async with self._lock:
93 while self._heap:
94 qt = heapq.heappop(self._heap)
95 task = self._pending.pop(qt.task.id, None)
96 if task and task.state == TaskState.PENDING:
97 return task
98 return None
100 async def peek(self) -> ScheduledTask | None:
101 async with self._lock:
102 if self._heap:
103 return self._heap[0].task
104 return None
106 def pending_count(self) -> int:
107 return len(self._heap)
109 def dead_count(self) -> int:
110 return len(self._dead)
112 async def move_to_dead(self, task: ScheduledTask):
113 async with self._lock:
114 task.state = TaskState.DEAD
115 self._dead.append(task)
116 self._pending.pop(task.id, None)
118 def stats(self) -> dict:
119 return {"pending": len(self._heap), "dead": len(self._dead), "max_size": self.max_size}
122class TaskQueue:
123 """任务队列管理器。"""
125 def __init__(self, queue: MemoryQueue | None = None, concurrency: int = 4):
126 self._queue = queue or MemoryQueue()
127 self._concurrency = concurrency
128 self._running: set[str] = set()
129 self._callbacks: dict[str, Callable] = {}
130 self._semaphore = asyncio.Semaphore(concurrency)
131 self._running_flag = False
133 def register_callback(self, task_name: str, handler: Callable):
134 """注册任务处理器。"""
135 self._callbacks[task_name] = handler
137 async def submit(self, task: ScheduledTask) -> str:
138 if task.name not in self._callbacks:
139 raise ValueError(f"No handler registered for task: {task.name}")
140 task_id = await self._queue.enqueue(task)
141 return task_id
143 async def start(self):
144 """启动Worker循环。"""
145 self._running_flag = True
146 while self._running_flag:
147 task = await self._queue.dequeue()
148 if not task:
149 await asyncio.sleep(0.1)
150 continue
151 asyncio.create_task(self._execute(task))
153 def stop(self):
154 self._running_flag = False
156 async def _execute(self, task: ScheduledTask):
157 async with self._semaphore:
158 self._running.add(task.id)
159 task.state = TaskState.RUNNING
160 task.started_at = time.time()
162 handler = self._callbacks.get(task.name)
163 if not handler:
164 task.state = TaskState.FAILED
165 task.error = f"No handler: {task.name}"
166 self._running.discard(task.id)
167 return
169 try:
170 result = handler(task.payload)
171 if asyncio.iscoroutine(result):
172 result = await asyncio.wait_for(result, timeout=task.timeout)
173 task.result = result
174 task.state = TaskState.SUCCESS
175 except TimeoutError:
176 task.error = f"Timeout after {task.timeout}s"
177 await self._handle_failure(task)
178 except Exception as e:
179 task.error = str(e)
180 await self._handle_failure(task)
181 finally:
182 task.completed_at = time.time()
183 self._running.discard(task.id)
185 async def _handle_failure(self, task: ScheduledTask):
186 if task.retry_count < task.max_retries:
187 task.retry_count += 1
188 task.state = TaskState.RETRYING
189 await asyncio.sleep(task.retry_delay * (2 ** (task.retry_count - 1))) # 指数退避
190 task.state = TaskState.PENDING
191 task.priority = TaskPriority(task.priority.value + 10) # 提升优先级
192 await self._queue.enqueue(task)
193 else:
194 task.state = TaskState.FAILED
195 await self._queue.move_to_dead(task)
197 def cancel(self, task_id: str):
198 """取消任务。"""
199 task = self._queue._pending.get(task_id)
200 if task and task.state in (TaskState.PENDING, TaskState.RETRYING):
201 task.state = TaskState.CANCELLED
203 def stats(self) -> dict:
204 return {
205 "running": len(self._running),
206 "concurrency": self._concurrency,
207 "queue": self._queue.stats(),
208 "handlers": list(self._callbacks.keys()),
209 }