فهرست منبع

feat: distinguish provider tool calls from text events

Problem: shared event kind discarded the source and provider arguments needed for direct ChatAgent tools. Risk: compatible providers may omit finish_reason; flush behavior is guarded against duplicate emission.
zhenyu.hu 2 هفته پیش
والد
کامیت
c51daebfd2
3فایلهای تغییر یافته به همراه194 افزوده شده و 8 حذف شده
  1. 8 0
      src/agent_lab/domain/messages.py
  2. 4 4
      src/agent_lab/infrastructure/openai_compatible.py
  3. 182 4
      tests/test_openai_stream_parser.py

+ 8 - 0
src/agent_lab/domain/messages.py

@@ -36,6 +36,14 @@ class StreamItem:
     def message_delta(cls, content: str) -> "StreamItem":
         return cls(kind="message_delta", content=content)
 
+    @classmethod
+    def text_event(cls, event: ToolCallEvent) -> "StreamItem":
+        return cls(kind="text_event", event=event)
+
+    @classmethod
+    def provider_tool_call(cls, event: ToolCallEvent) -> "StreamItem":
+        return cls(kind="provider_tool_call", event=event)
+
     @classmethod
     def usage_item(cls, usage: TokenUsage) -> "StreamItem":
         return cls(kind="usage", usage=usage)

+ 4 - 4
src/agent_lab/infrastructure/openai_compatible.py

@@ -44,7 +44,7 @@ class ChatCompletionStreamParser:
         return items
 
     def flush(self) -> list[StreamItem]:
-        return self._text_events.flush()
+        return [*self._drain_tool_events(), *self._text_events.flush()]
 
     def _accumulate_tool_call(self, tool_call: dict[str, Any]) -> None:
         index = int(tool_call["index"])
@@ -62,7 +62,7 @@ class ChatCompletionStreamParser:
 
         function = tool_call.get("function") or {}
         if function.get("name"):
-            state["name"] = function["name"]
+            state["name"] += function["name"]
         if "arguments" in function and function["arguments"] is not None:
             state["arguments"] += function["arguments"]
 
