89 lines
3.0 KiB
Python
89 lines
3.0 KiB
Python
from app.api.views.agent_run_sse import encode_events, recovery_events
|
|||
|
|
|
||
|
|
|
||
|
|
def test_first_connection_emits_unicode_deltas_only_after_commit():
|
||
|
|
content = "基金😀\n结果" * 100
|
||
|
|
events = recovery_events(run_id="r", trace_id="t", status="succeeded", error_code=None,
|
||
|
|
content=content, tool_calls=None, replay=False, chunk_size=11)
|
||
|
|
assert events[0][0] == "start" and events[-1][0] == "done"
|
||
|
|
assert "".join(payload["content"] for name, payload in events if name == "delta") == content
|
||
|
|
assert all(len(payload["content"]) <= 11 for name, payload in events if name == "delta")
|
||
|
|
ids = [line for event in encode_events("r", events)
|
||
|
|
for line in event.splitlines() if line.startswith("id:")]
|
||
|
|
assert len(set(ids)) == len(ids)
|
||
|
|
pending = recovery_events(run_id="r", trace_id="t", status="running", error_code=None,
|
||
|
|
content=content, tool_calls=None, replay=False)
|
||
|
|
assert [name for name, _ in pending] == ["start"]
|
||
|
|
|
||
|
|
|
||
|
|
def test_empty_result_still_emits_delta():
|
||
|
|
events = recovery_events(run_id="r", trace_id="t", status="succeeded", error_code=None,
|
||
|
|
content="", tool_calls=None, replay=False)
|
||
|
|
assert [name for name, _ in events] == ["start", "delta", "done"]
|
||
|
|
|
||
|
|
|
||
|
|
def test_pending_run_never_emits_result_events() -> None:
|
||
|
|
events = recovery_events(
|
||
|
|
run_id="run",
|
||
|
|
trace_id="trace",
|
||
|
|
status="running",
|
||
|
|
error_code=None,
|
||
|
|
content=None,
|
||
|
|
tool_calls=None,
|
||
|
|
)
|
||
|
|
assert [name for name, _ in events] == ["start"]
|
||
|
|
|
||
|
|
|
||
|
|
def test_succeeded_run_emits_tools_and_result_after_terminal_state() -> None:
|
||
|
|
events = recovery_events(
|
||
|
|
run_id="run",
|
||
|
|
trace_id="trace",
|
||
|
|
status="succeeded",
|
||
|
|
error_code=None,
|
||
|
|
content="answer",
|
||
|
|
tool_calls={"calls": []},
|
||
|
|
)
|
||
|
|
assert [name for name, _ in events] == ["start", "tools", "replace", "done"]
|
||
|
|
|
||
|
|
|
||
|
|
def test_failed_run_emits_error_without_result() -> None:
|
||
|
|
events = recovery_events(
|
||
|
|
run_id="run",
|
||
|
|
trace_id="trace",
|
||
|
|
status="failed",
|
||
|
|
error_code="MODEL_TIMEOUT",
|
||
|
|
content=None,
|
||
|
|
tool_calls=None,
|
||
|
|
)
|
||
|
|
assert [name for name, _ in events] == ["start", "error"]
|
||
|
|
|
||
|
|
|
||
|
|
def test_client_disconnect_does_not_change_recoverable_terminal_events() -> None:
|
||
|
|
events = recovery_events(
|
||
|
|
run_id="run",
|
||
|
|
trace_id="trace",
|
||
|
|
status="succeeded",
|
||
|
|
error_code=None,
|
||
|
|
content="answer",
|
||
|
|
tool_calls=None,
|
||
|
|
)
|
||
|
|
stream = encode_events("run", events)
|
||
|
|
first = next(stream)
|
||
|
|
assert "event: start" in first
|
||
|
|
del stream
|
||
|
|
recovered = list(
|
||
|
|
encode_events(
|
||
|
|
"run",
|
||
|
|
recovery_events(
|
||
|
|
run_id="run",
|
||
|
|
trace_id="trace",
|
||
|
|
status="succeeded",
|
||
|
|
error_code=None,
|
||
|
|
content="answer",
|
||
|
|
tool_calls=None,
|
||
|
|
),
|
||
|
|
)
|
||
|
|
)
|
||
|
|
assert any("event: replace" in item for item in recovered)
|
||
|
|
assert any("event: done" in item for item in recovered)
|