test_openai_stream_parser.py 3.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133
  1. from agent_lab.domain.events import ToolCallEvent
  2. from agent_lab.infrastructure.openai_compatible import ChatCompletionStreamParser
  3. def test_parser_emits_visible_content_usage_and_text_protocol_events():
  4. parser = ChatCompletionStreamParser()
  5. items = []
  6. items.extend(
  7. parser.feed(
  8. {
  9. "choices": [
  10. {
  11. "delta": {"content": "hello"},
  12. "finish_reason": None,
  13. }
  14. ]
  15. }
  16. )
  17. )
  18. items.extend(
  19. parser.feed(
  20. {
  21. "choices": [
  22. {
  23. "delta": {"content": "\n<agent_"},
  24. "finish_reason": None,
  25. }
  26. ]
  27. }
  28. )
  29. )
  30. items.extend(
  31. parser.feed(
  32. {
  33. "choices": [
  34. {
  35. "delta": {
  36. "content": "events>\nhandoff_note\nmock_search\n</agent_events>"
  37. },
  38. "finish_reason": "stop",
  39. }
  40. ],
  41. "usage": {
  42. "prompt_tokens": 10,
  43. "completion_tokens": 3,
  44. "total_tokens": 13,
  45. "prompt_tokens_details": {"cached_tokens": 4},
  46. },
  47. }
  48. )
  49. )
  50. assert [item.content for item in items if item.kind == "message_delta"] == [
  51. "hello",
  52. "\n",
  53. ]
  54. events = [item.event for item in items if item.kind == "event"]
  55. assert events == [
  56. ToolCallEvent(
  57. id="event_1",
  58. name="handoff_note",
  59. arguments={},
  60. raw_arguments="{}",
  61. ),
  62. ToolCallEvent(
  63. id="event_2",
  64. name="mock_search",
  65. arguments={},
  66. raw_arguments="{}",
  67. )
  68. ]
  69. usage = [item.usage for item in items if item.kind == "usage"][0]
  70. assert usage.total_tokens == 13
  71. assert usage.cached_tokens == 4
  72. def test_parser_emits_provider_tool_call_events_for_event_agent():
  73. parser = ChatCompletionStreamParser()
  74. items = []
  75. items.extend(
  76. parser.feed(
  77. {
  78. "choices": [
  79. {
  80. "delta": {
  81. "tool_calls": [
  82. {
  83. "index": 0,
  84. "id": "call_1",
  85. "type": "function",
  86. "function": {
  87. "name": "mock_search",
  88. "arguments": '{"query":"',
  89. },
  90. }
  91. ]
  92. },
  93. "finish_reason": None,
  94. }
  95. ]
  96. }
  97. )
  98. )
  99. items.extend(
  100. parser.feed(
  101. {
  102. "choices": [
  103. {
  104. "delta": {
  105. "tool_calls": [
  106. {
  107. "index": 0,
  108. "function": {"arguments": 'latency docs"}'},
  109. }
  110. ]
  111. },
  112. "finish_reason": "tool_calls",
  113. }
  114. ]
  115. }
  116. )
  117. )
  118. assert [item.event for item in items if item.kind == "event"] == [
  119. ToolCallEvent(
  120. id="call_1",
  121. name="mock_search",
  122. arguments={"query": "latency docs"},
  123. raw_arguments='{"query":"latency docs"}',
  124. )
  125. ]