117 lines
3.5 KiB
Python
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 == "大模型服务返回消息不完整"
|