Files
ZNJJ-api-server/test/api/test_chat_sse_contract.py
2026-07-27 17:21:29 +08:00

117 lines
3.5 KiB
Python

import json
import pytest
from src.api.endpoints import chat
from src.backends.chat import ChatInput, ChatResult, FormUpdate, TextDelta
from src.schemas.models import ProcessRequest_chat
def make_request(**overrides):
payload = {
"sessionId": "session-001",
"timeStamp": "20260726120000",
"text": "发生了交通事故",
"needFormUpdate": True,
}
payload.update(overrides)
return ProcessRequest_chat(**payload)
async def response_text(response):
chunks = []
async for chunk in response.body_iterator:
chunks.append(chunk.decode() if isinstance(chunk, bytes) else chunk)
return "".join(chunks)
def parse_sse(body):
events = []
for block in body.strip().split("\n\n"):
lines = block.splitlines()
event = lines[0].removeprefix("event: ")
data = json.loads(lines[1].removeprefix("data: "))
events.append((event, data))
return events
class OrderedBackend:
async def stream(self, chat_input: ChatInput):
yield TextDelta("<sta")
yield TextDelta("te>1002</state>")
yield FormUpdate({"jdcsl": 2})
yield TextDelta("第一句。")
yield TextDelta("第二句。")
async def complete(self, chat_input: ChatInput):
return ChatResult("<state>1002</state>第一句。第二句。", "1002", {"jdcsl": 2})
class MissingPrefixBackend:
async def stream(self, chat_input: ChatInput):
yield TextDelta("没有状态前缀")
async def complete(self, chat_input: ChatInput):
return ChatResult("没有状态前缀")
@pytest.mark.asyncio
async def test_sse_success_event_order_and_cardinality():
response = await chat(make_request(), stream=True, backend=OrderedBackend())
events = parse_sse(await response_text(response))
names = [name for name, _ in events]
assert names == [
"stage_code",
"formUpdate",
"text_delta",
"text_delta",
"done",
]
assert names.count("stage_code") == 1
assert names.count("formUpdate") == 1
assert names.count("done") == 1
assert "error" not in names
assert "".join(data["text"] for name, data in events if name == "text_delta") == (
"第一句。第二句。"
)
@pytest.mark.asyncio
async def test_use_text_chunk_only_changes_delta_boundaries():
request = make_request(useTextChunk=True)
response = await chat(request, stream=True, backend=OrderedBackend())
events = parse_sse(await response_text(response))
assert "".join(data["text"] for name, data in events if name == "text_delta") == (
"第一句。第二句。"
)
assert [name for name, _ in events].count("done") == 1
@pytest.mark.asyncio
async def test_missing_stream_prefix_characterizes_current_legacy_behavior():
response = await chat(
make_request(needFormUpdate=False),
stream=True,
backend=MissingPrefixBackend(),
)
events = parse_sse(await response_text(response))
assert [name for name, _ in events] == ["text_delta", "done"]
assert events[0][1]["text"] == "没有状态前缀"
@pytest.mark.asyncio
async def test_missing_non_stream_prefix_returns_compatible_business_error():
response = await chat(
make_request(needFormUpdate=False),
stream=False,
backend=MissingPrefixBackend(),
)
assert response.code == "500"
assert response.outputText == ""
assert response.nextStageCode == ""
assert response.msg == "大模型服务返回消息不完整"