first commit
This commit is contained in:
@@ -0,0 +1,434 @@
|
||||
# pyright: reportCallIssue=false, reportArgumentType=false
|
||||
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime
|
||||
import importlib
|
||||
import json
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
from contextlib import contextmanager
|
||||
|
||||
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 ReactEndEvent, ReactStartEvent, ReActEventType
|
||||
from SimpleLLMFunc.hooks.stream import EventOrigin, EventYield, ResponseYield
|
||||
|
||||
from web_interface.models import ChatCompletionRequest, ChatMessage
|
||||
|
||||
chat_router_module = importlib.import_module("web_interface.routers.chat_router")
|
||||
utils_module = importlib.import_module("web_interface.utils")
|
||||
from web_interface.routers.chat_router import stream_chat_completion, stream_chat_events
|
||||
from web_interface.utils import process_agent_response, validate_chat_request
|
||||
|
||||
|
||||
class _FakeConversation:
|
||||
def __init__(self):
|
||||
self.uuid = "conversation-1"
|
||||
self.context = type("Ctx", (), {"persist": self._persist})()
|
||||
self.sketch_pad = type("Sketch", (), {"persist": lambda self: None})()
|
||||
self.persisted = False
|
||||
|
||||
async def _persist(self):
|
||||
self.persisted = True
|
||||
return True
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
|
||||
class _FakeAgent:
|
||||
def __init__(self, outputs):
|
||||
self.outputs = outputs
|
||||
self.name = "fake-agent"
|
||||
self.queries = []
|
||||
self.raw_contents = []
|
||||
|
||||
async def run(self, query, raw_user_content=None):
|
||||
self.queries.append(query)
|
||||
self.raw_contents.append(raw_user_content)
|
||||
for output in self.outputs:
|
||||
yield output
|
||||
|
||||
|
||||
def _origin(seq: int = 1) -> EventOrigin:
|
||||
return EventOrigin(
|
||||
session_id="session-1",
|
||||
agent_call_id="agent-call-1",
|
||||
event_seq=seq,
|
||||
)
|
||||
|
||||
|
||||
def _request() -> ChatCompletionRequest:
|
||||
return _request_with_messages([_user_message("make a cube")], stream=True)
|
||||
|
||||
|
||||
def _user_message(content: Any) -> ChatMessage:
|
||||
return ChatMessage(
|
||||
role="user",
|
||||
content=content,
|
||||
name=None,
|
||||
tool_calls=None,
|
||||
tool_call_id=None,
|
||||
)
|
||||
|
||||
|
||||
def _request_with_messages(
|
||||
messages: list[ChatMessage],
|
||||
*,
|
||||
stream: bool = True,
|
||||
) -> ChatCompletionRequest:
|
||||
return ChatCompletionRequest(
|
||||
model="cadagent",
|
||||
messages=messages,
|
||||
temperature=1.0,
|
||||
top_p=1.0,
|
||||
n=1,
|
||||
stream=stream,
|
||||
stop=None,
|
||||
max_tokens=None,
|
||||
presence_penalty=0.0,
|
||||
frequency_penalty=0.0,
|
||||
logit_bias=None,
|
||||
user=None,
|
||||
tools=None,
|
||||
tool_choice=None,
|
||||
)
|
||||
|
||||
|
||||
def _parse_sse_lines(lines):
|
||||
current_event = "message"
|
||||
data_lines = []
|
||||
|
||||
def _flush_packet():
|
||||
nonlocal current_event, data_lines
|
||||
if not data_lines:
|
||||
return None
|
||||
|
||||
payload_text = "\n".join(data_lines)
|
||||
try:
|
||||
payload = json.loads(payload_text)
|
||||
except json.JSONDecodeError:
|
||||
payload = {"raw": payload_text}
|
||||
|
||||
packet = {"event": current_event, "data": payload}
|
||||
current_event = "message"
|
||||
data_lines = []
|
||||
return packet
|
||||
|
||||
for raw_line in lines:
|
||||
line = (
|
||||
raw_line.decode("utf-8") if isinstance(raw_line, bytes) else str(raw_line)
|
||||
)
|
||||
|
||||
if line == "":
|
||||
packet = _flush_packet()
|
||||
if packet is not None:
|
||||
yield packet
|
||||
continue
|
||||
|
||||
if line.startswith(":"):
|
||||
continue
|
||||
if line.startswith("event:"):
|
||||
current_event = line[6:].strip() or "message"
|
||||
continue
|
||||
if line.startswith("data:"):
|
||||
data_lines.append(line[5:].strip())
|
||||
|
||||
packet = _flush_packet()
|
||||
if packet is not None:
|
||||
yield packet
|
||||
|
||||
|
||||
def test_validate_chat_request_accepts_image_only_user_message():
|
||||
request = _request_with_messages(
|
||||
[
|
||||
_user_message(
|
||||
[
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,abcd"},
|
||||
}
|
||||
]
|
||||
)
|
||||
],
|
||||
stream=False,
|
||||
)
|
||||
|
||||
query, request_id, raw_user_content = validate_chat_request(request)
|
||||
|
||||
assert isinstance(query, list)
|
||||
assert request_id.startswith("chatcmpl-")
|
||||
assert isinstance(raw_user_content, list)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_chat_events_passes_multimodal_query_to_agent():
|
||||
request = _request_with_messages(
|
||||
[
|
||||
_user_message(
|
||||
[
|
||||
{"type": "text", "text": "analyze this"},
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": "data:image/png;base64,abcd"},
|
||||
},
|
||||
]
|
||||
)
|
||||
],
|
||||
stream=True,
|
||||
)
|
||||
conversation = _FakeConversation()
|
||||
agent = _FakeAgent([ResponseYield(response="ok", messages=[])])
|
||||
|
||||
_ = [
|
||||
packet
|
||||
async for packet in stream_chat_events(
|
||||
request, cast(Any, conversation), cast(Any, agent)
|
||||
)
|
||||
]
|
||||
|
||||
assert len(agent.queries) == 1
|
||||
assert isinstance(agent.queries[0], list)
|
||||
assert isinstance(agent.raw_contents[0], list)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_chat_completion_projects_only_response_packets():
|
||||
request = _request()
|
||||
conversation = _FakeConversation()
|
||||
outputs = [
|
||||
EventYield(
|
||||
event=ReactStartEvent(
|
||||
event_type=ReActEventType.REACT_START,
|
||||
timestamp=datetime(2026, 3, 18, 12, 0, 0),
|
||||
trace_id="trace-1",
|
||||
func_name="chat_impl",
|
||||
iteration=0,
|
||||
user_task_prompt="make a cube",
|
||||
initial_messages=[],
|
||||
available_tools=[],
|
||||
),
|
||||
origin=_origin(1),
|
||||
),
|
||||
ResponseYield(
|
||||
response=cast(
|
||||
Any,
|
||||
{
|
||||
"id": "chunk-1",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1,
|
||||
"model": "cadagent",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"role": "assistant",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "cad_code_generator",
|
||||
"arguments": "{}",
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
"finish_reason": None,
|
||||
}
|
||||
],
|
||||
},
|
||||
),
|
||||
messages=[],
|
||||
),
|
||||
]
|
||||
agent = _FakeAgent(outputs)
|
||||
|
||||
packets = [
|
||||
packet
|
||||
async for packet in stream_chat_completion(
|
||||
request,
|
||||
"chatcmpl-test",
|
||||
cast(Any, conversation),
|
||||
cast(Any, agent),
|
||||
)
|
||||
]
|
||||
|
||||
assert any('"tool_calls"' in packet for packet in packets)
|
||||
assert not any("react_start" in packet for packet in packets)
|
||||
assert packets[-1] == "data: [DONE]\n\n"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_chat_events_emits_named_sse_events_and_done():
|
||||
request = _request()
|
||||
conversation = _FakeConversation()
|
||||
outputs = [
|
||||
EventYield(
|
||||
event=ReactStartEvent(
|
||||
event_type=ReActEventType.REACT_START,
|
||||
timestamp=datetime(2026, 3, 18, 12, 0, 0),
|
||||
trace_id="trace-1",
|
||||
func_name="chat_impl",
|
||||
iteration=0,
|
||||
user_task_prompt="make a cube",
|
||||
initial_messages=[],
|
||||
available_tools=[],
|
||||
),
|
||||
origin=_origin(1),
|
||||
),
|
||||
ResponseYield(
|
||||
response="hello", messages=[{"role": "assistant", "content": "hello"}]
|
||||
),
|
||||
EventYield(
|
||||
event=ReactEndEvent(
|
||||
event_type=ReActEventType.REACT_END,
|
||||
timestamp=datetime(2026, 3, 18, 12, 0, 1),
|
||||
trace_id="trace-1",
|
||||
func_name="chat_impl",
|
||||
iteration=1,
|
||||
final_response="hello",
|
||||
final_messages=[{"role": "assistant", "content": "hello"}],
|
||||
total_iterations=1,
|
||||
total_execution_time=0.5,
|
||||
total_tool_calls=0,
|
||||
total_llm_calls=1,
|
||||
),
|
||||
origin=_origin(2),
|
||||
),
|
||||
]
|
||||
agent = _FakeAgent(outputs)
|
||||
|
||||
packets = [
|
||||
packet
|
||||
async for packet in stream_chat_events(
|
||||
request,
|
||||
cast(Any, conversation),
|
||||
cast(Any, agent),
|
||||
)
|
||||
]
|
||||
|
||||
assert packets[0].startswith("event: react_start\n")
|
||||
assert any(packet.startswith("event: response\n") for packet in packets)
|
||||
assert any('"delta_text": "hello"' in packet for packet in packets)
|
||||
assert packets[-1].startswith("event: done\n")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_agent_response_aggregates_text_from_response_yields_only():
|
||||
conversation = _FakeConversation()
|
||||
outputs = [
|
||||
EventYield(
|
||||
event=ReactStartEvent(
|
||||
event_type=ReActEventType.REACT_START,
|
||||
timestamp=datetime(2026, 3, 18, 12, 0, 0),
|
||||
trace_id="trace-1",
|
||||
func_name="chat_impl",
|
||||
iteration=0,
|
||||
user_task_prompt="make a cube",
|
||||
initial_messages=[],
|
||||
available_tools=[],
|
||||
),
|
||||
origin=_origin(1),
|
||||
),
|
||||
ResponseYield(response="hello", messages=[]),
|
||||
ResponseYield(response=" world", messages=[]),
|
||||
]
|
||||
agent = _FakeAgent(outputs)
|
||||
|
||||
full_response, prompt_tokens, completion_tokens = await process_agent_response(
|
||||
"make a cube", cast(Any, conversation), cast(Any, agent)
|
||||
)
|
||||
|
||||
assert full_response == "hello world"
|
||||
assert prompt_tokens is None
|
||||
assert completion_tokens is None
|
||||
assert conversation.persisted is True
|
||||
|
||||
|
||||
def test_api_client_parse_sse_lines_understands_event_and_data_frames():
|
||||
lines = [
|
||||
b"event: response",
|
||||
b'data: {"delta_text": "hello"}',
|
||||
b"",
|
||||
b"event: tool_call_start",
|
||||
b'data: {"event_type": "tool_call_start"}',
|
||||
b"",
|
||||
b"event: done",
|
||||
b'data: {"ok": true}',
|
||||
b"",
|
||||
]
|
||||
|
||||
packets = list(_parse_sse_lines(lines))
|
||||
|
||||
assert packets == [
|
||||
{"event": "response", "data": {"delta_text": "hello"}},
|
||||
{"event": "tool_call_start", "data": {"event_type": "tool_call_start"}},
|
||||
{"event": "done", "data": {"ok": True}},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_chat_events_propagates_conversation_session(monkeypatch):
|
||||
request = _request()
|
||||
conversation = _FakeConversation()
|
||||
outputs = [ResponseYield(response="hello", messages=[])]
|
||||
agent = _FakeAgent(outputs)
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
@contextmanager
|
||||
def fake_propagate_conversation_session(**kwargs):
|
||||
captured.update(kwargs)
|
||||
yield
|
||||
|
||||
monkeypatch.setattr(
|
||||
chat_router_module,
|
||||
"propagate_conversation_session",
|
||||
fake_propagate_conversation_session,
|
||||
)
|
||||
|
||||
packets = [
|
||||
packet
|
||||
async for packet in stream_chat_events(
|
||||
request,
|
||||
cast(Any, conversation),
|
||||
cast(Any, agent),
|
||||
)
|
||||
]
|
||||
|
||||
assert any(packet.startswith("event: response\n") for packet in packets)
|
||||
assert captured["conversation_id"] == "conversation-1"
|
||||
assert captured["tags"] == ["cadagent", "event_stream"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_agent_response_propagates_conversation_session(monkeypatch):
|
||||
conversation = _FakeConversation()
|
||||
agent = _FakeAgent([ResponseYield(response="hello", messages=[])])
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
@contextmanager
|
||||
def fake_propagate_conversation_session(**kwargs):
|
||||
captured.update(kwargs)
|
||||
yield
|
||||
|
||||
monkeypatch.setattr(
|
||||
utils_module,
|
||||
"propagate_conversation_session",
|
||||
fake_propagate_conversation_session,
|
||||
)
|
||||
|
||||
full_response, _, _ = await process_agent_response(
|
||||
"make a cube", cast(Any, conversation), cast(Any, agent)
|
||||
)
|
||||
|
||||
assert full_response == "hello"
|
||||
assert captured["conversation_id"] == "conversation-1"
|
||||
assert captured["tags"] == ["cadagent", "non_stream"]
|
||||
Reference in New Issue
Block a user