From 83b55c0c724d0902d71526cca6fcaa7e2ba7bbdb Mon Sep 17 00:00:00 2001 From: Mplan Date: Fri, 3 Jul 2026 19:16:57 +0800 Subject: [PATCH] =?UTF-8?q?=E6=B7=BB=E5=8A=A0=E8=8E=B7=E5=8F=96=E5=8E=86?= =?UTF-8?q?=E5=8F=B2=E5=AF=B9=E8=AF=9D=E6=95=B0=E6=8D=AE=E7=9A=84=E6=8E=A5?= =?UTF-8?q?=E5=8F=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 10 ++- app.py | 51 ++++++---------- models.py | 9 ++- service.py | 64 +++++-------------- tests/test_api.py | 153 ++++++++++++++++++---------------------------- 5 files changed, 104 insertions(+), 183 deletions(-) diff --git a/README.md b/README.md index d26ddd8..8619459 100644 --- a/README.md +++ b/README.md @@ -32,12 +32,10 @@ curl -X POST http://127.0.0.1:8000/sessions/SESSION_ID/messages \ -d '{"content":"你好,请记住我的名字是小明。"}' ``` -设置 `stream` 为 `true` 可通过 SSE 接收增量回复: +API 返回及存储的每条消息都包含 UTC `created_at`。只有模型完整响应成功后,本轮消息才会写入会话历史。 + +获取指定会话的完整历史: ```bash -curl -N -X POST http://127.0.0.1:8000/sessions/SESSION_ID/messages \ - -H 'Content-Type: application/json' \ - -d '{"content":"介绍一下你自己。","stream":true}' +curl http://127.0.0.1:8000/sessions/SESSION_ID/messages ``` - -流式响应依次返回 `delta` 事件、包含完整消息与 token 用量的 `done` 事件,最后以 `data: [DONE]` 结束。返回及存储的每条消息都包含 UTC `created_at`。只有完整响应成功结束后,本轮消息才会写入会话历史。 diff --git a/app.py b/app.py index 873a2cf..f881d09 100644 --- a/app.py +++ b/app.py @@ -1,10 +1,8 @@ from contextlib import asynccontextmanager from collections.abc import AsyncGenerator -import json from fastapi import FastAPI, HTTPException, Request, status from openai import AsyncOpenAI -from starlette.responses import JSONResponse, Response, StreamingResponse from config import Settings from models import ( @@ -12,23 +10,12 @@ from models import ( CreateSessionResponse, SendMessageRequest, SendMessageResponse, + SessionHistoryResponse, ) from service import ChatProviderError, ChatService from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError -async def encode_sse( - events: AsyncGenerator[dict[str, object], None], -) -> AsyncGenerator[str, None]: - try: - async for event in events: - yield f"data: {json.dumps(event, ensure_ascii=False)}\n\n" - except (ChatProviderError, SessionStorageError) as exc: - error = {"type": "error", "detail": str(exc)} - yield f"data: {json.dumps(error, ensure_ascii=False)}\n\n" - yield "data: [DONE]\n\n" - - def create_app( settings: Settings | None = None, client: AsyncOpenAI | None = None, @@ -83,30 +70,13 @@ def create_app( @app.post( "/sessions/{session_id}/messages", response_model=SendMessageResponse, - responses={ - 200: { - "content": {"text/event-stream": {}}, - "description": "JSON response or SSE stream when stream=true", - } - }, ) async def send_message( session_id: str, body: SendMessageRequest, request: Request - ) -> Response: + ) -> SendMessageResponse: service: ChatService = request.app.state.chat_service try: - if body.stream: - await service.ensure_session_exists(session_id) - return StreamingResponse( - encode_sse( - service.generate_response_stream(session_id, body.content) - ), - media_type="text/event-stream", - headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, - ) - - response = await service.generate_response(session_id, body.content) - return JSONResponse(content=response.model_dump(mode="json")) + return await service.generate_response(session_id, body.content) except SessionNotFoundError as exc: raise HTTPException(status_code=404, detail="session not found") from exc except ChatProviderError as exc: @@ -114,6 +84,21 @@ def create_app( except SessionStorageError as exc: raise HTTPException(status_code=500, detail=str(exc)) from exc + @app.get( + "/sessions/{session_id}/messages", + response_model=SessionHistoryResponse, + ) + async def get_session_history( + session_id: str, request: Request + ) -> SessionHistoryResponse: + service: ChatService = request.app.state.chat_service + try: + return await service.get_session_history(session_id) + except SessionNotFoundError as exc: + raise HTTPException(status_code=404, detail="session not found") from exc + except SessionStorageError as exc: + raise HTTPException(status_code=500, detail=str(exc)) from exc + return app diff --git a/models.py b/models.py index a924145..5164e22 100644 --- a/models.py +++ b/models.py @@ -57,7 +57,6 @@ class SendMessageRequest(BaseModel): model_config = ConfigDict(extra="forbid") content: str - stream: bool = False @field_validator("content") @classmethod @@ -77,3 +76,11 @@ class SendMessageResponse(BaseModel): session_id: str message: Message usage: TokenUsage | None + + +class SessionHistoryResponse(BaseModel): + session_id: str + system_prompt: str + created_at: datetime + updated_at: datetime + messages: list[Message] diff --git a/service.py b/service.py index 2de31a5..e90d0ea 100644 --- a/service.py +++ b/service.py @@ -1,9 +1,14 @@ import asyncio -from collections.abc import AsyncGenerator from datetime import UTC, datetime from openai import APIError, AsyncOpenAI -from models import Message, SendMessageResponse, Session, TokenUsage +from models import ( + Message, + SendMessageResponse, + Session, + SessionHistoryResponse, + TokenUsage, +) from storage import JsonSessionStorage @@ -36,7 +41,6 @@ class ChatService: completion = await self.client.chat.completions.create( model=self.model, messages=api_messages, # type: ignore[arg-type] - stream=False, extra_body={"thinking": {"type": "disabled"}}, ) assistant_content = completion.choices[0].message.content @@ -57,57 +61,17 @@ class ChatService: usage=usage, ) - async def ensure_session_exists(self, session_id: str) -> None: - await self.storage.read(session_id) - - async def generate_response_stream( - self, session_id: str, content: str - ) -> AsyncGenerator[dict[str, object], None]: + async def get_session_history(self, session_id: str) -> SessionHistoryResponse: lock = self._locks.setdefault(session_id, asyncio.Lock()) async with lock: - user_created_at = datetime.now(UTC) session = await self.storage.read(session_id) - api_messages = self._build_api_messages(session, content) - - try: - stream = await self.client.chat.completions.create( - model=self.model, - messages=api_messages, # type: ignore[arg-type] - stream=True, - stream_options={"include_usage": True}, - extra_body={"thinking": {"type": "disabled"}}, - ) - parts: list[str] = [] - usage: TokenUsage | None = None - - try: - async for chunk in stream: - chunk_usage = getattr(chunk, "usage", None) - if chunk_usage is not None: - usage = self._parse_usage(chunk_usage) - - for choice in chunk.choices: - delta = choice.delta.content - if delta: - parts.append(delta) - yield {"type": "delta", "content": delta} - finally: - await stream.close() - except (APIError, AttributeError, TypeError) as exc: - raise ChatProviderError("Upstream chat stream failed") from exc - - assistant_content = "".join(parts) - if not assistant_content: - raise ChatProviderError("Upstream chat provider returned an empty response") - - assistant_message = await self._save_exchange( - session, content, user_created_at, assistant_content + return SessionHistoryResponse( + session_id=session.session_id, + system_prompt=session.system_prompt, + created_at=session.created_at, + updated_at=session.updated_at, + messages=session.messages, ) - yield { - "type": "done", - "message": assistant_message.model_dump(mode="json"), - "usage": usage.model_dump() if usage is not None else None, - } @staticmethod def _build_api_messages( diff --git a/tests/test_api.py b/tests/test_api.py index c1a9c02..ac057b9 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -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()