feat: 新增多格式文件上传入库与认证体系
- 新增 JWT 认证模块,支持登录/注册/用户管理 - 新增文件上传接口,支持 .txt/.md/.html/.pdf/.docx 等格式解析入库 - 新增检索结果 AI 总结功能 - 新增文本去重缓存机制 - 新增全局认证夹具简化测试 - 新增配置项与环境变量支持 - 完善文档与测试覆盖
This commit is contained in:
@@ -0,0 +1,62 @@
|
||||
"""认证 API:POST /auth/login、POST /auth/register、GET /auth/me"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import structlog
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from app.api.response import ApiError, ok
|
||||
from app.config import settings
|
||||
from app.core.auth import (
|
||||
ERR_REGISTER_DISABLED,
|
||||
UserStore,
|
||||
create_access_token,
|
||||
get_current_user,
|
||||
get_user_store,
|
||||
)
|
||||
from app.models.auth import (
|
||||
AuthUser,
|
||||
LoginRequest,
|
||||
RegisterRequest,
|
||||
TokenResponse,
|
||||
)
|
||||
|
||||
logger = structlog.get_logger()
|
||||
|
||||
router = APIRouter(prefix="/api/v1", tags=["auth"])
|
||||
|
||||
|
||||
def _build_token_response(user, store: UserStore) -> dict[str, Any]:
|
||||
"""签发 token 并构造统一响应 data"""
|
||||
token, expires_in = create_access_token(user.username, user.role)
|
||||
auth_user = AuthUser(username=user.username, role=user.role, created_at=user.created_at)
|
||||
resp = TokenResponse(access_token=token, expires_in=expires_in, user=auth_user)
|
||||
return resp.model_dump(mode="json")
|
||||
|
||||
|
||||
@router.post("/auth/login")
|
||||
async def login(req: LoginRequest) -> dict[str, Any]:
|
||||
"""用户名密码登录,返回 JWT access_token"""
|
||||
store = get_user_store()
|
||||
user = await store.authenticate(req.username, req.password)
|
||||
logger.info("用户登录成功", username=user.username)
|
||||
return ok(_build_token_response(user, store))
|
||||
|
||||
|
||||
@router.post("/auth/register")
|
||||
async def register(req: RegisterRequest) -> dict[str, Any]:
|
||||
"""注册新用户(role=user),注册后自动签发 token
|
||||
|
||||
受 settings.auth_register_enabled 控制,关闭时返回 ERR_REGISTER_DISABLED。
|
||||
"""
|
||||
if not settings.auth_register_enabled:
|
||||
raise ApiError(ERR_REGISTER_DISABLED, "注册已关闭")
|
||||
store = get_user_store()
|
||||
user = await store.create(req.username, req.password, role="user")
|
||||
return ok(_build_token_response(user, store))
|
||||
|
||||
|
||||
@router.get("/auth/me")
|
||||
async def me(user: AuthUser = Depends(get_current_user)) -> dict[str, Any]:
|
||||
"""返回当前登录用户信息"""
|
||||
return ok(user.model_dump(mode="json"))
|
||||
+103
-7
@@ -1,13 +1,19 @@
|
||||
"""文档 API:POST /api/v1/documents 异步入库 + 任务查询 + 文档管理(列表/详情/删除)"""
|
||||
"""文档 API:POST /api/v1/documents 异步入库 + 任务查询 + 文档管理(列表/详情/删除)+ 文件上传入库"""
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import structlog
|
||||
from fastapi import APIRouter, Query
|
||||
from fastapi import APIRouter, Depends, File, Form, Query, UploadFile
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from app.api.response import ApiError, ok
|
||||
from app.config import Settings
|
||||
from app.config import Settings, settings
|
||||
from app.core.auth import AuthUser, get_current_user, require_admin
|
||||
from app.core.file_parser import parse_file
|
||||
from app.core.ingest_tasks import IngestTaskManager
|
||||
from app.core.ingestion import Ingester
|
||||
from app.models.document import DocumentInput
|
||||
@@ -55,8 +61,13 @@ def _get_task_manager() -> IngestTaskManager:
|
||||
return _task_manager
|
||||
|
||||
|
||||
def _allowed_extensions() -> set[str]:
|
||||
"""解析 settings.upload_allowed_extensions 逗号分隔字符串为扩展名集合(全小写、含点号)"""
|
||||
return {ext.strip().lower() for ext in settings.upload_allowed_extensions.split(",") if ext.strip()}
|
||||
|
||||
|
||||
@router.post("/documents")
|
||||
async def ingest_document(doc: DocumentInput) -> JSONResponse:
|
||||
async def ingest_document(doc: DocumentInput, user: AuthUser = Depends(get_current_user)) -> JSONResponse:
|
||||
"""文档入库入口:登记异步任务并返回 202 + task_id,入库结果经任务查询端点获取"""
|
||||
if not doc.text.strip():
|
||||
raise ApiError(1001, "文档内容不能为空")
|
||||
@@ -64,8 +75,92 @@ async def ingest_document(doc: DocumentInput) -> JSONResponse:
|
||||
return JSONResponse(status_code=202, content=ok({"task_id": task_id, "status": "pending"}))
|
||||
|
||||
|
||||
@router.post("/documents/upload")
|
||||
async def upload_document(
|
||||
file: UploadFile = File(..., description="上传的文件(.txt/.md/.html/.htm/.pdf/.docx)"),
|
||||
title: str = Form(default="", description="可选标题,默认取原文件名去扩展"),
|
||||
source: str = Form(default="", description="可选来源标识,默认 file:{原文件名}"),
|
||||
metadata: str = Form(default="", description='可选元数据 JSON 字符串,如 \'{"author":"x"}\''),
|
||||
user: AuthUser = Depends(get_current_user),
|
||||
) -> JSONResponse:
|
||||
"""文件上传入库入口:校验 → 提取文本 → 落盘 → 提交异步入库流水线
|
||||
|
||||
与 POST /documents 共用同一 IngestTaskManager;任务状态经
|
||||
GET /api/v1/documents/tasks/{task_id} 查询。
|
||||
"""
|
||||
original_filename = file.filename or "unnamed"
|
||||
ext = Path(original_filename).suffix.lower()
|
||||
|
||||
# 1. 扩展名校验
|
||||
allowed = _allowed_extensions()
|
||||
if ext not in allowed:
|
||||
raise ApiError(1001, f"不支持的文件类型: {ext or '(无扩展名)'}")
|
||||
|
||||
# 2. 读取字节并校验大小
|
||||
content = await file.read()
|
||||
max_bytes = settings.upload_max_size_mb * 1024 * 1024
|
||||
if len(content) > max_bytes:
|
||||
raise ApiError(1001, f"文件超过大小上限: {settings.upload_max_size_mb}MB")
|
||||
|
||||
# 3. 提取文本
|
||||
try:
|
||||
text = parse_file(original_filename, content)
|
||||
except ValueError as exc:
|
||||
raise ApiError(1001, str(exc)) from exc
|
||||
if not text.strip():
|
||||
raise ApiError(1001, "无法从文件提取文本")
|
||||
|
||||
# 4. 落盘(按 YYYY/MM 日期分片;失败仅 warning,不阻塞入库)
|
||||
doc_id = uuid.uuid4().hex
|
||||
saved_path = ""
|
||||
metadata_dict: dict[str, str] = {}
|
||||
try:
|
||||
upload_dir = Path(settings.upload_dir).resolve()
|
||||
shard_subdir = datetime.now(UTC).strftime("%Y/%m")
|
||||
target_dir = upload_dir / shard_subdir
|
||||
target_dir.mkdir(parents=True, exist_ok=True)
|
||||
target = target_dir / f"{doc_id}_{original_filename}"
|
||||
target.write_bytes(content)
|
||||
saved_path = str(target)
|
||||
metadata_dict.update(
|
||||
{
|
||||
"raw_file_path": saved_path,
|
||||
"original_filename": original_filename,
|
||||
"original_size_bytes": str(len(content)),
|
||||
}
|
||||
)
|
||||
logger.info("上传文件已落盘", doc_id=doc_id, saved_path=saved_path, size=len(content))
|
||||
except Exception:
|
||||
logger.warning("上传文件落盘失败,仅做文本入库", doc_id=doc_id, filename=original_filename, exc_info=True)
|
||||
|
||||
# 5. 合并用户传入的 metadata(落盘元数据优先级更高,不与用户键冲突)
|
||||
if metadata.strip():
|
||||
try:
|
||||
user_meta = json.loads(metadata)
|
||||
if isinstance(user_meta, dict):
|
||||
for k, v in user_meta.items():
|
||||
metadata_dict.setdefault(str(k), str(v))
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
logger.warning("metadata 不是合法 JSON,已忽略", raw=metadata)
|
||||
|
||||
# 6. 默认 title / source
|
||||
if not title.strip():
|
||||
title = Path(original_filename).stem
|
||||
if not source.strip():
|
||||
source = f"file:{original_filename}"
|
||||
|
||||
# 7. 提交入库流水线
|
||||
doc_input = DocumentInput(text=text, title=title, source=source, metadata=metadata_dict)
|
||||
task_id = await _get_task_manager().submit(doc_input)
|
||||
|
||||
return JSONResponse(
|
||||
status_code=202,
|
||||
content=ok({"task_id": task_id, "status": "pending", "saved_path": saved_path}),
|
||||
)
|
||||
|
||||
|
||||
@router.get("/documents/tasks/{task_id}")
|
||||
async def get_ingest_task(task_id: str) -> dict[str, Any]:
|
||||
async def get_ingest_task(task_id: str, user: AuthUser = Depends(get_current_user)) -> dict[str, Any]:
|
||||
"""查询入库任务状态:含 task_id/status/created_at/updated_at,done 附 result,failed 附 error"""
|
||||
task = await _get_task_manager().get(task_id)
|
||||
if task is None:
|
||||
@@ -77,6 +172,7 @@ async def get_ingest_task(task_id: str) -> dict[str, Any]:
|
||||
async def list_documents(
|
||||
limit: int = Query(default=20, ge=1, le=100),
|
||||
offset: str | None = None,
|
||||
user: AuthUser = Depends(get_current_user),
|
||||
) -> dict[str, Any]:
|
||||
"""分页列出文档(L1 摘要),返回 items 与下一页游标 next_offset"""
|
||||
try:
|
||||
@@ -88,7 +184,7 @@ async def list_documents(
|
||||
|
||||
|
||||
@router.get("/documents/{doc_id}")
|
||||
async def get_document(doc_id: str) -> dict[str, Any]:
|
||||
async def get_document(doc_id: str, user: AuthUser = Depends(get_current_user)) -> dict[str, Any]:
|
||||
"""获取文档详情:L1 记录 + L2/L3 节点 + chunks 数量"""
|
||||
try:
|
||||
detail = await _get_qdrant().get_doc_detail(doc_id)
|
||||
@@ -101,7 +197,7 @@ async def get_document(doc_id: str) -> dict[str, Any]:
|
||||
|
||||
|
||||
@router.delete("/documents/{doc_id}")
|
||||
async def delete_document(doc_id: str) -> dict[str, Any]:
|
||||
async def delete_document(doc_id: str, user: AuthUser = Depends(require_admin)) -> dict[str, Any]:
|
||||
"""删除文档:四层集合中该 doc_id 的所有点;幂等,不存在也返回成功(删除数全 0)"""
|
||||
try:
|
||||
deleted = await _get_qdrant().delete_by_doc_id(doc_id)
|
||||
|
||||
@@ -4,10 +4,11 @@ from functools import lru_cache
|
||||
from typing import Any
|
||||
|
||||
import structlog
|
||||
from fastapi import APIRouter
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from app.api.response import ApiError, ok
|
||||
from app.config import settings
|
||||
from app.core.auth import AuthUser, get_current_user
|
||||
from app.models.knowledge import UNCATEGORIZED, TaxonomyCategory, load_taxonomy
|
||||
from app.services.qdrant import ALL_COLLECTIONS, COLLECTION_L1, QdrantService
|
||||
|
||||
@@ -36,14 +37,14 @@ def _get_taxonomy() -> list[TaxonomyCategory]:
|
||||
|
||||
|
||||
@router.get("/knowledge/categories")
|
||||
async def list_categories() -> dict[str, Any]:
|
||||
async def list_categories(user: AuthUser = Depends(get_current_user)) -> dict[str, Any]:
|
||||
"""返回完整知识分类类目集"""
|
||||
categories = _get_taxonomy()
|
||||
return ok({"categories": [c.model_dump() for c in categories], "count": len(categories)})
|
||||
|
||||
|
||||
@router.get("/knowledge/stats")
|
||||
async def knowledge_stats() -> dict[str, Any]:
|
||||
async def knowledge_stats(user: AuthUser = Depends(get_current_user)) -> dict[str, Any]:
|
||||
"""返回四层集合规模与 L1 类目分布统计"""
|
||||
service = _get_qdrant()
|
||||
try:
|
||||
|
||||
@@ -4,9 +4,10 @@ from hashlib import sha256
|
||||
from typing import Any
|
||||
|
||||
import structlog
|
||||
from fastapi import APIRouter
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from app.api.response import ApiError, ok
|
||||
from app.core.auth import AuthUser, get_current_user
|
||||
from app.core.retriever import Retriever
|
||||
from app.models.search import SearchRequest
|
||||
from app.services.redis import get_cache
|
||||
@@ -27,13 +28,13 @@ def _get_retriever() -> Retriever:
|
||||
|
||||
|
||||
def _cache_key(request: SearchRequest) -> str:
|
||||
"""检索缓存键:query + top_k 的短哈希(top_k 影响结果集,需参与键计算)"""
|
||||
digest = sha256((request.query + "|" + str(request.top_k)).encode()).hexdigest()[:16]
|
||||
"""检索缓存键:query + top_k + summarize 的短哈希(三者均影响结果集,需参与键计算)"""
|
||||
digest = sha256((request.query + "|" + str(request.top_k) + "|" + str(request.summarize)).encode()).hexdigest()[:16]
|
||||
return f"search:{digest}"
|
||||
|
||||
|
||||
@router.post("/search")
|
||||
async def search(request: SearchRequest) -> dict[str, Any]:
|
||||
async def search(request: SearchRequest, user: AuthUser = Depends(get_current_user)) -> dict[str, Any]:
|
||||
"""分层检索入口,返回统一包装的 SearchResponse
|
||||
|
||||
先查 Redis 缓存:命中直接返回缓存的响应;未命中走检索流程并回写缓存。
|
||||
|
||||
Reference in New Issue
Block a user