第一次提交
This commit is contained in:
@@ -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())
|
||||
Reference in New Issue
Block a user