This commit is contained in:
amb
2026-09-01 14:49:07 +08:00
parent 4db2d1a411
commit 70ece67910
8 changed files with 495 additions and 9 deletions
+104
View File
@@ -0,0 +1,104 @@
"""文档分类 API 路由。"""
from fastapi import APIRouter, Depends
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
from app.api.deps import get_current_user, get_db
from app.core.errors import NotFoundError
from app.models.document_category import DocumentCategory
from app.models.user import User
from app.repositories.kb_repo import KnowledgeBaseRepository
router = APIRouter(prefix="/knowledge-bases/{kb_id}/categories", tags=["categories"])
class CategoryRequest(BaseModel):
name: str = Field(min_length=1, max_length=255)
sort_order: int = 0
class CategoryResponse(BaseModel):
id: str
name: str
sort_order: int
model_config = {"from_attributes": True}
@router.get("", response_model=list[CategoryResponse])
def list_categories(
kb_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> list[CategoryResponse]:
_check_kb_owner(kb_id, user, db)
from sqlalchemy import select
stmt = (
select(DocumentCategory)
.where(DocumentCategory.knowledge_base_id == kb_id)
.order_by(DocumentCategory.sort_order)
)
cats = list(db.scalars(stmt).all())
return [CategoryResponse.model_validate(c) for c in cats]
@router.post("", response_model=CategoryResponse, status_code=201)
def create_category(
kb_id: str,
body: CategoryRequest,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> CategoryResponse:
_check_kb_owner(kb_id, user, db)
cat = DocumentCategory(
knowledge_base_id=kb_id,
name=body.name,
sort_order=body.sort_order,
)
db.add(cat)
db.commit()
db.refresh(cat)
return CategoryResponse.model_validate(cat)
@router.put("/{cat_id}", response_model=CategoryResponse)
def update_category(
kb_id: str,
cat_id: str,
body: CategoryRequest,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> CategoryResponse:
_check_kb_owner(kb_id, user, db)
cat = db.get(DocumentCategory, cat_id)
if cat is None or cat.knowledge_base_id != kb_id:
raise NotFoundError("分类不存在。")
cat.name = body.name
cat.sort_order = body.sort_order
db.commit()
db.refresh(cat)
return CategoryResponse.model_validate(cat)
@router.delete("/{cat_id}", status_code=204)
def delete_category(
kb_id: str,
cat_id: str,
user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> None:
_check_kb_owner(kb_id, user, db)
cat = db.get(DocumentCategory, cat_id)
if cat is None or cat.knowledge_base_id != kb_id:
raise NotFoundError("分类不存在。")
db.delete(cat)
db.commit()
def _check_kb_owner(kb_id: str, user: User, db: Session) -> None:
repo = KnowledgeBaseRepository(db)
kb = repo.get_by_id(kb_id)
if kb is None or kb.user_id != user.id or kb.status == "DELETED":
raise NotFoundError("知识库不存在。")
+2
View File
@@ -11,6 +11,7 @@ from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from app.api.auth import router as auth_router
from app.api.categories import router as cat_router
from app.api.documents import router as doc_router
from app.api.health import router as health_router
from app.api.knowledge_bases import router as kb_router
@@ -84,6 +85,7 @@ def create_app() -> FastAPI:
app.include_router(auth_router, prefix="/api", tags=["auth"])
app.include_router(kb_router, prefix="/api", tags=["knowledge-bases"])
app.include_router(doc_router, prefix="/api", tags=["documents"])
app.include_router(cat_router, prefix="/api", tags=["categories"])
app.include_router(public_router, tags=["public"])
return app
+24
View File
@@ -18,11 +18,27 @@ from app.api.deps import get_db
from app.core.errors import NotFoundError, RateLimitedError
from app.core.rate_limit import check_rate_limit
from app.core.security import decrypt_token
from app.services.access_log_service import AccessLogService
from app.services.kb_public_service import KbPublicService
router = APIRouter(prefix="/k", tags=["public"])
def _log_access(db: Session, kb_id: str, path: str, request: Request, doc_id: str | None = None, req_type: str | None = None) -> None:
"""记录访问日志(best-effort,不阻塞响应)。"""
try:
ua = request.headers.get("user-agent", "")
AccessLogService(db).record(
knowledge_base_id=kb_id,
document_id=doc_id,
path=path,
user_agent=ua,
request_type=req_type,
)
except Exception:
pass # 日志失败不影响响应
def _rate_limit(request: Request, token: str) -> None:
"""限流检查。"""
ip = request.client.host if request.client else "unknown"
@@ -51,6 +67,7 @@ def kb_index_markdown(
_rate_limit(request, token)
svc = KbPublicService(db)
kb = svc.get_kb_by_token(token)
_log_access(db, kb.id, f"/k/{token}.md", request, req_type="md")
docs, _ = svc.list_documents(kb)
lines = [f"# {kb.name}", ""]
@@ -87,6 +104,7 @@ def kb_index_text(
_rate_limit(request, token)
svc = KbPublicService(db)
kb = svc.get_kb_by_token(token)
_log_access(db, kb.id, f"/k/{token}.txt", request, req_type="txt")
docs, _ = svc.list_documents(kb)
lines = [kb.name, "=" * len(kb.name), ""]
@@ -119,6 +137,7 @@ def kb_index_json(
_rate_limit(request, token)
svc = KbPublicService(db)
kb = svc.get_kb_by_token(token)
_log_access(db, kb.id, f"/k/{token}.json", request, req_type="json")
docs, _ = svc.list_documents(kb)
doc_list = []
@@ -156,6 +175,7 @@ def kb_index_html(
_rate_limit(request, token)
svc = KbPublicService(db)
kb = svc.get_kb_by_token(token)
_log_access(db, kb.id, f"/k/{token}", request, req_type="html")
docs, total = svc.list_documents(kb, page=page, page_size=50)
doc_rows = ""
@@ -232,6 +252,7 @@ def doc_page_html(
svc = KbPublicService(db)
kb = svc.get_kb_by_token(token)
doc = svc.get_document_by_token(kb, doc_token)
_log_access(db, kb.id, f"/k/{token}/doc/{doc_token}", request, doc_id=doc.id, req_type="doc_html")
markdown_content = svc.get_document_markdown(doc)
# Markdown → HTML(简单转换)
@@ -291,6 +312,7 @@ def doc_page_markdown(
svc = KbPublicService(db)
kb = svc.get_kb_by_token(token)
doc = svc.get_document_by_token(kb, doc_token)
_log_access(db, kb.id, f"/k/{token}/doc/{doc_token}.md", request, doc_id=doc.id, req_type="doc_md")
markdown_content = svc.get_document_markdown(doc)
return PlainTextResponse(content=markdown_content, media_type="text/markdown")
@@ -307,6 +329,7 @@ def doc_page_text(
svc = KbPublicService(db)
kb = svc.get_kb_by_token(token)
doc = svc.get_document_by_token(kb, doc_token)
_log_access(db, kb.id, f"/k/{token}/doc/{doc_token}.txt", request, doc_id=doc.id, req_type="doc_txt")
markdown_content = svc.get_document_markdown(doc)
# 去掉 Markdown 标记
@@ -332,6 +355,7 @@ def search_html(
_rate_limit(request, token)
svc = KbPublicService(db)
kb = svc.get_kb_by_token(token)
_log_access(db, kb.id, f"/k/{token}/search?q={q}", request, req_type="search")
results, total = svc.search_documents(kb, q, page=page)
result_items = ""
@@ -0,0 +1,32 @@
"""访问日志服务。"""
from datetime import datetime, timezone
from app.models.access_log import AccessLog
from sqlalchemy.orm import Session
class AccessLogService:
def __init__(self, session: Session) -> None:
self._session = session
def record(
self,
*,
knowledge_base_id: str,
document_id: str | None = None,
path: str,
user_agent: str | None = None,
request_type: str | None = None,
) -> None:
"""记录一次访问。"""
log = AccessLog(
knowledge_base_id=knowledge_base_id,
document_id=document_id,
path=path,
accessed_at=datetime.now(timezone.utc).isoformat(),
user_agent=user_agent[:500] if user_agent else None,
request_type=request_type,
)
self._session.add(log)
self._session.commit()
+61
View File
@@ -0,0 +1,61 @@
"""Phase 15 访问日志测试。"""
from fastapi.testclient import TestClient
from app.main import app
def _client() -> TestClient:
return TestClient(app, raise_server_exceptions=False)
def test_access_log_recorded() -> None:
"""访问公共页面后应有访问日志。"""
with _client() as client:
# 注册+创建知识库
client.post("/api/auth/register", json={"username": "log_user", "email": "log@example.com", "password": "password123"})
resp = client.post("/api/knowledge-bases", json={"name": "Log Test"})
kb_id = resp.json()["id"]
# 获取 AI 链接
resp = client.get(f"/api/knowledge-bases/{kb_id}/link")
token = resp.json()["token"]
# 访问公共页面
client.get(f"/k/{token}")
client.get(f"/k/{token}.json")
# 检查访问日志
from app.core.db import get_session_factory
from sqlalchemy import text
factory = get_session_factory()
with factory() as session:
result = session.execute(text("SELECT COUNT(*) FROM access_logs WHERE knowledge_base_id = :kb_id"), {"kb_id": kb_id})
count = result.scalar()
assert count >= 2
def test_access_log_with_kb() -> None:
"""访问知识库首页应记录访问日志。"""
with _client() as client:
client.post("/api/auth/register", json={"username": "log_doc", "email": "log_doc@example.com", "password": "password123"})
resp = client.post("/api/knowledge-bases", json={"name": "Log Test"})
kb_id = resp.json()["id"]
# 获取 AI 链接
resp = client.get(f"/api/knowledge-bases/{kb_id}/link")
token = resp.json()["token"]
# 访问公共页面
client.get(f"/k/{token}")
# 检查访问日志
from app.core.db import get_session_factory
from sqlalchemy import text
factory = get_session_factory()
with factory() as session:
result = session.execute(text("SELECT COUNT(*) FROM access_logs WHERE knowledge_base_id = :kb_id"), {"kb_id": kb_id})
count = result.scalar()
assert count >= 1
+77
View File
@@ -0,0 +1,77 @@
"""Phase 15 分类测试。"""
from fastapi.testclient import TestClient
from app.main import app
def _client() -> TestClient:
return TestClient(app, raise_server_exceptions=False)
def _setup(client: TestClient) -> str:
"""注册+创建知识库,返回 kb_id。"""
client.post("/api/auth/register", json={
"username": "cat_user",
"email": "cat@example.com",
"password": "password123",
})
resp = client.post("/api/knowledge-bases", json={"name": "Cat Test KB"})
return resp.json()["id"]
def test_create_category() -> None:
with _client() as client:
kb_id = _setup(client)
resp = client.post(f"/api/knowledge-bases/{kb_id}/categories", json={
"name": "技术文档",
"sort_order": 1,
})
assert resp.status_code == 201
assert resp.json()["name"] == "技术文档"
assert resp.json()["sort_order"] == 1
def test_list_categories() -> None:
with _client() as client:
kb_id = _setup(client)
client.post(f"/api/knowledge-bases/{kb_id}/categories", json={"name": "A", "sort_order": 2})
client.post(f"/api/knowledge-bases/{kb_id}/categories", json={"name": "B", "sort_order": 1})
resp = client.get(f"/api/knowledge-bases/{kb_id}/categories")
assert resp.status_code == 200
cats = resp.json()
assert len(cats) == 2
# 按 sort_order 排序
assert cats[0]["name"] == "B"
assert cats[1]["name"] == "A"
def test_update_category() -> None:
with _client() as client:
kb_id = _setup(client)
resp = client.post(f"/api/knowledge-bases/{kb_id}/categories", json={"name": "Old"})
cat_id = resp.json()["id"]
resp = client.put(f"/api/knowledge-bases/{kb_id}/categories/{cat_id}", json={"name": "New"})
assert resp.status_code == 200
assert resp.json()["name"] == "New"
def test_delete_category() -> None:
with _client() as client:
kb_id = _setup(client)
resp = client.post(f"/api/knowledge-bases/{kb_id}/categories", json={"name": "To Delete"})
cat_id = resp.json()["id"]
resp = client.delete(f"/api/knowledge-bases/{kb_id}/categories/{cat_id}")
assert resp.status_code == 204
def test_category_idor() -> None:
"""用户A 不能操作用户B 的分类。"""
with _client() as client:
kb_id = _setup(client)
resp = client.post(f"/api/knowledge-bases/{kb_id}/categories", json={"name": "A's Cat"})
cat_id = resp.json()["id"]
# 用户 B 登录
client.post("/api/auth/register", json={"username": "cat_b", "email": "cat_b@example.com", "password": "password123"})
resp = client.get(f"/api/knowledge-bases/{kb_id}/categories")
assert resp.status_code == 404