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