添加获取历史对话数据的接口
This commit is contained in:
@@ -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`。只有完整响应成功结束后,本轮消息才会写入会话历史。
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
+60
-93
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user