增加鉴权系统

This commit is contained in:
2026-07-03 21:31:25 +08:00
parent 8073dd63a7
commit 7e65fadd83
12 changed files with 301 additions and 20 deletions
+1
View File
@@ -4,6 +4,7 @@ DEEPSEEK_BASE_URL=https://api.deepseek.com
DEEPSEEK_MODEL=deepseek-v4-flash DEEPSEEK_MODEL=deepseek-v4-flash
DEFAULT_SYSTEM_PROMPT=You are a helpful assistant. DEFAULT_SYSTEM_PROMPT=You are a helpful assistant.
CHAT_DATA_DIR=data CHAT_DATA_DIR=data
API_KEY_WHITELIST={"11111111-1111-4111-8111-111111111111":"replace-with-a-secret-key"}
MAX_TOOL_ROUNDS=5 MAX_TOOL_ROUNDS=5
MAX_TOOL_CALLS_PER_TURN=10 MAX_TOOL_CALLS_PER_TURN=10
TOOL_TIMEOUT_SECONDS=5 TOOL_TIMEOUT_SECONDS=5
+2
View File
@@ -5,6 +5,7 @@ build/
dist/ dist/
wheels/ wheels/
*.egg-info *.egg-info
.pytest_cache
# Virtual environments # Virtual environments
.venv .venv
@@ -12,3 +13,4 @@ wheels/
# Local configuration and persisted conversations # Local configuration and persisted conversations
.env .env
data/ data/
+20 -1
View File
@@ -5,6 +5,7 @@
## 代码结构 ## 代码结构
- `routes.py`:集中定义所有 HTTP 路由。 - `routes.py`:集中定义所有 HTTP 路由。
- `auth.py`:校验 API key 白名单并识别当前用户。
- `schemas.py`:集中定义 API 请求和响应结构。 - `schemas.py`:集中定义 API 请求和响应结构。
- `domain.py`:定义会话、消息和工具调用的持久化模型。 - `domain.py`:定义会话、消息和工具调用的持久化模型。
- `tools.py`:集中定义工具注册表、参数结构和工具函数。 - `tools.py`:集中定义工具注册表、参数结构和工具函数。
@@ -24,12 +25,19 @@ uv run uvicorn main:app --reload
服务会自动加载项目根目录的 `.env`,已有系统环境变量优先级更高。可配置 `DEEPSEEK_BASE_URL``DEEPSEEK_MODEL``DEFAULT_SYSTEM_PROMPT``CHAT_DATA_DIR` 和工具调用限制,完整示例见 `.env.example`。第一版应只使用一个 Uvicorn worker。 服务会自动加载项目根目录的 `.env`,已有系统环境变量优先级更高。可配置 `DEEPSEEK_BASE_URL``DEEPSEEK_MODEL``DEFAULT_SYSTEM_PROMPT``CHAT_DATA_DIR` 和工具调用限制,完整示例见 `.env.example`。第一版应只使用一个 Uvicorn worker。
所有业务接口都必须通过 `X-API-Key` 请求头提供访问密钥。白名单使用用户到密钥的 JSON 映射:
```text
API_KEY_WHITELIST={"11111111-1111-4111-8111-111111111111":"replace-with-a-secret-key"}
```
## 使用 ## 使用
创建会话: 创建会话:
```bash ```bash
curl -X POST http://127.0.0.1:8000/sessions \ curl -X POST http://127.0.0.1:8000/sessions \
-H 'X-API-Key: replace-with-a-secret-key' \
-H 'Content-Type: application/json' \ -H 'Content-Type: application/json' \
-d '{"system_prompt":"你是一个简洁的中文助手。"}' -d '{"system_prompt":"你是一个简洁的中文助手。"}'
``` ```
@@ -38,6 +46,7 @@ curl -X POST http://127.0.0.1:8000/sessions \
```bash ```bash
curl -X POST http://127.0.0.1:8000/sessions/SESSION_ID/messages \ curl -X POST http://127.0.0.1:8000/sessions/SESSION_ID/messages \
-H 'X-API-Key: replace-with-a-secret-key' \
-H 'Content-Type: application/json' \ -H 'Content-Type: application/json' \
-d '{"content":"你好,请记住我的名字是小明。"}' -d '{"content":"你好,请记住我的名字是小明。"}'
``` ```
@@ -53,5 +62,15 @@ API 返回及存储的每条消息都包含 UTC `created_at`。只有模型完
获取指定会话的完整历史: 获取指定会话的完整历史:
```bash ```bash
curl http://127.0.0.1:8000/sessions/SESSION_ID/messages curl http://127.0.0.1:8000/sessions/SESSION_ID/messages \
-H 'X-API-Key: replace-with-a-secret-key'
``` ```
查询当前用户的累计 API 调用次数和 token 用量:
```bash
curl http://127.0.0.1:8000/usage \
-H 'X-API-Key: replace-with-a-secret-key'
```
每个会话 JSON 顶层保存所属 `user_id`、成功聊天次数及累计 token。`/usage` 会扫描并汇总当前用户的全部会话,不使用独立统计文件。创建会话、查询历史和查询用量不计入 `api_calls`
+4
View File
@@ -4,6 +4,7 @@ from contextlib import asynccontextmanager
from fastapi import FastAPI from fastapi import FastAPI
from openai import AsyncOpenAI from openai import AsyncOpenAI
from auth import ApiKeyAuthenticator
from config import Settings from config import Settings
from routes import router from routes import router
from service import ChatService from service import ChatService
@@ -37,6 +38,9 @@ def create_app(
max_tool_rounds=resolved_settings.max_tool_rounds, max_tool_rounds=resolved_settings.max_tool_rounds,
max_tool_calls_per_turn=resolved_settings.max_tool_calls_per_turn, max_tool_calls_per_turn=resolved_settings.max_tool_calls_per_turn,
) )
app.state.authenticator = ApiKeyAuthenticator(
resolved_settings.api_key_whitelist
)
yield yield
if client is None: if client is None:
+45
View File
@@ -0,0 +1,45 @@
from dataclasses import dataclass
import hmac
from typing import Annotated
from fastapi import Depends, HTTPException, Request, Security, status
from fastapi.security import APIKeyHeader
api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False)
@dataclass(frozen=True, slots=True)
class AuthenticatedUser:
user_id: str
class ApiKeyAuthenticator:
def __init__(self, whitelist: dict[str, str]) -> None:
if not whitelist:
raise ValueError("API_KEY_WHITELIST must contain at least one user")
self._whitelist = whitelist.copy()
def authenticate(self, api_key: str | None) -> AuthenticatedUser | None:
if not api_key:
return None
for user_id, allowed_key in self._whitelist.items():
if hmac.compare_digest(api_key, allowed_key):
return AuthenticatedUser(user_id=user_id)
return None
async def require_api_key(
request: Request,
api_key: Annotated[str | None, Security(api_key_header)],
) -> AuthenticatedUser:
authenticator: ApiKeyAuthenticator = request.app.state.authenticator
user = authenticator.authenticate(api_key)
if user is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="api_key is required",
)
return user
CurrentUser = Annotated[AuthenticatedUser, Depends(require_api_key)]
+31
View File
@@ -1,6 +1,8 @@
from dataclasses import dataclass from dataclasses import dataclass
import json
from pathlib import Path from pathlib import Path
import os import os
from uuid import UUID
from dotenv import load_dotenv from dotenv import load_dotenv
@@ -13,6 +15,7 @@ class Settings:
model: str model: str
default_system_prompt: str default_system_prompt: str
data_dir: Path data_dir: Path
api_key_whitelist: dict[str, str]
max_tool_rounds: int = 5 max_tool_rounds: int = 5
max_tool_calls_per_turn: int = 10 max_tool_calls_per_turn: int = 10
tool_timeout_seconds: float = 5.0 tool_timeout_seconds: float = 5.0
@@ -29,9 +32,37 @@ class Settings:
"DEFAULT_SYSTEM_PROMPT", "You are a helpful assistant." "DEFAULT_SYSTEM_PROMPT", "You are a helpful assistant."
), ),
data_dir=Path(os.getenv("CHAT_DATA_DIR", "data")), data_dir=Path(os.getenv("CHAT_DATA_DIR", "data")),
api_key_whitelist=_parse_api_key_whitelist(
os.getenv("API_KEY_WHITELIST", "{}")
),
max_tool_rounds=int(os.getenv("MAX_TOOL_ROUNDS", "5")), max_tool_rounds=int(os.getenv("MAX_TOOL_ROUNDS", "5")),
max_tool_calls_per_turn=int( max_tool_calls_per_turn=int(
os.getenv("MAX_TOOL_CALLS_PER_TURN", "10") os.getenv("MAX_TOOL_CALLS_PER_TURN", "10")
), ),
tool_timeout_seconds=float(os.getenv("TOOL_TIMEOUT_SECONDS", "5")), tool_timeout_seconds=float(os.getenv("TOOL_TIMEOUT_SECONDS", "5")),
) )
def _parse_api_key_whitelist(value: str) -> dict[str, str]:
try:
parsed = json.loads(value)
except json.JSONDecodeError as exc:
raise ValueError("API_KEY_WHITELIST must be valid JSON") from exc
if not isinstance(parsed, dict):
raise ValueError("API_KEY_WHITELIST must map user UUIDs to API keys")
normalized: dict[str, str] = {}
for user_id, api_key in parsed.items():
if not isinstance(user_id, str) or not isinstance(api_key, str) or not api_key:
raise ValueError("API_KEY_WHITELIST must map user UUIDs to API keys")
try:
normalized_user_id = str(UUID(user_id))
except ValueError as exc:
raise ValueError("API_KEY_WHITELIST user IDs must be valid UUIDs") from exc
if normalized_user_id in normalized:
raise ValueError("API_KEY_WHITELIST contains duplicate user UUIDs")
normalized[normalized_user_id] = api_key
if len(set(normalized.values())) != len(normalized):
raise ValueError("API keys in API_KEY_WHITELIST must be unique")
return normalized
+5
View File
@@ -46,10 +46,15 @@ Message = Annotated[
class Session(BaseModel): class Session(BaseModel):
model_config = ConfigDict(extra="forbid") model_config = ConfigDict(extra="forbid")
user_id: str | None = None
session_id: str session_id: str
system_prompt: str system_prompt: str
created_at: datetime created_at: datetime
updated_at: datetime updated_at: datetime
api_calls: int = 0
prompt_tokens: int = 0
completion_tokens: int = 0
total_tokens: int = 0
messages: list[Message] = Field(default_factory=list) messages: list[Message] = Field(default_factory=list)
@model_validator(mode="before") @model_validator(mode="before")
+34 -8
View File
@@ -1,17 +1,19 @@
from fastapi import APIRouter, HTTPException, Request, status from fastapi import APIRouter, HTTPException, Request, status
from auth import CurrentUser
from schemas import ( from schemas import (
CreateSessionRequest, CreateSessionRequest,
CreateSessionResponse, CreateSessionResponse,
SendMessageRequest, SendMessageRequest,
SendMessageResponse, SendMessageResponse,
SessionHistoryResponse, SessionHistoryResponse,
UserUsageResponse,
) )
from service import ChatProviderError, ChatService from service import ChatProviderError, ChatService
from storage import SessionNotFoundError, SessionStorageError from storage import SessionNotFoundError, SessionStorageError
router = APIRouter(prefix="/sessions", tags=["sessions"]) router = APIRouter()
def get_chat_service(request: Request) -> ChatService: def get_chat_service(request: Request) -> ChatService:
@@ -19,32 +21,41 @@ def get_chat_service(request: Request) -> ChatService:
@router.post( @router.post(
"", "/sessions",
response_model=CreateSessionResponse, response_model=CreateSessionResponse,
status_code=status.HTTP_201_CREATED, status_code=status.HTTP_201_CREATED,
tags=["sessions"],
) )
async def create_session( async def create_session(
request: Request, request: Request,
_user: CurrentUser,
body: CreateSessionRequest | None = None, body: CreateSessionRequest | None = None,
) -> CreateSessionResponse: ) -> CreateSessionResponse:
try: try:
return await get_chat_service(request).create_session( return await get_chat_service(request).create_session(
body.system_prompt if body is not None else None _user.user_id,
body.system_prompt if body is not None else None,
) )
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
@router.post("/{session_id}/messages", response_model=SendMessageResponse) @router.post(
"/sessions/{session_id}/messages",
response_model=SendMessageResponse,
tags=["sessions"],
)
async def send_message( async def send_message(
session_id: str, session_id: str,
body: SendMessageRequest, body: SendMessageRequest,
request: Request, request: Request,
user: CurrentUser,
) -> SendMessageResponse: ) -> SendMessageResponse:
try: try:
return await get_chat_service(request).generate_response( response = await get_chat_service(request).generate_response(
session_id, body.content session_id, body.content, user.user_id
) )
return response
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:
@@ -53,14 +64,29 @@ async def send_message(
raise HTTPException(status_code=500, detail=str(exc)) from exc raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.get("/{session_id}/messages", response_model=SessionHistoryResponse) @router.get(
"/sessions/{session_id}/messages",
response_model=SessionHistoryResponse,
tags=["sessions"],
)
async def get_session_history( async def get_session_history(
session_id: str, session_id: str,
request: Request, request: Request,
_user: CurrentUser,
) -> SessionHistoryResponse: ) -> SessionHistoryResponse:
try: try:
return await get_chat_service(request).get_session_history(session_id) return await get_chat_service(request).get_session_history(
session_id, _user.user_id
)
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 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
@router.get("/usage", response_model=UserUsageResponse, tags=["usage"])
async def get_current_user_usage(
request: Request,
user: CurrentUser,
) -> UserUsageResponse:
return await get_chat_service(request).get_user_usage(user.user_id)
+9
View File
@@ -63,3 +63,12 @@ class SessionHistoryResponse(BaseModel):
created_at: datetime created_at: datetime
updated_at: datetime updated_at: datetime
messages: list[Message] messages: list[Message]
class UserUsageResponse(BaseModel):
user_id: str
api_calls: int
prompt_tokens: int
completion_tokens: int
total_tokens: int
updated_at: datetime
+46 -7
View File
@@ -19,8 +19,9 @@ from schemas import (
SendMessageResponse, SendMessageResponse,
SessionHistoryResponse, SessionHistoryResponse,
TokenUsage, TokenUsage,
UserUsageResponse,
) )
from storage import JsonSessionStorage from storage import JsonSessionStorage, SessionNotFoundError
from tools import ToolRegistry from tools import ToolRegistry
@@ -51,10 +52,11 @@ class ChatService:
self._locks: dict[str, asyncio.Lock] = {} self._locks: dict[str, asyncio.Lock] = {}
async def create_session( async def create_session(
self, system_prompt: str | None = None self, user_id: str, system_prompt: str | None = None
) -> CreateSessionResponse: ) -> CreateSessionResponse:
session = await self.storage.create( session = await self.storage.create(
system_prompt or self.default_system_prompt system_prompt or self.default_system_prompt,
user_id,
) )
return CreateSessionResponse( return CreateSessionResponse(
session_id=session.session_id, session_id=session.session_id,
@@ -63,12 +65,12 @@ class ChatService:
) )
async def generate_response( async def generate_response(
self, session_id: str, content: str self, session_id: str, content: str, user_id: str | None = None
) -> SendMessageResponse: ) -> SendMessageResponse:
lock = self._locks.setdefault(session_id, asyncio.Lock()) lock = self._locks.setdefault(session_id, asyncio.Lock())
async with lock: async with lock:
user_message = UserMessage(content=content, created_at=datetime.now(UTC)) user_message = UserMessage(content=content, created_at=datetime.now(UTC))
session = await self.storage.read(session_id) session = await self._read_owned_session(session_id, user_id)
api_messages = self._build_api_messages(session) api_messages = self._build_api_messages(session)
api_messages.append(self._message_to_api(user_message)) api_messages.append(self._message_to_api(user_message))
pending_messages: list[Message] = [user_message] pending_messages: list[Message] = [user_message]
@@ -98,6 +100,11 @@ class ChatService:
pending_messages.append(final_message) pending_messages.append(final_message)
session.messages.extend(pending_messages) session.messages.extend(pending_messages)
session.updated_at = final_message.created_at session.updated_at = final_message.created_at
session.api_calls += 1
if has_usage:
session.prompt_tokens += usage.prompt_tokens
session.completion_tokens += usage.completion_tokens
session.total_tokens += usage.total_tokens
await self.storage.write(session) await self.storage.write(session)
return SendMessageResponse( return SendMessageResponse(
session_id=session_id, session_id=session_id,
@@ -150,10 +157,12 @@ class ChatService:
raise ChatProviderError("Tool round limit exceeded") raise ChatProviderError("Tool round limit exceeded")
async def get_session_history(self, session_id: str) -> SessionHistoryResponse: async def get_session_history(
self, session_id: str, user_id: str | None = None
) -> SessionHistoryResponse:
lock = self._locks.setdefault(session_id, asyncio.Lock()) lock = self._locks.setdefault(session_id, asyncio.Lock())
async with lock: async with lock:
session = await self.storage.read(session_id) session = await self._read_owned_session(session_id, user_id)
return SessionHistoryResponse( return SessionHistoryResponse(
session_id=session.session_id, session_id=session.session_id,
system_prompt=session.system_prompt, system_prompt=session.system_prompt,
@@ -162,6 +171,36 @@ class ChatService:
messages=session.messages, messages=session.messages,
) )
async def get_user_usage(self, user_id: str) -> UserUsageResponse:
sessions = await self.storage.list_by_user(user_id)
updated_at = max(
(session.updated_at for session in sessions),
default=datetime.now(UTC),
)
return UserUsageResponse(
user_id=user_id,
api_calls=sum(session.api_calls for session in sessions),
prompt_tokens=sum(session.prompt_tokens for session in sessions),
completion_tokens=sum(
session.completion_tokens for session in sessions
),
total_tokens=sum(session.total_tokens for session in sessions),
updated_at=updated_at,
)
async def _read_owned_session(
self, session_id: str, user_id: str | None
) -> Session:
session = await self.storage.read(session_id)
if user_id is None:
return session
if session.user_id is None:
session.user_id = user_id
await self.storage.write(session)
elif session.user_id != user_id:
raise SessionNotFoundError(session_id)
return session
async def _request_completion( async def _request_completion(
self, messages: list[dict[str, Any]] self, messages: list[dict[str, Any]]
) -> Any: ) -> Any:
+23 -3
View File
@@ -23,8 +23,8 @@ class JsonSessionStorage:
def __init__(self, data_dir: Path) -> None: def __init__(self, data_dir: Path) -> None:
self.data_dir = data_dir self.data_dir = data_dir
async def create(self, system_prompt: str) -> Session: async def create(self, system_prompt: str, user_id: str | None = None) -> Session:
return await asyncio.to_thread(self._create, system_prompt) return await asyncio.to_thread(self._create, system_prompt, user_id)
async def read(self, session_id: str) -> Session: async def read(self, session_id: str) -> Session:
return await asyncio.to_thread(self._read, session_id) return await asyncio.to_thread(self._read, session_id)
@@ -32,6 +32,9 @@ class JsonSessionStorage:
async def write(self, session: Session) -> None: async def write(self, session: Session) -> None:
await asyncio.to_thread(self._write, session) await asyncio.to_thread(self._write, session)
async def list_by_user(self, user_id: str) -> list[Session]:
return await asyncio.to_thread(self._list_by_user, user_id)
def _path(self, session_id: str) -> Path: def _path(self, session_id: str) -> Path:
try: try:
normalized = str(UUID(session_id)) normalized = str(UUID(session_id))
@@ -39,12 +42,13 @@ class JsonSessionStorage:
raise SessionNotFoundError(session_id) from exc raise SessionNotFoundError(session_id) from exc
return self.data_dir / f"{normalized}.json" return self.data_dir / f"{normalized}.json"
def _create(self, system_prompt: str) -> Session: def _create(self, system_prompt: str, user_id: str | None) -> Session:
try: try:
self.data_dir.mkdir(parents=True, exist_ok=True) self.data_dir.mkdir(parents=True, exist_ok=True)
while True: while True:
now = datetime.now(UTC) now = datetime.now(UTC)
session = Session( session = Session(
user_id=user_id,
session_id=str(uuid4()), session_id=str(uuid4()),
system_prompt=system_prompt, system_prompt=system_prompt,
created_at=now, created_at=now,
@@ -75,6 +79,22 @@ class JsonSessionStorage:
except (OSError, json.JSONDecodeError, ValidationError) as exc: except (OSError, json.JSONDecodeError, ValidationError) as exc:
raise SessionStorageError("could not read session") from exc raise SessionStorageError("could not read session") from exc
def _list_by_user(self, user_id: str) -> list[Session]:
sessions: list[Session] = []
try:
for path in self.data_dir.glob("*.json"):
try:
UUID(path.stem)
except ValueError:
continue
with path.open(encoding="utf-8") as file:
session = Session.model_validate(json.load(file))
if session.user_id == user_id:
sessions.append(session)
except (OSError, json.JSONDecodeError, ValidationError) as exc:
raise SessionStorageError("could not list sessions") from exc
return sessions
def _write(self, session: Session) -> None: def _write(self, session: Session) -> None:
path = self._path(session.session_id) path = self._path(session.session_id)
if not path.exists(): if not path.exists():
+81 -1
View File
@@ -10,11 +10,15 @@ from openai import APIConnectionError
import pytest import pytest
from app import create_app from app import create_app
from config import Settings from config import Settings, _parse_api_key_whitelist
from service import ChatService from service import ChatService
from storage import JsonSessionStorage from storage import JsonSessionStorage
TEST_USER_ID = "11111111-1111-4111-8111-111111111111"
OTHER_USER_ID = "22222222-2222-4222-8222-222222222222"
class FakeCompletions: class FakeCompletions:
def __init__(self) -> None: def __init__(self) -> None:
self.calls: list[dict[str, Any]] = [] self.calls: list[dict[str, Any]] = []
@@ -55,9 +59,14 @@ def client_and_provider(tmp_path: Path) -> tuple[TestClient, FakeClient]:
model="deepseek-v4-flash", model="deepseek-v4-flash",
default_system_prompt="default prompt", default_system_prompt="default prompt",
data_dir=tmp_path, data_dir=tmp_path,
api_key_whitelist={
TEST_USER_ID: "test-api-key",
OTHER_USER_ID: "other-api-key",
},
) )
app = create_app(settings=settings, client=provider) # type: ignore[arg-type] app = create_app(settings=settings, client=provider) # type: ignore[arg-type]
with TestClient(app) as client: with TestClient(app) as client:
client.headers["X-API-Key"] = "test-api-key"
yield client, provider yield client, provider
@@ -79,6 +88,54 @@ def test_settings_load_api_key_from_dotenv(
assert settings.api_key == "test-key" assert settings.api_key == "test-key"
def test_api_key_whitelist_requires_and_normalizes_uuid() -> None:
parsed = _parse_api_key_whitelist(
'{"7A007F92-86D0-4F51-931F-007B9F0B8B0E":"key"}'
)
assert parsed == {
"7a007f92-86d0-4f51-931f-007b9f0b8b0e": "key"
}
with pytest.raises(ValueError, match="valid UUIDs"):
_parse_api_key_whitelist('{"custom-name":"key"}')
def test_all_api_endpoints_require_valid_api_key(
client_and_provider: tuple[TestClient, FakeClient]
) -> None:
client, _ = client_and_provider
del client.headers["X-API-Key"]
missing = client.post("/sessions")
invalid = client.get("/usage", headers={"X-API-Key": "invalid"})
assert missing.status_code == 401
assert missing.json() == {"detail": "api_key is required"}
assert invalid.status_code == 401
assert invalid.json() == {"detail": "api_key is required"}
def test_usage_endpoint_returns_current_user_totals(
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
) -> None:
client, _ = client_and_provider
session_id = create_session(client)
message = client.post(
f"/sessions/{session_id}/messages", json={"content": "hello"}
)
response = client.get("/usage")
assert message.status_code == 200
assert response.status_code == 200
assert response.json()["user_id"] == TEST_USER_ID
assert response.json()["api_calls"] == 1
assert response.json()["prompt_tokens"] == 10
assert response.json()["completion_tokens"] == 2
assert response.json()["total_tokens"] == 12
assert not (tmp_path / "usage.json").exists()
def test_create_session_uses_default_prompt( def test_create_session_uses_default_prompt(
client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path client_and_provider: tuple[TestClient, FakeClient], tmp_path: Path
) -> None: ) -> None:
@@ -86,6 +143,7 @@ def test_create_session_uses_default_prompt(
session_id = create_session(client) session_id = create_session(client)
data = json.loads((tmp_path / f"{session_id}.json").read_text()) data = json.loads((tmp_path / f"{session_id}.json").read_text())
assert data["user_id"] == TEST_USER_ID
assert data["system_prompt"] == "default prompt" assert data["system_prompt"] == "default prompt"
assert data["messages"] == [] assert data["messages"] == []
@@ -204,6 +262,7 @@ def test_legacy_messages_without_timestamp_remain_usable(
session_id = create_session(client) session_id = create_session(client)
path = tmp_path / f"{session_id}.json" path = tmp_path / f"{session_id}.json"
data = json.loads(path.read_text()) data = json.loads(path.read_text())
data.pop("user_id")
data["messages"] = [{"role": "user", "content": "legacy"}] data["messages"] = [{"role": "user", "content": "legacy"}]
path.write_text(json.dumps(data)) path.write_text(json.dumps(data))
@@ -213,9 +272,30 @@ def test_legacy_messages_without_timestamp_remain_usable(
assert response.status_code == 200 assert response.status_code == 200
migrated = json.loads(path.read_text()) migrated = json.loads(path.read_text())
assert migrated["user_id"] == TEST_USER_ID
assert all("created_at" in message for message in migrated["messages"]) assert all("created_at" in message for message in migrated["messages"])
def test_session_cannot_be_accessed_by_another_user(
client_and_provider: tuple[TestClient, FakeClient]
) -> None:
client, _ = client_and_provider
session_id = create_session(client)
other_headers = {"X-API-Key": "other-api-key"}
history = client.get(
f"/sessions/{session_id}/messages", headers=other_headers
)
message = client.post(
f"/sessions/{session_id}/messages",
headers=other_headers,
json={"content": "hello"},
)
assert history.status_code == 404
assert message.status_code == 404
def test_stream_parameter_is_rejected( def test_stream_parameter_is_rejected(
client_and_provider: tuple[TestClient, FakeClient] client_and_provider: tuple[TestClient, FakeClient]
) -> None: ) -> None: