Coverage for agentos/tests/test_a2a.py: 100%

236 statements  

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

1"""测试 A2A 协议 — Task, Message, Handoff, Client, Server。""" 

2 

3import time 

4 

5import pytest 

6 

7from agentos.protocols.a2a import ( 

8 A2AArtifact, 

9 A2AHandoff, 

10 A2AMessage, 

11 A2AServer, 

12 A2ASession, 

13 A2ATask, 

14 DataPart, 

15 FilePart, 

16 MessageRole, 

17 TaskState, 

18 TextPart, 

19 new_handoff, 

20 new_task, 

21 part_from_dict, 

22) 

23 

24 

25class TestA2AParts: 

26 def test_text_part_roundtrip(self): 

27 tp = TextPart(text="hello", meta={"lang": "en"}) 

28 d = tp.to_dict() 

29 assert d["type"] == "text" 

30 tp2 = TextPart.from_dict(d) 

31 assert tp2.text == "hello" 

32 assert tp2.meta == {"lang": "en"} 

33 

34 def test_file_part_roundtrip(self): 

35 fp = FilePart( 

36 url="https://ex.com/f.pdf", 

37 filename="report.pdf", 

38 mime_type="application/pdf", 

39 size=1024, 

40 ) 

41 d = fp.to_dict() 

42 fp2 = FilePart.from_dict(d) 

43 assert fp2.filename == "report.pdf" 

44 assert fp2.mime_type == "application/pdf" 

45 

46 def test_data_part_roundtrip(self): 

47 dp = DataPart(data={"count": 42}, schema_uri="https://schema.org/result") 

48 d = dp.to_dict() 

49 dp2 = DataPart.from_dict(d) 

50 assert dp2.data["count"] == 42 

51 

52 def test_part_from_dict_dispatcher(self): 

53 d = {"type": "text", "text": "hi"} 

54 p = part_from_dict(d) 

55 assert isinstance(p, TextPart) 

56 assert p.text == "hi" 

57 

58 d = {"type": "file", "filename": "x.txt"} 

59 p = part_from_dict(d) 

60 assert isinstance(p, FilePart) 

61 

62 d = {"type": "data", "data": {"a": 1}} 

63 p = part_from_dict(d) 

64 assert isinstance(p, DataPart) 

65 

66 

67class TestA2AArtifact: 

68 def test_roundtrip(self): 

69 art = A2AArtifact( 

70 name="result.json", mime_type="application/json", blob=b'{"a":1}', size=8 

71 ) 

72 d = art.to_dict() 

73 art2 = A2AArtifact.from_dict(d) 

74 assert art2.name == "result.json" 

75 assert art2.blob == b'{"a":1}' 

76 

77 def test_url_artifact(self): 

78 art = A2AArtifact(name="image.png", url="https://cdn.ex/img.png") 

79 d = art.to_dict() 

80 assert "url" in d 

81 art2 = A2AArtifact.from_dict(d) 

82 assert art2.url == "https://cdn.ex/img.png" 

83 

84 

85class TestA2AMessage: 

86 def test_user_text(self): 

87 msg = A2AMessage.user_text("hello world") 

88 assert msg.role == MessageRole.USER 

89 assert len(msg.parts) == 1 

90 assert msg.parts[0].text == "hello world" 

91 

92 def test_agent_text(self): 

93 msg = A2AMessage.agent_text("done") 

94 assert msg.role == MessageRole.AGENT 

95 assert msg.get_text() == "done" 

96 

97 def test_multipart_roundtrip(self): 

98 msg = A2AMessage( 

99 role=MessageRole.USER, 

100 parts=[ 

101 TextPart(text="analyze this"), 

102 FilePart(filename="data.csv"), 

103 DataPart(data={"options": {"method": "pca"}}), 

104 ], 

105 ) 

106 d = msg.to_dict() 

107 msg2 = A2AMessage.from_dict(d) 

108 assert msg2.role == MessageRole.USER 

109 assert len(msg2.parts) == 3 

110 assert isinstance(msg2.parts[0], TextPart) 

111 assert isinstance(msg2.parts[1], FilePart) 

