Coverage for agentos/queue/task_queue.py: 43%

149 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-09 09:19 +0800

1""" 

2AgentOS v0.40 Task Queue — 异步任务调度与重试。 

3支持:内存队列(开发)/ Redis队列(生产)、优先级、重试、死信队列。 

4""" 

5 

6from __future__ import annotations 

7 

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 

16 

17 

18class TaskState(StrEnum): 

19 """任务状态枚举。""" 

20 

21 PENDING = "pending" 

22 RUNNING = "running" 

23 SUCCESS = "success" 

24 FAILED = "failed" 

25 RETRYING = "retrying" 

26 CANCELLED = "cancelled" 

27 DEAD = "dead" # 死信 

28 

29 

30class TaskPriority(int, Enum): 

31 """任务优先级枚举。""" 

32 

33 LOW = 0 

34 NORMAL = 50 

35 HIGH = 100 

36 CRITICAL = 200 

37 

38 

39@dataclass(order=True) 

40class QueuedTask: 

41 """带优先级的任务节点(priority取负以实现最大堆)。""" 

42 

43 priority: int # -priority for max-heap 

44 created_at: float = field(compare=False) 

45 task: ScheduledTask = field(compare=False) 

46 

47 

48@dataclass 

49class ScheduledTask: 

50 """调度任务。""" 

51 

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) 

68 

69 

70class MemoryQueue: 

71 """基于堆内存的任务队列 — 开发环境默认。""" 

72 

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() 

79 

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 

90 

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 

99 

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 

105 

106 def pending_count(self) -> int: 

107 return len(self._heap) 

108 

109 def dead_count(self) -> int: 

110 return len(self._dead) 

111 

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) 

117 

118 def stats(self) -> dict: 

119 return {"pending": len(self._heap), "dead": len(self._dead), "max_size": self.max_size} 

120 

121 

122class TaskQueue: 

123 """任务队列管理器。""" 

124 

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 

132 

133 def register_callback(self, task_name: str, handler: Callable): 

134 """注册任务处理器。""" 

135 self._callbacks[task_name] = handler 

136 

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 

142 

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)) 

152 

153 def stop(self): 

154 self._running_flag = False 

155 

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() 

161 

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 

168 

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) 

184 

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) 

196 

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 

202 

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 }