添加获取历史对话数据的接口

This commit is contained in:
2026-07-03 19:16:57 +08:00
parent 92497d35a5
commit 83b55c0c72
5 changed files with 104 additions and 183 deletions
+4 -6
View File
@@ -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`。只有完整响应成功结束后,本轮消息才会写入会话历史。
+18 -33
View File
@@ -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
+8 -1
View File
@@ -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]
+14 -50
View File
@@ -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"}},
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,
)
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
)
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(
+60 -93
View File
@@ -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()