80 lines
2.2 KiB
Python
80 lines
2.2 KiB
Python
|
|
"""Tests for `app.engine.sse.SseEmitter`."""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import asyncio
|
||
|
|
|
||
|
|
import pytest
|
||
|
|
|
||
|
|
from app.engine.sse import SseEmitter
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_emitter_yields_emitted_events():
|
||
|
|
emitter = SseEmitter()
|
||
|
|
await emitter.emit("tool_call", {"tool": "calc"})
|
||
|
|
await emitter.emit("phase_end", {"phase": 1})
|
||
|
|
await emitter.done({"result": "ok"})
|
||
|
|
|
||
|
|
events = []
|
||
|
|
async for evt in emitter.stream():
|
||
|
|
events.append(evt)
|
||
|
|
assert len(events) == 3
|
||
|
|
assert events[0]["event"] == "tool_call"
|
||
|
|
assert events[1]["event"] == "phase_end"
|
||
|
|
assert events[2]["event"] == "done"
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_emitter_event_ids_increment():
|
||
|
|
emitter = SseEmitter()
|
||
|
|
await emitter.emit("a", {})
|
||
|
|
await emitter.emit("b", {})
|
||
|
|
events = []
|
||
|
|
async for evt in emitter.stream():
|
||
|
|
events.append(evt)
|
||
|
|
# IDs are evt_1, evt_2 (sentinel None doesn't yield an event)
|
||
|
|
assert events[0]["id"] == "evt_1"
|
||
|
|
assert events[1]["id"] == "evt_2"
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_emitter_error_closes_stream():
|
||
|
|
emitter = SseEmitter()
|
||
|
|
await emitter.error("test_error", "something went wrong", details={"k": "v"})
|
||
|
|
events = []
|
||
|
|
async for evt in emitter.stream():
|
||
|
|
events.append(evt)
|
||
|
|
assert len(events) == 1
|
||
|
|
assert events[0]["event"] == "error"
|
||
|
|
import json
|
||
|
|
payload = json.loads(events[0]["data"])
|
||
|
|
assert payload["code"] == "test_error"
|
||
|
|
assert payload["details"] == {"k": "v"}
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_emitter_done_with_no_result():
|
||
|
|
emitter = SseEmitter()
|
||
|
|
await emitter.done()
|
||
|
|
events = []
|
||
|
|
async for evt in emitter.stream():
|
||
|
|
events.append(evt)
|
||
|
|
assert len(events) == 1
|
||
|
|
assert events[0]["event"] == "done"
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.asyncio
|
||
|
|
async def test_emitter_serializes_non_serializable_data():
|
||
|
|
emitter = SseEmitter()
|
||
|
|
# A set is not JSON-serializable
|
||
|
|
await emitter.emit("test", {"data": {1, 2, 3}})
|
||
|
|
await emitter.done()
|
||
|
|
events = []
|
||
|
|
async for evt in emitter.stream():
|
||
|
|
events.append(evt)
|
||
|
|
# The first event should still have valid JSON (with serialization fallback)
|
||
|
|
import json
|
||
|
|
payload = json.loads(events[0]["data"])
|
||
|
|
assert "error" in payload or "data" in payload
|