first commit
This commit is contained in:
@@ -0,0 +1,220 @@
|
||||
import os
|
||||
import sys
|
||||
import importlib
|
||||
from datetime import datetime
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
if PROJECT_ROOT not in sys.path:
|
||||
sys.path.insert(0, PROJECT_ROOT)
|
||||
|
||||
from SimpleLLMFunc.hooks.events import (
|
||||
CustomEvent,
|
||||
ReactEndEvent,
|
||||
ReactIterationStartEvent,
|
||||
ReActEventType,
|
||||
)
|
||||
from SimpleLLMFunc.hooks.stream import EventOrigin, EventYield, ResponseYield
|
||||
|
||||
from agent.BaseAgent import BaseAgent
|
||||
from react_stream import format_sse, serialize_react_output
|
||||
|
||||
|
||||
base_agent_module = importlib.import_module("agent.BaseAgent")
|
||||
|
||||
|
||||
class _FakeDelta:
|
||||
def __init__(self, content, reasoning=None):
|
||||
self.content = content
|
||||
self.reasoning = reasoning
|
||||
|
||||
|
||||
class _FakeChoice:
|
||||
def __init__(self, content, reasoning=None):
|
||||
self.delta = _FakeDelta(content, reasoning)
|
||||
|
||||
|
||||
class _FakeChunk:
|
||||
def __init__(self, content, reasoning=None):
|
||||
self.choices = [_FakeChoice(content, reasoning)]
|
||||
|
||||
|
||||
class _DummyAgent(BaseAgent):
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def get_toolkit(self):
|
||||
return []
|
||||
|
||||
def chat_impl(self, history, query, sketch_pad_summary):
|
||||
if False:
|
||||
yield history, query, sketch_pad_summary
|
||||
|
||||
async def run(self, query, raw_user_content=None):
|
||||
if False:
|
||||
yield query, raw_user_content
|
||||
|
||||
|
||||
class _FakeContext:
|
||||
def __init__(self):
|
||||
self.messages = []
|
||||
|
||||
async def store_message(self, message):
|
||||
self.messages.append(message)
|
||||
|
||||
|
||||
def _build_origin() -> EventOrigin:
|
||||
return EventOrigin(
|
||||
session_id="session-1",
|
||||
agent_call_id="agent-call-1",
|
||||
event_seq=1,
|
||||
)
|
||||
|
||||
|
||||
def test_serialize_react_output_normalizes_response_and_event_payloads():
|
||||
response_output = ResponseYield(
|
||||
response=cast(Any, _FakeChunk("hello world", reasoning="thinking")),
|
||||
messages=[{"role": "assistant", "content": "hello world"}],
|
||||
)
|
||||
response_payload = serialize_react_output(response_output, delta_consumer="web")
|
||||
|
||||
assert response_payload["type"] == "response"
|
||||
assert response_payload["delta_text"] == "hello world"
|
||||
assert response_payload["delta_reasoning"] == "thinking"
|
||||
assert response_payload["messages"][0]["content"] == "hello world"
|
||||
|
||||
event = ReactEndEvent(
|
||||
event_type=ReActEventType.REACT_END,
|
||||
timestamp=datetime(2026, 3, 18, 12, 0, 0),
|
||||
trace_id="trace-1",
|
||||
func_name="chat_impl",
|
||||
iteration=1,
|
||||
final_response="done",
|
||||
final_messages=[{"role": "assistant", "content": "done"}],
|
||||
total_iterations=1,
|
||||
total_execution_time=0.5,
|
||||
total_tool_calls=0,
|
||||
total_llm_calls=1,
|
||||
)
|
||||
event_output = EventYield(event=event, origin=_build_origin())
|
||||
event_payload = serialize_react_output(event_output)
|
||||
|
||||
assert event_payload["type"] == "event"
|
||||
assert event_payload["event_type"] == "react_end"
|
||||
assert event_payload["event"]["timestamp"] == "2026-03-18T12:00:00"
|
||||
assert event_payload["origin"]["session_id"] == "session-1"
|
||||
|
||||
sse_packet = format_sse("response", response_payload)
|
||||
assert sse_packet.startswith("event: response\n")
|
||||
assert '"delta_text": "hello world"' in sse_packet
|
||||
|
||||
|
||||
def test_custom_event_uses_event_name_for_stream_routing():
|
||||
event = CustomEvent(
|
||||
event_type=ReActEventType.CUSTOM_EVENT,
|
||||
timestamp=datetime(2026, 3, 18, 12, 0, 0),
|
||||
trace_id="trace-1",
|
||||
func_name="chat_impl",
|
||||
iteration=1,
|
||||
event_name="subagent_status",
|
||||
data={"phase": "started"},
|
||||
)
|
||||
event_output = EventYield(event=event, origin=_build_origin())
|
||||
|
||||
payload = serialize_react_output(event_output)
|
||||
|
||||
assert payload["event_type"] == "subagent_status"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_and_persist_ignores_events_and_preserves_message_order(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_context = _FakeContext()
|
||||
monkeypatch.setattr(base_agent_module, "get_current_context", lambda: fake_context)
|
||||
|
||||
tool_call = [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {"name": "lookup", "arguments": "{}"},
|
||||
}
|
||||
]
|
||||
|
||||
async def output_stream():
|
||||
yield ResponseYield(
|
||||
response=cast(Any, _FakeChunk("Hello")),
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
yield EventYield(
|
||||
event=ReactIterationStartEvent(
|
||||
event_type=ReActEventType.REACT_ITERATION_START,
|
||||
timestamp=datetime(2026, 3, 18, 12, 0, 0),
|
||||
trace_id="trace-1",
|
||||
func_name="chat_impl",
|
||||
iteration=1,
|
||||
current_messages=[{"role": "user", "content": "hi"}],
|
||||
),
|
||||
origin=_build_origin(),
|
||||
)
|
||||
yield ResponseYield(
|
||||
response=cast(Any, _FakeChunk("")),
|
||||
messages=cast(
|
||||
Any,
|
||||
[
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": None, "tool_calls": tool_call},
|
||||
],
|
||||
),
|
||||
)
|
||||
yield ResponseYield(
|
||||
response=cast(Any, _FakeChunk("")),
|
||||
messages=cast(
|
||||
Any,
|
||||
[
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": None, "tool_calls": tool_call},
|
||||
{
|
||||
"role": "tool",
|
||||
"content": "lookup result",
|
||||
"tool_call_id": "call_1",
|
||||
},
|
||||
],
|
||||
),
|
||||
)
|
||||
yield ResponseYield(
|
||||
response=cast(Any, _FakeChunk(" world")),
|
||||
messages=cast(
|
||||
Any,
|
||||
[
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": None, "tool_calls": tool_call},
|
||||
{
|
||||
"role": "tool",
|
||||
"content": "lookup result",
|
||||
"tool_call_id": "call_1",
|
||||
},
|
||||
{"role": "assistant", "content": "Hello world"},
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
agent = _DummyAgent()
|
||||
|
||||
yielded = []
|
||||
async for output in agent._stream_and_persist(output_stream()):
|
||||
yielded.append(output)
|
||||
|
||||
assert len(yielded) == 5
|
||||
assert [message.role for message in fake_context.messages] == [
|
||||
"assistant",
|
||||
"assistant",
|
||||
"tool",
|
||||
"assistant",
|
||||
]
|
||||
assert fake_context.messages[0].content == "Hello"
|
||||
assert fake_context.messages[1].tool_calls[0].id == "call_1"
|
||||
assert fake_context.messages[2].tool_call_id == "call_1"
|
||||
assert fake_context.messages[3].content == " world"
|
||||
Reference in New Issue
Block a user