112 assert isinstance(msg2.parts[2], DataPart) 

113 assert msg2.get_text() == "analyze this" 

114 

115 

116class TestA2ATask: 

117 def test_lifecycle(self): 

118 task = A2ATask(input=A2AMessage.user_text("do something")) 

119 assert task.state == TaskState.SUBMITTED 

120 

121 task.start_working() 

122 assert task.state == TaskState.WORKING 

123 

124 task.complete(A2AMessage.agent_text("done")) 

125 assert task.state == TaskState.COMPLETED 

126 assert task.output.get_text() == "done" 

127 assert task.is_terminal() 

128 

129 def test_fail(self): 

130 task = A2ATask(input=A2AMessage.user_text("bad")) 

131 task.start_working() 

132 task.fail("something went wrong") 

133 assert task.state == TaskState.FAILED 

134 assert task.error == "something went wrong" 

135 assert task.is_terminal() 

136 

137 def test_cancel(self): 

138 task = A2ATask() 

139 assert not task.is_terminal() 

140 task.cancel() 

141 assert task.state == TaskState.CANCELLED 

142 assert task.is_terminal() 

143 

144 def test_cannot_start_non_submitted(self): 

145 task = A2ATask() 

146 task.start_working() 

147 with pytest.raises(ValueError): 

148 task.start_working() 

149 

150 def test_cannot_complete_non_working(self): 

151 task = A2ATask() 

152 with pytest.raises(ValueError): 

153 task.complete() 

154 

155 def test_cannot_cancel_completed(self): 

156 task = A2ATask() 

157 task.start_working() 

158 task.complete() 

159 with pytest.raises(ValueError): 

160 task.cancel() 

161 

162 def test_artifact_attachment(self): 

163 task = A2ATask() 

164 task.add_artifact(A2AArtifact(name="out.csv")) 

165 task.add_artifact(A2AArtifact(name="out.png")) 

166 assert len(task.artifacts) == 2 

167 

168 def test_json_roundtrip(self): 

169 task = A2ATask(input=A2AMessage.user_text("hello")) 

170 task.start_working() 

171 task.complete(A2AMessage.agent_text("result")) 

172 task.add_artifact(A2AArtifact(name="out.json", blob=b"{}")) 

173 

174 json_str = task.to_json() 

175 task2 = A2ATask.from_json(json_str) 

176 assert task2.task_id == task.task_id 

177 assert task2.state == TaskState.COMPLETED 

178 assert task2.input.get_text() == "hello" 

179 assert task2.artifacts[0].name == "out.json" 

180 

181 def test_state_history(self): 

182 task = A2ATask() 

183 task.start_working() 

184 task.complete() 

185 assert len(task._state_history) == 2 

186 assert task._state_history[0][0] == TaskState.SUBMITTED 

187 assert task._state_history[1][0] == TaskState.WORKING 

188 

189 

190class TestA2AHandoff: 

191 def test_roundtrip(self): 

192 task = A2ATask(input=A2AMessage.user_text("do x")) 

193 ho = A2AHandoff( 

194 source_agent="coordinator", 

195 target_agent="worker", 

196 task=task, 

197 reason="delegation", 

198 ) 

199 d = ho.to_dict() 

200 ho2 = A2AHandoff.from_dict(d) 

201 assert ho2.source_agent == "coordinator" 

202 assert ho2.target_agent == "worker" 

203 assert ho2.task.task_id == task.task_id 

204 assert ho2.reason == "delegation" 

205 

206 def test_json_roundtrip(self): 

207 task = A2ATask(input=A2AMessage.user_text("test")) 

208 ho = A2AHandoff(source_agent="a", target_agent="b", task=task) 

209 ho2 = A2AHandoff.from_json(ho.to_json()) 

210 assert ho2.source_agent == "a" 

211 assert ho2.handoff_id == ho.handoff_id 

212 

213 

214class TestA2ASession: 

215 def test_basic(self): 

216 sess = A2ASession() 

217 sess.add_message(A2AMessage.user_text("hi")) 

218 sess.add_message(A2AMessage.agent_text("hello")) 

219 sess.add_task(A2ATask()) 