@@ -72,7 +72,7 @@ class ChatCompletionStreamParser:
             state = self._tool_calls[index]
             raw_arguments = state["arguments"]
             events.append(
-                StreamItem.event(
+                StreamItem.provider_tool_call(
                     ToolCallEvent(
                         id=state["id"],
                         name=state["name"],
@@ -193,7 +193,7 @@ class TextEventProtocolParser:
                 continue
             self._event_count += 1
             events.append(
-                StreamItem.event(
+                StreamItem.text_event(
                     ToolCallEvent(
                         id=f"event_{self._event_count}",
                         name=name,

+ 182 - 4
tests/test_openai_stream_parser.py

@@ -1,4 +1,5 @@
 from agent_lab.domain.events import ToolCallEvent
+from agent_lab.domain.messages import StreamItem
 from agent_lab.infrastructure.openai_compatible import ChatCompletionStreamParser
 
 
@@ -55,7 +56,7 @@ def test_parser_emits_visible_content_usage_and_text_protocol_events():
         "hello",
         "\n",
     ]
-    events = [item.event for item in items if item.kind == "event"]
+    events = [item.event for item in items if item.kind == "text_event"]
     assert events == [
         ToolCallEvent(
             id="event_1",
@@ -75,10 +76,26 @@ def test_parser_emits_visible_content_usage_and_text_protocol_events():
     assert usage.cached_tokens == 4
 
 
-def test_parser_emits_provider_tool_call_events_for_event_agent():
+def test_parser_distinguishes_text_events_from_provider_tool_calls():
     parser = ChatCompletionStreamParser()
 
     items = []
+    items.extend(
+        parser.feed(
+            {
+                "choices": [
+                    {
+                        "delta": {
+                            "content": (
+                                "visible<agent_events>mock_search</agent_events>"
+                            )
+                        },
+                        "finish_reason": None,
+                    }
+                ]
+            }
+        )
+    )
     items.extend(
         parser.feed(
             {
@@ -123,7 +140,20 @@ def test_parser_emits_provider_tool_call_events_for_event_agent():
         )
     )
 
-    assert [item.event for item in items if item.kind == "event"] == [
+    assert [item.content for item in items if item.kind == "message_delta"] == [
+        "visible"
+    ]
+    assert [item.event for item in items if item.kind == "text_event"] == [
+        ToolCallEvent(
+            id="event_1",
+            name="mock_search",
+            arguments={},
+            raw_arguments="{}",
+        )
+    ]
+    assert [
+        item.event for item in items if item.kind == "provider_tool_call"
+    ] == [
         ToolCallEvent(
             id="call_1",
             name="mock_search",
@@ -162,7 +192,104 @@ def test_parser_drains_provider_tool_call_events_on_any_terminal_finish_reason()
         )
     )
 
-    assert [item.event for item in items if item.kind == "event"] == [
+    assert [
+        item.event for item in items if item.kind == "provider_tool_call"
+    ] == [
+        ToolCallEvent(
+            id="call_1",
+            name="mock_search",
+            arguments={"query": "latency docs"},
+            raw_arguments='{"query":"latency docs"}',
+        )
+    ]
+
+
+def test_parser_flushes_provider_tool_calls_without_finish_reason():
+    parser = ChatCompletionStreamParser()
+
+    parser.feed(
+        {
+            "choices": [
+                {
+                    "delta": {
+                        "tool_calls": [
+                            {
+                                "index": 0,
+                                "id": "call_1",
+                                "function": {
+                                    "name": "mock_",
+                                    "arguments": '{"query":"',
+                                },
+                            }
+                        ]
+                    },
+                    "finish_reason": None,
+                }
+            ]
+        }
+    )
+    parser.feed(
+        {
+            "choices": [
+                {
+                    "delta": {
+                        "tool_calls": [
+                            {
+                                "index": 0,
+                                "function": {
+                                    "name": "search",
+                                    "arguments": 'latency docs"}',
+                                },
+                            }
+                        ]
+                    },
+                    "finish_reason": None,
+                }
+            ]
+        }
+    )
+
+    assert parser.flush() == [
+        StreamItem.provider_tool_call(
+            ToolCallEvent(
+                id="call_1",
+                name="mock_search",
+                arguments={"query": "latency docs"},
+                raw_arguments='{"query":"latency docs"}',
+            )
+        )
+    ]
+
+
+def test_parser_does_not_emit_provider_tool_call_twice_after_finish_and_flush():
+    parser = ChatCompletionStreamParser()
+
+    items = parser.feed(
+        {
+            "choices": [
+                {
+                    "delta": {
+                        "tool_calls": [
+                            {
+                                "index": 0,
+                                "id": "call_1",
+                                "function": {
+                                    "name": "mock_search",
+                                    "arguments": '{"query":"latency docs"}',
+                                },
+                            }
+                        ]
+                    },
+                    "finish_reason": "tool_calls",
+                }
+            ]
+        }
+    )
+    items.extend(parser.flush())
+
+    assert [
+        item.event for item in items if item.kind == "provider_tool_call"
+    ] == [
         ToolCallEvent(
             id="call_1",
             name="mock_search",
@@ -170,3 +297,54 @@ def test_parser_drains_provider_tool_call_events_on_any_terminal_finish_reason()
             raw_arguments='{"query":"latency docs"}',
         )
     ]
+
+
+def test_parser_preserves_multiple_provider_tool_call_order():
+    parser = ChatCompletionStreamParser()
+
+    items = parser.feed(
+        {
+            "choices": [
+                {
+                    "delta": {
+                        "tool_calls": [
+                            {
+                                "index": 1,
+                                "id": "call_2",
+                                "function": {
+                                    "name": "second",
+                                    "arguments": '{"value":2}',
+                                },
+                            },
+                            {
+                                "index": 0,
+                                "id": "call_1",
+                                "function": {
+                                    "name": "first",
+                                    "arguments": '{"value":1}',
+                                },
+                            },
+                        ]
+                    },
+                    "finish_reason": "tool_calls",
+                }
+            ]
+        }
+    )
+
+    assert [
+        item.event for item in items if item.kind == "provider_tool_call"
+    ] == [
+        ToolCallEvent(
+            id="call_1",
+            name="first",
+            arguments={"value": 1},
+            raw_arguments='{"value":1}',
+        ),
+        ToolCallEvent(
+            id="call_2",
+            name="second",
+            arguments={"value": 2},
+            raw_arguments='{"value":2}',
+        ),
+    ]