第一次提交

This commit is contained in:
2026-06-29 23:48:53 +08:00
commit 8714fb07c4
13 changed files with 1471 additions and 0 deletions
+321
View File
@@ -0,0 +1,321 @@
import asyncio
from datetime import datetime
import json
from pathlib import Path
from types import SimpleNamespace
from typing import Any
from fastapi.testclient import TestClient
from openai import APIConnectionError
import pytest
from app import create_app
from config import Settings
from service import ChatService
from storage import JsonSessionStorage
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:
self.calls.append(kwargs)
if self.delay:
await asyncio.sleep(self.delay)
if self.error:
raise self.error
if kwargs.get("stream"):
return FakeStream(self.stream_error)
return SimpleNamespace(
choices=[
SimpleNamespace(
message=SimpleNamespace(content=f"reply-{len(self.calls)}")
)
],
usage=SimpleNamespace(
prompt_tokens=10,
completion_tokens=2,
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()
self.chat = SimpleNamespace(completions=self.completions)
@pytest.fixture
def client_and_deepseek(tmp_path: Path) -> tuple[TestClient, FakeClient]:
deepseek = FakeClient()
settings = Settings(
api_key=None,
base_url="https://api.deepseek.com",
model="deepseek-v4-flash",
default_system_prompt="default prompt",
data_dir=tmp_path,
)
app = create_app(settings=settings, client=deepseek) # type: ignore[arg-type]
with TestClient(app) as client:
yield client, deepseek
def create_session(client: TestClient, **body: str) -> str:
response = client.post("/sessions", json=body)
assert response.status_code == 201
return response.json()["session_id"]
def test_create_session_uses_default_prompt(
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, _ = client_and_deepseek
session_id = create_session(client)
data = json.loads((tmp_path / f"{session_id}.json").read_text())
assert data["system_prompt"] == "default prompt"
assert data["messages"] == []
def test_create_session_without_body(
client_and_deepseek: tuple[TestClient, FakeClient]
) -> None:
client, _ = client_and_deepseek
response = client.post("/sessions")
assert response.status_code == 201
assert response.json()["system_prompt"] == "default prompt"
def test_create_session_accepts_custom_prompt(
client_and_deepseek: tuple[TestClient, FakeClient]
) -> None:
client, _ = client_and_deepseek
response = client.post("/sessions", json={"system_prompt": "回答中文"})
assert response.status_code == 201
assert response.json()["system_prompt"] == "回答中文"
def test_multi_turn_request_contains_full_history(
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, deepseek = client_and_deepseek
session_id = create_session(client, system_prompt="system")
first = client.post(
f"/sessions/{session_id}/messages", json={"content": "first"}
)
second = client.post(
f"/sessions/{session_id}/messages", json={"content": "second"}
)
assert first.status_code == second.status_code == 200
assert deepseek.completions.calls[1]["messages"] == [
{"role": "system", "content": "system"},
{"role": "user", "content": "first"},
{"role": "assistant", "content": "reply-1"},
{"role": "user", "content": "second"},
]
assert second.json()["usage"] == {
"prompt_tokens": 10,
"completion_tokens": 2,
"total_tokens": 12,
}
datetime.fromisoformat(second.json()["message"]["created_at"])
data = json.loads((tmp_path / f"{session_id}.json").read_text())
assert [message["content"] for message in data["messages"]] == [
"first",
"reply-1",
"second",
"reply-2",
]
assert all(message["created_at"].endswith("Z") for message in data["messages"])
def test_invalid_session_and_blank_message(
client_and_deepseek: tuple[TestClient, FakeClient]
) -> None:
client, _ = client_and_deepseek
missing = client.post(
"/sessions/00000000-0000-0000-0000-000000000000/messages",
json={"content": "hello"},
)
blank = client.post(
"/sessions/00000000-0000-0000-0000-000000000000/messages",
json={"content": " "},
)
assert missing.status_code == 404
assert blank.status_code == 422
def test_upstream_failure_does_not_change_history(
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, deepseek = client_and_deepseek
session_id = create_session(client)
before = (tmp_path / f"{session_id}.json").read_text()
deepseek.completions.error = APIConnectionError(request=object()) # type: ignore[arg-type]
response = client.post(
f"/sessions/{session_id}/messages", json={"content": "hello"}
)
assert response.status_code == 502
assert (tmp_path / f"{session_id}.json").read_text() == before
def test_corrupt_session_returns_500(
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, _ = client_and_deepseek
session_id = create_session(client)
(tmp_path / f"{session_id}.json").write_text("not-json")
response = client.post(
f"/sessions/{session_id}/messages", json={"content": "hello"}
)
assert response.status_code == 500
def test_streaming_response_is_sse_and_persists_complete_message(
client_and_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, deepseek = client_and_deepseek
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 deepseek.completions.calls[0]["stream"] is True
assert deepseek.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_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, deepseek = client_and_deepseek
session_id = create_session(client)
before = (tmp_path / f"{session_id}.json").read_text()
deepseek.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_deepseek: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, _ = client_and_deepseek
session_id = create_session(client)
path = tmp_path / f"{session_id}.json"
data = json.loads(path.read_text())
data["messages"] = [{"role": "user", "content": "legacy"}]
path.write_text(json.dumps(data))
response = client.post(
f"/sessions/{session_id}/messages", json={"content": "new"}
)
assert response.status_code == 200
migrated = json.loads(path.read_text())
assert all("created_at" in message for message in migrated["messages"])
def test_same_session_concurrent_messages_are_serialized(tmp_path: Path) -> None:
async def scenario() -> None:
deepseek = FakeClient()
deepseek.completions.delay = 0.01
storage = JsonSessionStorage(tmp_path)
session = await storage.create("system")
service = ChatService(
storage=storage,
client=deepseek, # type: ignore[arg-type]
model="deepseek-v4-flash",
)
await asyncio.gather(
service.generate_response(session.session_id, "first"),
service.generate_response(session.session_id, "second"),
)
assert deepseek.completions.calls[1]["messages"] == [
{"role": "system", "content": "system"},
{"role": "user", "content": "first"},
{"role": "assistant", "content": "reply-1"},
{"role": "user", "content": "second"},
]
asyncio.run(scenario())