|
@@ -1,4 +1,5 @@
|
|
|
from agent_lab.domain.events import ToolCallEvent
|
|
from agent_lab.domain.events import ToolCallEvent
|
|
|
|
|
+from agent_lab.domain.messages import StreamItem
|
|
|
from agent_lab.infrastructure.openai_compatible import ChatCompletionStreamParser
|
|
from agent_lab.infrastructure.openai_compatible import ChatCompletionStreamParser
|
|
|
|
|
|
|
|
|
|
|
|
@@ -55,7 +56,7 @@ def test_parser_emits_visible_content_usage_and_text_protocol_events():
|
|
|
"hello",
|
|
"hello",
|
|
|
"\n",
|
|
"\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 == [
|
|
assert events == [
|
|
|
ToolCallEvent(
|
|
ToolCallEvent(
|
|
|
id="event_1",
|
|
id="event_1",
|
|
@@ -75,10 +76,26 @@ def test_parser_emits_visible_content_usage_and_text_protocol_events():
|
|
|
assert usage.cached_tokens == 4
|
|
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()
|
|
parser = ChatCompletionStreamParser()
|
|
|
|
|
|
|
|
items = []
|
|
items = []
|
|
|
|
|
+ items.extend(
|
|
|
|
|
+ parser.feed(
|
|
|
|
|
+ {
|
|
|
|
|
+ "choices": [
|
|
|
|
|
+ {
|
|
|
|
|
+ "delta": {
|
|
|
|
|
+ "content": (
|
|
|
|
|
+ "visible<agent_events>mock_search</agent_events>"
|
|
|
|
|
+ )
|
|
|
|
|
+ },
|
|
|
|
|
+ "finish_reason": None,
|
|
|
|
|
+ }
|
|
|
|
|
+ ]
|
|
|
|
|
+ }
|
|
|
|
|
+ )
|
|
|
|
|
+ )
|
|
|
items.extend(
|
|
items.extend(
|
|
|
parser.feed(
|
|
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(
|
|
ToolCallEvent(
|
|
|
id="call_1",
|
|
id="call_1",
|
|
|
name="mock_search",
|
|
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(
|
|
ToolCallEvent(
|
|
|
id="call_1",
|
|
id="call_1",
|
|
|
name="mock_search",
|
|
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"}',
|
|
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}',
|
|
|
|
|
+ ),
|
|
|
|
|
+ ]
|