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_provider(tmp_path: Path) -> tuple[TestClient, FakeClient]: provider = 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=provider) # type: ignore[arg-type] with TestClient(app) as client: yield client, provider 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_provider: tuple[TestClient, FakeClient], tmp_path: Path ) -> None: client, _ = client_and_provider 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_provider: tuple[TestClient, FakeClient] ) -> None: client, _ = client_and_provider 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_provider: tuple[TestClient, FakeClient] ) -> None: client, _ = client_and_provider 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_provider: tuple[TestClient, FakeClient], tmp_path: Path ) -> None: client, provider = client_and_provider 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 provider.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_provider: tuple[TestClient, FakeClient] ) -> None: client, _ = client_and_provider 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_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.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_provider: tuple[TestClient, FakeClient], tmp_path: Path ) -> None: client, _ = client_and_provider 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_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: client, _ = client_and_provider 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: provider = FakeClient() provider.completions.delay = 0.01 storage = JsonSessionStorage(tmp_path) session = await storage.create("system") service = ChatService( storage=storage, client=provider, # 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 provider.completions.calls[1]["messages"] == [ {"role": "system", "content": "system"}, {"role": "user", "content": "first"}, {"role": "assistant", "content": "reply-1"}, {"role": "user", "content": "second"}, ] asyncio.run(scenario())