test_openai_stream_parser.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350
  1. from agent_lab.domain.events import ToolCallEvent
  2. from agent_lab.domain.messages import StreamItem
  3. from agent_lab.infrastructure.openai_compatible import ChatCompletionStreamParser
  4. def test_parser_emits_visible_content_usage_and_text_protocol_events():
  5. parser = ChatCompletionStreamParser()
  6. items = []
  7. items.extend(
  8. parser.feed(
  9. {
  10. "choices": [
  11. {
  12. "delta": {"content": "hello"},
  13. "finish_reason": None,
  14. }
  15. ]
  16. }
  17. )
  18. )
  19. items.extend(
  20. parser.feed(
  21. {
  22. "choices": [
  23. {
  24. "delta": {"content": "\n<agent_"},
  25. "finish_reason": None,
  26. }
  27. ]
  28. }
  29. )
  30. )
  31. items.extend(
  32. parser.feed(
  33. {
  34. "choices": [
  35. {
  36. "delta": {
  37. "content": "events>\nhandoff_note\nmock_search\n</agent_events>"
  38. },
  39. "finish_reason": "stop",
  40. }
  41. ],
  42. "usage": {
  43. "prompt_tokens": 10,
  44. "completion_tokens": 3,
  45. "total_tokens": 13,
  46. "prompt_tokens_details": {"cached_tokens": 4},
  47. },
  48. }
  49. )
  50. )
  51. assert [item.content for item in items if item.kind == "message_delta"] == [
  52. "hello",
  53. "\n",
  54. ]
  55. events = [item.event for item in items if item.kind == "text_event"]
  56. assert events == [
  57. ToolCallEvent(
  58. id="event_1",
  59. name="handoff_note",
  60. arguments={},
  61. raw_arguments="{}",
  62. ),
  63. ToolCallEvent(
  64. id="event_2",
  65. name="mock_search",
  66. arguments={},
  67. raw_arguments="{}",
  68. )
  69. ]
  70. usage = [item.usage for item in items if item.kind == "usage"][0]
  71. assert usage.total_tokens == 13
  72. assert usage.cached_tokens == 4
  73. def test_parser_distinguishes_text_events_from_provider_tool_calls():
  74. parser = ChatCompletionStreamParser()
  75. items = []
  76. items.extend(
  77. parser.feed(
  78. {
  79. "choices": [
  80. {
  81. "delta": {
  82. "content": (
  83. "visible<agent_events>mock_search</agent_events>"
  84. )
  85. },
  86. "finish_reason": None,
  87. }
  88. ]
  89. }
  90. )
  91. )
  92. items.extend(
  93. parser.feed(
  94. {
  95. "choices": [
  96. {
  97. "delta": {
  98. "tool_calls": [
  99. {
  100. "index": 0,
  101. "id": "call_1",
  102. "type": "function",
  103. "function": {
  104. "name": "mock_search",
  105. "arguments": '{"query":"',
  106. },
  107. }
  108. ]
  109. },
  110. "finish_reason": None,
  111. }
  112. ]
  113. }
  114. )
  115. )
  116. items.extend(
  117. parser.feed(
  118. {
  119. "choices": [
  120. {
  121. "delta": {
  122. "tool_calls": [
  123. {
  124. "index": 0,
  125. "function": {"arguments": 'latency docs"}'},
  126. }
  127. ]
  128. },
  129. "finish_reason": "tool_calls",
  130. }
  131. ]
  132. }
  133. )
  134. )
  135. assert [item.content for item in items if item.kind == "message_delta"] == [
  136. "visible"
  137. ]
  138. assert [item.event for item in items if item.kind == "text_event"] == [
  139. ToolCallEvent(
  140. id="event_1",
  141. name="mock_search",
  142. arguments={},
  143. raw_arguments="{}",
  144. )
  145. ]
  146. assert [
  147. item.event for item in items if item.kind == "provider_tool_call"
  148. ] == [
  149. ToolCallEvent(
  150. id="call_1",
  151. name="mock_search",
  152. arguments={"query": "latency docs"},
  153. raw_arguments='{"query":"latency docs"}',
  154. )
  155. ]
  156. def test_parser_drains_provider_tool_call_events_on_any_terminal_finish_reason():
  157. parser = ChatCompletionStreamParser()
  158. items = []
  159. items.extend(
  160. parser.feed(
  161. {
  162. "choices": [
  163. {
  164. "delta": {
  165. "tool_calls": [
  166. {
  167. "index": 0,
  168. "id": "call_1",
  169. "type": "function",
  170. "function": {
  171. "name": "mock_search",
  172. "arguments": '{"query":"latency docs"}',
  173. },
  174. }
  175. ]
  176. },
  177. "finish_reason": "stop",
  178. }
  179. ]
  180. }
  181. )
  182. )
  183. assert [
  184. item.event for item in items if item.kind == "provider_tool_call"
  185. ] == [
  186. ToolCallEvent(
  187. id="call_1",
  188. name="mock_search",
  189. arguments={"query": "latency docs"},
  190. raw_arguments='{"query":"latency docs"}',
  191. )
  192. ]
  193. def test_parser_flushes_provider_tool_calls_without_finish_reason():
  194. parser = ChatCompletionStreamParser()
  195. parser.feed(
  196. {
  197. "choices": [
  198. {
  199. "delta": {
  200. "tool_calls": [
  201. {
  202. "index": 0,
  203. "id": "call_1",
  204. "function": {
  205. "name": "mock_",
  206. "arguments": '{"query":"',
  207. },
  208. }
  209. ]
  210. },
  211. "finish_reason": None,
  212. }
  213. ]
  214. }
  215. )
  216. parser.feed(
  217. {
  218. "choices": [
  219. {
  220. "delta": {
  221. "tool_calls": [
  222. {
  223. "index": 0,
  224. "function": {
  225. "name": "search",
  226. "arguments": 'latency docs"}',
  227. },
  228. }
  229. ]
  230. },
  231. "finish_reason": None,
  232. }
  233. ]
  234. }
  235. )
  236. assert parser.flush() == [
  237. StreamItem.provider_tool_call(
  238. ToolCallEvent(
  239. id="call_1",
  240. name="mock_search",
  241. arguments={"query": "latency docs"},
  242. raw_arguments='{"query":"latency docs"}',
  243. )
  244. )
  245. ]
  246. def test_parser_does_not_emit_provider_tool_call_twice_after_finish_and_flush():
  247. parser = ChatCompletionStreamParser()
  248. items = parser.feed(
  249. {
  250. "choices": [
  251. {
  252. "delta": {
  253. "tool_calls": [
  254. {
  255. "index": 0,
  256. "id": "call_1",
  257. "function": {
  258. "name": "mock_search",
  259. "arguments": '{"query":"latency docs"}',
  260. },
  261. }
  262. ]
  263. },
  264. "finish_reason": "tool_calls",
  265. }
  266. ]
  267. }
  268. )
  269. items.extend(parser.flush())
  270. assert [
  271. item.event for item in items if item.kind == "provider_tool_call"
  272. ] == [
  273. ToolCallEvent(
  274. id="call_1",
  275. name="mock_search",
  276. arguments={"query": "latency docs"},
  277. raw_arguments='{"query":"latency docs"}',
  278. )
  279. ]
  280. def test_parser_preserves_multiple_provider_tool_call_order():
  281. parser = ChatCompletionStreamParser()
  282. items = parser.feed(
  283. {
  284. "choices": [
  285. {
  286. "delta": {
  287. "tool_calls": [
  288. {
  289. "index": 1,
  290. "id": "call_2",
  291. "function": {
  292. "name": "second",
  293. "arguments": '{"value":2}',
  294. },
  295. },
  296. {
  297. "index": 0,
  298. "id": "call_1",
  299. "function": {
  300. "name": "first",
  301. "arguments": '{"value":1}',
  302. },
  303. },
  304. ]
  305. },
  306. "finish_reason": "tool_calls",
  307. }
  308. ]
  309. }
  310. )
  311. assert [
  312. item.event for item in items if item.kind == "provider_tool_call"
  313. ] == [
  314. ToolCallEvent(
  315. id="call_1",
  316. name="first",
  317. arguments={"value": 1},
  318. raw_arguments='{"value":1}',
  319. ),
  320. ToolCallEvent(
  321. id="call_2",
  322. name="second",
  323. arguments={"value": 2},
  324. raw_arguments='{"value":2}',
  325. ),
  326. ]