添加获取历史对话数据的接口
This commit is contained in:
+60
-93
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user