添加获取历史对话数据的接口

This commit is contained in:
2026-07-03 19:16:57 +08:00
parent 92497d35a5
commit 83b55c0c72
5 changed files with 104 additions and 183 deletions
+60 -93
View File
@@ -19,7 +19,6 @@ class FakeCompletions:
def __init__(self) -> None:
self.calls: list[dict[str, Any]] = []
self.error: Exception | None = None
self.stream_error: Exception | None = None
self.delay = 0.0
async def create(self, **kwargs: Any) -> SimpleNamespace:
@@ -28,8 +27,6 @@ class FakeCompletions:
await asyncio.sleep(self.delay)
if self.error:
raise self.error
if kwargs.get("stream"):
return FakeStream(self.stream_error)
return SimpleNamespace(
choices=[
SimpleNamespace(
@@ -42,37 +39,6 @@ class FakeCompletions:
total_tokens=12,
),
)
class FakeStream:
def __init__(self, error: Exception | None = None) -> None:
self.error = error
self.closed = False
async def close(self) -> None:
self.closed = True
async def __aiter__(self):
yield SimpleNamespace(
choices=[SimpleNamespace(delta=SimpleNamespace(content="streamed "))],
usage=None,
)
if self.error:
raise self.error
yield SimpleNamespace(
choices=[SimpleNamespace(delta=SimpleNamespace(content="reply"))],
usage=None,
)
yield SimpleNamespace(
choices=[],
usage=SimpleNamespace(
prompt_tokens=11,
completion_tokens=3,
total_tokens=14,
),
)
class FakeClient:
def __init__(self) -> None:
self.completions = FakeCompletions()
@@ -216,65 +182,6 @@ def test_corrupt_session_returns_500(
assert response.status_code == 500
def test_streaming_response_is_sse_and_persists_complete_message(
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, provider = client_and_provider
session_id = create_session(client, system_prompt="system")
response = client.post(
f"/sessions/{session_id}/messages",
json={"content": "hello", "stream": True},
)
assert response.status_code == 200
assert response.headers["content-type"].startswith("text/event-stream")
payloads = [
line.removeprefix("data: ")
for line in response.text.splitlines()
if line.startswith("data: ")
]
assert [json.loads(payload)["type"] for payload in payloads[:-1]] == [
"delta",
"delta",
"done",
]
assert json.loads(payloads[-2])["usage"]["total_tokens"] == 14
assert payloads[-1] == "[DONE]"
assert provider.completions.calls[0]["stream"] is True
assert provider.completions.calls[0]["stream_options"] == {
"include_usage": True
}
data = json.loads((tmp_path / f"{session_id}.json").read_text())
assert data["messages"][-1]["role"] == "assistant"
assert data["messages"][-1]["content"] == "streamed reply"
assert data["messages"][-1]["created_at"].endswith("Z")
done_message = json.loads(payloads[-2])["message"]
assert done_message["created_at"] == data["messages"][-1]["created_at"]
def test_streaming_failure_emits_error_and_does_not_persist(
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, provider = client_and_provider
session_id = create_session(client)
before = (tmp_path / f"{session_id}.json").read_text()
provider.completions.stream_error = APIConnectionError( # type: ignore[arg-type]
request=object()
)
response = client.post(
f"/sessions/{session_id}/messages",
json={"content": "hello", "stream": True},
)
assert response.status_code == 200
assert '"type": "error"' in response.text
assert response.text.rstrip().endswith("data: [DONE]")
assert (tmp_path / f"{session_id}.json").read_text() == before
def test_legacy_messages_without_timestamp_remain_usable(
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
@@ -294,6 +201,66 @@ def test_legacy_messages_without_timestamp_remain_usable(
assert all("created_at" in message for message in migrated["messages"])
def test_stream_parameter_is_rejected(
client_and_provider: tuple[TestClient, FakeClient]
) -> None:
client, provider = client_and_provider
session_id = create_session(client)
response = client.post(
f"/sessions/{session_id}/messages",
json={"content": "hello", "stream": True},
)
assert response.status_code == 422
assert provider.completions.calls == []
def test_get_session_history_returns_all_messages(
client_and_provider: tuple[TestClient, FakeClient]
) -> None:
client, _ = client_and_provider
session_id = create_session(client, system_prompt="history prompt")
client.post(
f"/sessions/{session_id}/messages", json={"content": "first"}
)
client.post(
f"/sessions/{session_id}/messages", json={"content": "second"}
)
response = client.get(f"/sessions/{session_id}/messages")
assert response.status_code == 200
history = response.json()
assert history["session_id"] == session_id
assert history["system_prompt"] == "history prompt"
assert history["created_at"].endswith("Z")
assert history["updated_at"].endswith("Z")
assert [message["content"] for message in history["messages"]] == [
"first",
"reply-1",
"second",
"reply-2",
]
assert all(message["created_at"].endswith("Z") for message in history["messages"])
def test_get_session_history_handles_empty_and_missing_sessions(
client_and_provider: tuple[TestClient, FakeClient]
) -> None:
client, _ = client_and_provider
session_id = create_session(client)
empty = client.get(f"/sessions/{session_id}/messages")
missing = client.get(
"/sessions/00000000-0000-0000-0000-000000000000/messages"
)
assert empty.status_code == 200
assert empty.json()["messages"] == []
assert missing.status_code == 404
def test_same_session_concurrent_messages_are_serialized(tmp_path: Path) -> None:
async def scenario() -> None:
provider = FakeClient()