220 assert len(sess.history) == 2 

221 assert len(sess.tasks) == 1 

222 

223 def test_get_last_n(self): 

224 sess = A2ASession() 

225 for i in range(5): 

226 sess.add_message(A2AMessage.user_text(f"msg{i}")) 

227 last3 = sess.get_last_n_messages(3) 

228 assert len(last3) == 3 

229 assert last3[-1].get_text() == "msg4" 

230 

231 

232class TestA2Server: 

233 @pytest.mark.asyncio 

234 async def test_process_task_success(self): 

235 server = A2AServer() 

236 

237 async def handler(task: A2ATask): 

238 return A2AMessage.agent_text(f"processed: {task.input.get_text()}") 

239 

240 server.register_handler("worker", handler) 

241 task = new_task("hello test", target_agent="worker") 

242 result = await server.process_task(task.to_dict()) 

243 assert result["state"] == "completed" 

244 assert "processed: hello test" in result["output"]["parts"][0]["text"] 

245 

246 @pytest.mark.asyncio 

247 async def test_process_task_no_handler(self): 

248 server = A2AServer() 

249 task = new_task("hello", target_agent="nonexistent") 

250 result = await server.process_task(task.to_dict()) 

251 assert result["state"] == "failed" 

252 assert "No handler" in result["error"] 

253 

254 @pytest.mark.asyncio 

255 async def test_process_task_handler_error(self): 

256 server = A2AServer() 

257 

258 async def bad_handler(task): 

259 raise ValueError("simulated error") 

260 

261 server.register_handler("bad", bad_handler) 

262 task = new_task("test", target_agent="bad") 

263 result = await server.process_task(task.to_dict()) 

264 assert result["state"] == "failed" 

265 assert "simulated error" in result["error"] 

266 

267 def test_get_task(self): 

268 from agentos.protocols.a2a_store import InMemoryTaskStore 

269 

270 store = InMemoryTaskStore() 

271 server = A2AServer(task_store=store) 

272 task = A2ATask(task_id="task-001") 

273 store.save_task(task) 

274 assert server.get_task("task-001").task_id == "task-001" 

275 assert server.get_task("nonexistent") is None 

276 

277 def test_list_tasks_by_state(self): 

278 from agentos.protocols.a2a_store import InMemoryTaskStore 

279 

280 store = InMemoryTaskStore() 

281 server = A2AServer(task_store=store) 

282 t1 = A2ATask(task_id="t1") 

283 t2 = A2ATask(task_id="t2") 

284 t2.start_working() 

285 t2.complete() 

286 t3 = A2ATask(task_id="t3") 

287 t3.fail("err") 

288 for t in [t1, t2, t3]: 

289 store.save_task(t) 

290 assert len(server.list_tasks()) == 3 

291 assert len(server.list_tasks(TaskState.COMPLETED)) == 1 

292 assert len(server.list_tasks(TaskState.FAILED)) == 1 

293 

294 def test_cleanup(self): 

295 from agentos.protocols.a2a_store import InMemoryTaskStore 

296 

297 store = InMemoryTaskStore() 

298 server = A2AServer(task_store=store) 

299 old = A2ATask(task_id="old") 

300 old.start_working() 

301 old.complete() 

302 old._updated = time.time() - 4000 # fake old 

303 fresh = A2ATask(task_id="fresh") 

304 store.save_task(old) 

305 store.save_task(fresh) 

306 n = server.cleanup_old(max_age_seconds=3600) 

307 assert n == 1 

308 assert store.get_task("old") is None 

309 assert store.get_task("fresh") is not None 

310 

311 

312class TestConvenience: 

313 def test_new_task(self): 

314 t = new_task("my task", target_agent="worker", priority="high") 

315 assert t.input.get_text() == "my task" 

316 assert t.meta["target_agent"] == "worker" 

317 assert t.meta["priority"] == "high" 

318 

319 def test_new_handoff(self): 

320 t = new_task("delegate me") 

321 ho = new_handoff(t, source="a", target="b", reason="overload") 

322 assert ho.source_agent == "a" 

323 assert ho.target_agent == "b" 

324 assert ho.reason == "overload"