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

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":"你好,请记住我的名字是小明。"}' -d '{"content":"你好,请记住我的名字是小明。"}'
``` ```
设置 `stream``true` 可通过 SSE 接收增量回复: API 返回及存储的每条消息都包含 UTC `created_at`。只有模型完整响应成功后,本轮消息才会写入会话历史。
获取指定会话的完整历史:
```bash ```bash
curl -N -X POST http://127.0.0.1:8000/sessions/SESSION_ID/messages \ curl http://127.0.0.1:8000/sessions/SESSION_ID/messages
-H 'Content-Type: application/json' \
-d '{"content":"介绍一下你自己。","stream":true}'
``` ```
流式响应依次返回 `delta` 事件、包含完整消息与 token 用量的 `done` 事件,最后以 `data: [DONE]` 结束。返回及存储的每条消息都包含 UTC `created_at`。只有完整响应成功结束后,本轮消息才会写入会话历史。
+18 -33
View File
@@ -1,10 +1,8 @@
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from collections.abc import AsyncGenerator from collections.abc import AsyncGenerator
import json
from fastapi import FastAPI, HTTPException, Request, status from fastapi import FastAPI, HTTPException, Request, status
from openai import AsyncOpenAI from openai import AsyncOpenAI
from starlette.responses import JSONResponse, Response, StreamingResponse
from config import Settings from config import Settings
from models import ( from models import (
@@ -12,23 +10,12 @@ from models import (
CreateSessionResponse, CreateSessionResponse,
SendMessageRequest, SendMessageRequest,
SendMessageResponse, SendMessageResponse,
SessionHistoryResponse,
) )
from service import ChatProviderError, ChatService from service import ChatProviderError, ChatService
from storage import JsonSessionStorage, SessionNotFoundError, SessionStorageError 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( def create_app(
settings: Settings | None = None, settings: Settings | None = None,
client: AsyncOpenAI | None = None, client: AsyncOpenAI | None = None,
@@ -83,30 +70,13 @@ def create_app(
@app.post( @app.post(
"/sessions/{session_id}/messages", "/sessions/{session_id}/messages",
response_model=SendMessageResponse, response_model=SendMessageResponse,
responses={
200: {
"content": {"text/event-stream": {}},
"description": "JSON response or SSE stream when stream=true",
}
},
) )
async def send_message( async def send_message(
session_id: str, body: SendMessageRequest, request: Request session_id: str, body: SendMessageRequest, request: Request
) -> Response: ) -> SendMessageResponse:
service: ChatService = request.app.state.chat_service service: ChatService = request.app.state.chat_service
try: try:
if body.stream: return await service.generate_response(session_id, body.content)
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"))
except SessionNotFoundError as exc: except SessionNotFoundError as exc:
raise HTTPException(status_code=404, detail="session not found") from exc raise HTTPException(status_code=404, detail="session not found") from exc
except ChatProviderError as exc: except ChatProviderError as exc:
@@ -114,6 +84,21 @@ def create_app(
except SessionStorageError as exc: except SessionStorageError as exc:
raise HTTPException(status_code=500, detail=str(exc)) from 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 return app
+8 -1
View File
@@ -57,7 +57,6 @@ class SendMessageRequest(BaseModel):
model_config = ConfigDict(extra="forbid") model_config = ConfigDict(extra="forbid")
content: str content: str
stream: bool = False
@field_validator("content") @field_validator("content")
@classmethod @classmethod
@@ -77,3 +76,11 @@ class SendMessageResponse(BaseModel):
session_id: str session_id: str
message: Message message: Message
usage: TokenUsage | None 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 import asyncio
from collections.abc import AsyncGenerator
from datetime import UTC, datetime from datetime import UTC, datetime
from openai import APIError, AsyncOpenAI from openai import APIError, AsyncOpenAI
from models import Message, SendMessageResponse, Session, TokenUsage from models import (
Message,
SendMessageResponse,
Session,
SessionHistoryResponse,
TokenUsage,
)
from storage import JsonSessionStorage from storage import JsonSessionStorage
@@ -36,7 +41,6 @@ class ChatService:
completion = await self.client.chat.completions.create( completion = await self.client.chat.completions.create(
model=self.model, model=self.model,
messages=api_messages, # type: ignore[arg-type] messages=api_messages, # type: ignore[arg-type]
stream=False,
extra_body={"thinking": {"type": "disabled"}}, extra_body={"thinking": {"type": "disabled"}},
) )
assistant_content = completion.choices[0].message.content assistant_content = completion.choices[0].message.content
@@ -57,57 +61,17 @@ class ChatService:
usage=usage, usage=usage,
) )
async def ensure_session_exists(self, session_id: str) -> None: async def get_session_history(self, session_id: str) -> SessionHistoryResponse:
await self.storage.read(session_id)
async def generate_response_stream(
self, session_id: str, content: str
) -> AsyncGenerator[dict[str, object], None]:
lock = self._locks.setdefault(session_id, asyncio.Lock()) lock = self._locks.setdefault(session_id, asyncio.Lock())
async with lock: async with lock:
user_created_at = datetime.now(UTC)
session = await self.storage.read(session_id) session = await self.storage.read(session_id)
api_messages = self._build_api_messages(session, content) return SessionHistoryResponse(
session_id=session.session_id,
try: system_prompt=session.system_prompt,
stream = await self.client.chat.completions.create( created_at=session.created_at,
model=self.model, updated_at=session.updated_at,
messages=api_messages, # type: ignore[arg-type] messages=session.messages,
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
) )
yield {
"type": "done",
"message": assistant_message.model_dump(mode="json"),
"usage": usage.model_dump() if usage is not None else None,
}
@staticmethod @staticmethod
def _build_api_messages( def _build_api_messages(
+60 -93
View File
@@ -19,7 +19,6 @@ class FakeCompletions:
def __init__(self) -> None: def __init__(self) -> None:
self.calls: list[dict[str, Any]] = [] self.calls: list[dict[str, Any]] = []
self.error: Exception | None = None self.error: Exception | None = None
self.stream_error: Exception | None = None
self.delay = 0.0 self.delay = 0.0
async def create(self, **kwargs: Any) -> SimpleNamespace: async def create(self, **kwargs: Any) -> SimpleNamespace:
@@ -28,8 +27,6 @@ class FakeCompletions:
await asyncio.sleep(self.delay) await asyncio.sleep(self.delay)
if self.error: if self.error:
raise self.error raise self.error
if kwargs.get("stream"):
return FakeStream(self.stream_error)
return SimpleNamespace( return SimpleNamespace(
choices=[ choices=[
SimpleNamespace( SimpleNamespace(
@@ -42,37 +39,6 @@ class FakeCompletions:
total_tokens=12, 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: class FakeClient:
def __init__(self) -> None: def __init__(self) -> None:
self.completions = FakeCompletions() self.completions = FakeCompletions()
@@ -216,65 +182,6 @@ def test_corrupt_session_returns_500(
assert response.status_code == 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( def test_legacy_messages_without_timestamp_remain_usable(
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
) -> None: ) -> None:
@@ -294,6 +201,66 @@ def test_legacy_messages_without_timestamp_remain_usable(
assert all("created_at" in message for message in migrated["messages"]) 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: def test_same_session_concurrent_messages_are_serialized(tmp_path: Path) -> None:
async def scenario() -> None: async def scenario() -> None:
provider = FakeClient() provider = FakeClient()