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 缓存:命中直接返回缓存的响应;未命中走检索流程并回写缓存。
|
||||
|
||||
@@ -54,6 +54,27 @@ class Settings(BaseSettings):
|
||||
# 分块参数
|
||||
chunk_max_chars: int = 800 # chunk 超长二次切分阈值
|
||||
|
||||
# 认证(JWT)
|
||||
jwt_secret_key: str = "" # JWT 签名密钥,为空时启动自动生成(仅开发,生产必填)
|
||||
jwt_algorithm: str = "HS256"
|
||||
jwt_expire_minutes: int = 1440 # token 有效期(分钟),默认 24 小时
|
||||
auth_register_enabled: bool = True # 是否开放 POST /auth/register
|
||||
default_admin_username: str = "admin" # 启动时自动创建的默认管理员用户名
|
||||
default_admin_password: str = "" # 默认管理员密码,为空则不创建默认管理员
|
||||
|
||||
# 文件上传
|
||||
upload_dir: str = "./uploads" # 原始文件保存目录(相对路径以工作目录为基)
|
||||
upload_max_size_mb: int = 20 # 单文件大小上限(MB)
|
||||
upload_allowed_extensions: str = ".txt,.md,.html,.htm,.pdf,.docx" # 允许上传的扩展名(逗号分隔)
|
||||
|
||||
# PDF OCR(图片型/扫描件降级,pypdf extract_text 为空时触发)
|
||||
pdf_ocr_enabled: bool = True # 是否启用 OCR 降级(关闭则扫描件按"无法提取文本"拒绝入库)
|
||||
pdf_ocr_max_pages: int = 30 # 单文件 OCR 页数上限,超过仅前 N 页
|
||||
pdf_ocr_dpi: int = 200 # 渲染 DPI(越高越准但越慢,72~300 合理)
|
||||
|
||||
# 检索结果 AI 总结
|
||||
result_summary_max_hits: int = 5 # 参与总结的最大 hit 条数(控制 prompt 长度)
|
||||
|
||||
model_config = {"env_prefix": "", "case_sensitive": False}
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,213 @@
|
||||
"""JWT 认证核心:密码哈希、token 签发/验签、用户存储、FastAPI 鉴权依赖
|
||||
|
||||
用户持久化在 Redis(key: auth:user:{username})。登录/注册需 Redis 可用;
|
||||
鉴权(get_current_user)以 JWT 自包含信息为主,Redis 不可用时降级可用。
|
||||
"""
|
||||
|
||||
from datetime import UTC, datetime, timedelta
|
||||
|
||||
import bcrypt
|
||||
import jwt
|
||||
import structlog
|
||||
from fastapi import Depends
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
from redis import asyncio as redis_async
|
||||
|
||||
from app.api.response import ApiError
|
||||
from app.config import settings
|
||||
from app.models.auth import AuthUser, StoredUser
|
||||
|
||||
logger = structlog.get_logger()
|
||||
|
||||
# 错误码
|
||||
ERR_UNAUTHORIZED = 1003 # 未认证(需要登录)
|
||||
ERR_TOKEN_INVALID = 1005 # token 无效或已过期
|
||||
ERR_FORBIDDEN = 1006 # 权限不足
|
||||
ERR_USER_EXISTS = 1007 # 用户名已存在
|
||||
ERR_BAD_CREDENTIALS = 1008 # 用户名或密码错误
|
||||
ERR_REGISTER_DISABLED = 1009 # 注册已关闭
|
||||
|
||||
# Bearer token 提取器(auto_error=False,统一由 ApiError 处理)
|
||||
_bearer = HTTPBearer(auto_error=False)
|
||||
|
||||
# 模块级 JWT secret(settings 为空时自动生成,仅开发用)
|
||||
_auto_secret: str | None = None
|
||||
|
||||
# 模块级用户存储单例
|
||||
_user_store: "UserStore | None" = None
|
||||
|
||||
|
||||
def _get_secret() -> str:
|
||||
"""获取 JWT 签名密钥,settings 为空时生成一次性随机密钥(仅开发)"""
|
||||
global _auto_secret
|
||||
if settings.jwt_secret_key:
|
||||
return settings.jwt_secret_key
|
||||
if _auto_secret is None:
|
||||
import secrets
|
||||
|
||||
_auto_secret = secrets.token_urlsafe(48)
|
||||
logger.warning(
|
||||
"JWT_SECRET_KEY 未配置,已自动生成随机密钥(仅开发用,重启后旧 token 失效)"
|
||||
)
|
||||
return _auto_secret
|
||||
|
||||
|
||||
def hash_password(password: str) -> str:
|
||||
"""密码 bcrypt 哈希(返回 utf-8 字符串形式的哈希)"""
|
||||
return bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt()).decode("utf-8")
|
||||
|
||||
|
||||
def verify_password(password: str, hashed: str) -> bool:
|
||||
"""校验明文密码与哈希是否匹配,异常或不匹配均返回 False"""
|
||||
try:
|
||||
return bcrypt.checkpw(password.encode("utf-8"), hashed.encode("utf-8"))
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def create_access_token(username: str, role: str) -> tuple[str, int]:
|
||||
"""签发 JWT,返回 (token, expires_in_seconds)"""
|
||||
now = datetime.now(UTC)
|
||||
expire = now + timedelta(minutes=settings.jwt_expire_minutes)
|
||||
payload = {
|
||||
"sub": username,
|
||||
"role": role,
|
||||
"iat": int(now.timestamp()),
|
||||
"exp": int(expire.timestamp()),
|
||||
}
|
||||
token = jwt.encode(payload, _get_secret(), algorithm=settings.jwt_algorithm)
|
||||
return token, settings.jwt_expire_minutes * 60
|
||||
|
||||
|
||||
def decode_token(token: str) -> dict:
|
||||
"""验签并返回 payload,失败抛 ApiError(ERR_TOKEN_INVALID)"""
|
||||
try:
|
||||
return jwt.decode(token, _get_secret(), algorithms=[settings.jwt_algorithm])
|
||||
except jwt.ExpiredSignatureError as exc:
|
||||
raise ApiError(ERR_TOKEN_INVALID, "token 已过期") from exc
|
||||
except jwt.InvalidTokenError as exc:
|
||||
raise ApiError(ERR_TOKEN_INVALID, "token 无效") from exc
|
||||
|
||||
|
||||
class UserStore:
|
||||
"""Redis 用户存储"""
|
||||
|
||||
def __init__(self, redis_url: str | None = None) -> None:
|
||||
self._redis_url = redis_url or settings.redis_url
|
||||
self._client: redis_async.Redis | None = None
|
||||
|
||||
def _get_client(self) -> redis_async.Redis:
|
||||
if self._client is None:
|
||||
self._client = redis_async.from_url(self._redis_url, decode_responses=True)
|
||||
return self._client
|
||||
|
||||
@staticmethod
|
||||
def _key(username: str) -> str:
|
||||
return f"auth:user:{username}"
|
||||
|
||||
async def get(self, username: str) -> StoredUser | None:
|
||||
"""读取用户,Redis 异常或数据损坏均返回 None"""
|
||||
try:
|
||||
raw = await self._get_client().get(self._key(username))
|
||||
except Exception:
|
||||
logger.error("Redis 读取用户失败", username=username, exc_info=True)
|
||||
return None
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
return StoredUser.model_validate_json(raw)
|
||||
except Exception:
|
||||
logger.error("用户数据损坏", username=username, exc_info=True)
|
||||
return None
|
||||
|
||||
async def exists(self, username: str) -> bool:
|
||||
try:
|
||||
return bool(await self._get_client().exists(self._key(username)))
|
||||
except Exception:
|
||||
logger.error("Redis 检查用户存在性失败", username=username, exc_info=True)
|
||||
return False
|
||||
|
||||
async def create(
|
||||
self, username: str, password: str, role: str = "user"
|
||||
) -> StoredUser:
|
||||
"""创建用户,已存在抛 ApiError(ERR_USER_EXISTS),Redis 写失败抛 ApiError(2000)"""
|
||||
if await self.exists(username):
|
||||
raise ApiError(ERR_USER_EXISTS, f"用户名已存在: {username}")
|
||||
user = StoredUser(
|
||||
username=username,
|
||||
role=role,
|
||||
created_at=datetime.now(UTC),
|
||||
hashed_password=hash_password(password),
|
||||
)
|
||||
try:
|
||||
await self._get_client().set(self._key(username), user.model_dump_json())
|
||||
except Exception as exc:
|
||||
logger.error("Redis 写入用户失败", username=username, exc_info=True)
|
||||
raise ApiError(2000, f"用户创建失败: {exc}") from exc
|
||||
logger.info("用户创建成功", username=username, role=role)
|
||||
return user
|
||||
|
||||
async def authenticate(self, username: str, password: str) -> StoredUser:
|
||||
"""验证用户名密码,失败抛 ApiError(ERR_BAD_CREDENTIALS)"""
|
||||
user = await self.get(username)
|
||||
if user is None or not verify_password(password, user.hashed_password):
|
||||
raise ApiError(ERR_BAD_CREDENTIALS, "用户名或密码错误")
|
||||
return user
|
||||
|
||||
|
||||
def get_user_store() -> UserStore:
|
||||
"""全局 UserStore 单例"""
|
||||
global _user_store
|
||||
if _user_store is None:
|
||||
_user_store = UserStore()
|
||||
return _user_store
|
||||
|
||||
|
||||
async def ensure_default_admin() -> None:
|
||||
"""启动时按配置创建默认管理员(密码为空则跳过)"""
|
||||
pwd = settings.default_admin_password
|
||||
if not pwd:
|
||||
return
|
||||
store = get_user_store()
|
||||
try:
|
||||
if await store.exists(settings.default_admin_username):
|
||||
return
|
||||
await store.create(settings.default_admin_username, pwd, role="admin")
|
||||
logger.info("默认管理员已创建", username=settings.default_admin_username)
|
||||
except ApiError:
|
||||
raise
|
||||
except Exception:
|
||||
logger.warning("默认管理员创建失败,跳过", exc_info=True)
|
||||
|
||||
|
||||
async def get_current_user(
|
||||
credentials: HTTPAuthorizationCredentials | None = Depends(_bearer),
|
||||
) -> AuthUser:
|
||||
"""FastAPI 鉴权依赖:解析 Bearer token,返回 AuthUser
|
||||
|
||||
优先查 Redis 获取最新用户信息(支持删除/改角色生效);
|
||||
Redis 不可用或用户未命中时降级为 JWT payload 构造 AuthUser,保证鉴权可用。
|
||||
"""
|
||||
if credentials is None or (credentials.scheme or "").lower() != "bearer":
|
||||
raise ApiError(ERR_UNAUTHORIZED, "未提供认证凭证,请先登录")
|
||||
payload = decode_token(credentials.credentials)
|
||||
username = payload.get("sub")
|
||||
role = payload.get("role", "user")
|
||||
if not username:
|
||||
raise ApiError(ERR_TOKEN_INVALID, "token 缺少用户标识")
|
||||
|
||||
stored = await get_user_store().get(username)
|
||||
if stored is not None:
|
||||
return AuthUser(
|
||||
username=stored.username, role=stored.role, created_at=stored.created_at
|
||||
)
|
||||
# Redis 不可用或用户不存在(可能已删除):降级用 JWT payload
|
||||
logger.warning("用户存储查询未命中,降级使用 JWT payload", username=username)
|
||||
return AuthUser(username=username, role=role, created_at=datetime.now(UTC))
|
||||
|
||||
|
||||
async def require_admin(user: AuthUser = Depends(get_current_user)) -> AuthUser:
|
||||
"""要求管理员角色"""
|
||||
if user.role != "admin":
|
||||
raise ApiError(ERR_FORBIDDEN, "需要管理员权限")
|
||||
return user
|
||||
@@ -0,0 +1,204 @@
|
||||
"""多格式文件文本提取:按扩展名分发到对应解析器
|
||||
|
||||
支持的扩展名:
|
||||
- .txt / .md:UTF-8 解码(errors="replace" 兜底)
|
||||
- .html / .htm:标准库 html.parser 剥离标签提取可见文本
|
||||
- .pdf:pypdf 逐页 extract_text 拼接;文本层为空(扫描件/图片型)时
|
||||
自动降级为 OCR(pypdfium2 渲染 + rapidocr-onnxruntime 识别)
|
||||
- .docx:python-docx 段落文本拼接(不含表格/页眉页脚)
|
||||
|
||||
未识别扩展名抛 ValueError("不支持的文件类型: {ext}");
|
||||
解析异常统一包装为 ValueError("文件解析失败: {detail}"),原异常链式保留。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
from collections.abc import Callable
|
||||
from html.parser import HTMLParser
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import structlog
|
||||
|
||||
from app.config import settings
|
||||
|
||||
logger = structlog.get_logger()
|
||||
|
||||
# 模块级懒加载 OCR 引擎单例(首次调用时初始化,避免无扫描件场景白白下载模型)
|
||||
_ocr_engine: Any | None = None
|
||||
_ocr_unavailable: bool = False # 标记 OCR 依赖不可用,后续直接跳过避免重复尝试
|
||||
|
||||
|
||||
class _VisibleTextExtractor(HTMLParser):
|
||||
"""HTMLParser 子类:累积可见文本,跳过 script/style 内容"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__(convert_charrefs=True)
|
||||
self._parts: list[str] = []
|
||||
self._skip_depth = 0 # 在 script/style 标签内时 > 0
|
||||
|
||||
def handle_starttag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None:
|
||||
if tag.lower() in {"script", "style"}:
|
||||
self._skip_depth += 1
|
||||
|
||||
def handle_endtag(self, tag: str) -> None:
|
||||
if tag.lower() in {"script", "style"} and self._skip_depth > 0:
|
||||
self._skip_depth -= 1
|
||||
|
||||
def handle_data(self, data: str) -> None:
|
||||
if self._skip_depth == 0:
|
||||
self._parts.append(data)
|
||||
|
||||
def get_text(self) -> str:
|
||||
# 块级标签间用空格连接,再折叠多余空白
|
||||
text = " ".join(self._parts)
|
||||
return " ".join(text.split())
|
||||
|
||||
|
||||
def _decode_html(content: bytes) -> str:
|
||||
"""HTML 内容解码:优先 utf-8(带 BOM),失败回退 latin-1"""
|
||||
try:
|
||||
return content.decode("utf-8-sig")
|
||||
except UnicodeDecodeError:
|
||||
return content.decode("latin-1", errors="replace")
|
||||
|
||||
|
||||
def _parse_text(content: bytes) -> str:
|
||||
"""UTF-8 解码(errors=replace 兜底),保留原字符"""
|
||||
return content.decode("utf-8", errors="replace")
|
||||
|
||||
|
||||
def _parse_html(content: bytes) -> str:
|
||||
"""HTML 剥离标签,保留可见文本"""
|
||||
parser = _VisibleTextExtractor()
|
||||
parser.feed(_decode_html(content))
|
||||
parser.close()
|
||||
return parser.get_text()
|
||||
|
||||
|
||||
def _parse_pdf(content: bytes) -> str:
|
||||
"""PDF 解析:优先 pypdf extract_text;文本层为空(扫描件)时降级 OCR
|
||||
|
||||
OCR 流程:pypdfium2 渲染每页为 PIL Image → rapidocr-onnxruntime 识别 →
|
||||
拼接每页识别出的文本。受 settings.pdf_ocr_* 控制:开关、最大页数、DPI。
|
||||
OCR 依赖未安装或运行异常时降级返回空字符串(由上游 upload 端点拒绝入库)。
|
||||
"""
|
||||
from pypdf import PdfReader
|
||||
|
||||
reader = PdfReader(io.BytesIO(content))
|
||||
parts: list[str] = []
|
||||
for page in reader.pages:
|
||||
text = page.extract_text() or ""
|
||||
if text:
|
||||
parts.append(text)
|
||||
text_layer = "\n".join(parts).strip()
|
||||
|
||||
# 文本层非空:直接返回
|
||||
if text_layer:
|
||||
return text_layer
|
||||
|
||||
# 文本层为空 → 尝试 OCR 降级
|
||||
if not settings.pdf_ocr_enabled:
|
||||
return ""
|
||||
|
||||
ocr_text = _ocr_pdf(content, max_pages=settings.pdf_ocr_max_pages, dpi=settings.pdf_ocr_dpi)
|
||||
return ocr_text
|
||||
|
||||
|
||||
def _ocr_pdf(content: bytes, max_pages: int, dpi: int) -> str:
|
||||
"""对扫描件 PDF 跑 OCR:渲染每页 → 识别 → 拼接
|
||||
|
||||
返回空字符串的场景:依赖未安装 / 渲染或识别异常 / 无识别结果。
|
||||
任何异常仅告警不抛出,由上游按"无法提取文本"处理。
|
||||
"""
|
||||
global _ocr_engine, _ocr_unavailable
|
||||
|
||||
if _ocr_unavailable:
|
||||
return ""
|
||||
|
||||
# 1. 懒加载 OCR 引擎
|
||||
if _ocr_engine is None:
|
||||
try:
|
||||
from rapidocr_onnxruntime import RapidOCR
|
||||
|
||||
_ocr_engine = RapidOCR()
|
||||
logger.info("PDF OCR 引擎已初始化", dpi=dpi)
|
||||
except Exception:
|
||||
_ocr_unavailable = True
|
||||
logger.warning("OCR 依赖不可用,扫描件 PDF 将无法提取文本", exc_info=True)
|
||||
return ""
|
||||
|
||||
# 2. 渲染并识别
|
||||
try:
|
||||
import pypdfium2 as pdfium
|
||||
|
||||
scale = max(1.0, dpi / 72.0)
|
||||
pdf = pdfium.PdfDocument(io.BytesIO(content))
|
||||
total = min(len(pdf), max(1, max_pages))
|
||||
page_texts: list[str] = []
|
||||
for i in range(total):
|
||||
page = pdf[i]
|
||||
pil_image = page.render(scale=scale).to_pil()
|
||||
result, _ = _ocr_engine(pil_image)
|
||||
if result:
|
||||
# result: [[box, text, score], ...],按行拼接
|
||||
lines = [item[1] for item in result if item and len(item) >= 2 and item[1]]
|
||||
if lines:
|
||||
page_texts.append("\n".join(lines))
|
||||
pdf.close()
|
||||
return "\n".join(page_texts).strip()
|
||||
except Exception as exc:
|
||||
logger.warning("PDF OCR 失败,降级返回空文本", error=str(exc), exc_info=True)
|
||||
return ""
|
||||
|
||||
|
||||
def _parse_docx(content: bytes) -> str:
|
||||
"""DOCX 段落文本拼接(不含表格/页眉页脚)"""
|
||||
from docx import Document # type: ignore[import-untyped]
|
||||
|
||||
document = Document(io.BytesIO(content))
|
||||
parts = [p.text for p in document.paragraphs if p.text and p.text.strip()]
|
||||
return "\n".join(parts).strip()
|
||||
|
||||
|
||||
# 扩展名 → 解析函数映射(启动时构建,避免每次请求重复构造)
|
||||
_PARSERS: dict[str, Callable[[bytes], str]] = {
|
||||
".txt": _parse_text,
|
||||
".md": _parse_text,
|
||||
".html": _parse_html,
|
||||
".htm": _parse_html,
|
||||
".pdf": _parse_pdf,
|
||||
".docx": _parse_docx,
|
||||
}
|
||||
|
||||
|
||||
def parse_file(filename: str, content: bytes) -> str:
|
||||
"""按扩展名分发解析器提取文本
|
||||
|
||||
Args:
|
||||
filename: 文件名(用于判定扩展名)
|
||||
content: 文件二进制内容
|
||||
|
||||
Returns:
|
||||
提取的纯文本
|
||||
|
||||
Raises:
|
||||
ValueError: 未识别扩展名 → "不支持的文件类型: {ext}"
|
||||
解析失败 → "文件解析失败: {detail}"(保留原异常链)
|
||||
"""
|
||||
ext = Path(filename).suffix.lower()
|
||||
parser = _PARSERS.get(ext)
|
||||
if parser is None:
|
||||
raise ValueError(f"不支持的文件类型: {ext or '(无扩展名)'}")
|
||||
try:
|
||||
return parser(content)
|
||||
except ValueError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise ValueError(f"文件解析失败: {exc}") from exc
|
||||
|
||||
|
||||
def supported_extensions() -> set[str]:
|
||||
"""返回当前支持的扩展名集合(含点号,全小写)"""
|
||||
return set(_PARSERS.keys())
|
||||
@@ -9,6 +9,7 @@ Redis 不可用或写入失败仅记录 warning,不影响任务执行。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import uuid
|
||||
from datetime import UTC, datetime
|
||||
from enum import StrEnum
|
||||
@@ -25,6 +26,8 @@ logger = structlog.get_logger()
|
||||
|
||||
# Redis 任务状态 key 前缀
|
||||
REDIS_KEY_PREFIX = "ingest_task:"
|
||||
# Redis 文本去重 key 前缀(value 为已入库文档的 IngestionResult JSON)
|
||||
DEDUP_KEY_PREFIX = "dedup:sha256:"
|
||||
|
||||
|
||||
class IngestTaskStatus(StrEnum):
|
||||
@@ -62,9 +65,39 @@ class IngestTaskManager:
|
||||
self._mirror_tasks: set[asyncio.Task[None]] = set()
|
||||
|
||||
async def submit(self, doc: DocumentInput) -> str:
|
||||
"""登记入库任务并后台执行,立即返回 task_id"""
|
||||
"""登记入库任务并后台执行,立即返回 task_id
|
||||
|
||||
文本去重:基于 doc.text 的 sha256 在 Redis 中查重;命中则直接复用旧
|
||||
IngestionResult(仅置 deduplicated=True),不重跑流水线;未命中走原
|
||||
异步入库流程,完成后写入去重记录供后续命中复用。Redis 不可用时跳过
|
||||
去重,按原流程执行,不影响主流程。
|
||||
"""
|
||||
task_id = uuid.uuid4().hex
|
||||
now = _utc_now_iso()
|
||||
text_hash = hashlib.sha256(doc.text.encode("utf-8")).hexdigest()
|
||||
|
||||
# 1. 去重命中:直接置 done,复用旧结果,不调 _run
|
||||
dedup_record = await self._lookup_dedup(text_hash)
|
||||
if dedup_record is not None:
|
||||
result_dict = dict(dedup_record)
|
||||
result_dict["deduplicated"] = True
|
||||
self._tasks[task_id] = {
|
||||
"task_id": task_id,
|
||||
"status": IngestTaskStatus.DONE,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
"result": result_dict,
|
||||
"error": None,
|
||||
}
|
||||
self._schedule_mirror(task_id)
|
||||
logger.info(
|
||||
"入库任务命中去重,复用既有文档",
|
||||
task_id=task_id,
|
||||
document_id=result_dict.get("document_id"),
|
||||
)
|
||||
return task_id
|
||||
|
||||
# 2. 未命中:登记 pending 并后台跑流水线
|
||||
self._tasks[task_id] = {
|
||||
"task_id": task_id,
|
||||
"status": IngestTaskStatus.PENDING,
|
||||
@@ -74,7 +107,7 @@ class IngestTaskManager:
|
||||
"error": None,
|
||||
}
|
||||
self._schedule_mirror(task_id)
|
||||
background = asyncio.create_task(self._run(task_id, doc))
|
||||
background = asyncio.create_task(self._run(task_id, doc, text_hash))
|
||||
self._background_tasks.add(background)
|
||||
background.add_done_callback(self._background_tasks.discard)
|
||||
logger.info("入库任务已登记", task_id=task_id, title=doc.title)
|
||||
@@ -110,8 +143,8 @@ class IngestTaskManager:
|
||||
raise TimeoutError(f"入库任务 {task_id} 在 {timeout}s 内未进入终态")
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
async def _run(self, task_id: str, doc: DocumentInput) -> None:
|
||||
"""后台执行入库:并发限流 + 阶段状态推进 + 结果/错误落账"""
|
||||
async def _run(self, task_id: str, doc: DocumentInput, text_hash: str) -> None:
|
||||
"""后台执行入库:并发限流 + 阶段状态推进 + 结果/错误落账 + 去重记录写入"""
|
||||
async with self._semaphore:
|
||||
try:
|
||||
result = await self._ingester.ingest(doc, progress_cb=lambda stage: self._on_progress(task_id, stage))
|
||||
@@ -128,14 +161,39 @@ class IngestTaskManager:
|
||||
except Exception as exc:
|
||||
self._finish_failed(task_id, {"stage": "unknown", "message": str(exc), "partial_summary": None})
|
||||
else:
|
||||
result_dict = result.model_dump(mode="json")
|
||||
self._tasks[task_id].update(
|
||||
status=IngestTaskStatus.DONE,
|
||||
updated_at=_utc_now_iso(),
|
||||
result=result.model_dump(mode="json"),
|
||||
result=result_dict,
|
||||
)
|
||||
self._schedule_mirror(task_id)
|
||||
await self._record_dedup(text_hash, result_dict)
|
||||
logger.info("入库任务完成", task_id=task_id)
|
||||
|
||||
async def _lookup_dedup(self, text_hash: str) -> dict[str, Any] | None:
|
||||
"""查询文本去重记录;Redis 不可用或异常时降级为未命中"""
|
||||
if self._redis is None:
|
||||
return None
|
||||
try:
|
||||
return await self._redis.get_json(f"{DEDUP_KEY_PREFIX}{text_hash}")
|
||||
except Exception:
|
||||
logger.warning("去重记录查询失败,降级为未命中", text_hash=text_hash, exc_info=True)
|
||||
return None
|
||||
|
||||
async def _record_dedup(self, text_hash: str, result_dict: dict[str, Any]) -> None:
|
||||
"""写入文本去重记录(含完整 IngestionResult),供后续命中复用;失败仅告警"""
|
||||
if self._redis is None:
|
||||
return
|
||||
try:
|
||||
await self._redis.set_json(
|
||||
f"{DEDUP_KEY_PREFIX}{text_hash}",
|
||||
result_dict,
|
||||
ttl=self._settings.ingest_task_ttl_done,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("去重记录写入失败", text_hash=text_hash, exc_info=True)
|
||||
|
||||
def _finish_failed(self, task_id: str, error: dict[str, Any]) -> None:
|
||||
"""将任务置为 failed 并记录错误信息"""
|
||||
self._tasks[task_id].update(status=IngestTaskStatus.FAILED, updated_at=_utc_now_iso(), error=error)
|
||||
|
||||
@@ -40,6 +40,9 @@ class ParsedQuery(BaseModel):
|
||||
rewrite: str = Field(description="rewrite 后的 query")
|
||||
keywords: list[str] = Field(default_factory=list, description="提取的关键词")
|
||||
categories: list[CategoryHit] = Field(default_factory=list, description="命中类目列表")
|
||||
entities: list[str] = Field(default_factory=list, description="提取的实体(人名/产品/技术等)")
|
||||
intent: str = Field(default="", description="查询意图分类(如 查询/对比/操作/定义/排查)")
|
||||
time_range: str = Field(default="", description="时间范围(如 '2023年'、'最近一个月'),空表示无")
|
||||
parse_failed: bool = Field(default=False, description="LLM 输出解析是否失败")
|
||||
|
||||
|
||||
@@ -136,11 +139,28 @@ class QueryParser:
|
||||
logger.warning("query 解析失败:JSON 必填字段缺失或类型错误", query=query, output=raw[:200])
|
||||
return ParsedQuery(raw_query=query, rewrite=query, parse_failed=True)
|
||||
|
||||
# 扩展字段:实体/意图/时间范围(缺失或类型错误时降级为默认值,不触发 parse_failed)
|
||||
entities = data.get("entities", [])
|
||||
if not isinstance(entities, list):
|
||||
logger.warning("query 解析 entities 非列表,置空", query=query)
|
||||
entities = []
|
||||
intent = data.get("intent", "")
|
||||
if not isinstance(intent, str):
|
||||
logger.warning("query 解析 intent 非字符串,置空", query=query)
|
||||
intent = ""
|
||||
time_range = data.get("time_range", "")
|
||||
if not isinstance(time_range, str):
|
||||
logger.warning("query 解析 time_range 非字符串,置空", query=query)
|
||||
time_range = ""
|
||||
|
||||
return ParsedQuery(
|
||||
raw_query=query,
|
||||
rewrite=rewrite,
|
||||
keywords=[str(k) for k in keywords],
|
||||
categories=self._validate_categories(categories, query),
|
||||
entities=[str(e) for e in entities],
|
||||
intent=intent,
|
||||
time_range=time_range,
|
||||
)
|
||||
|
||||
async def parse_and_route(self, query: str) -> RouteDecision:
|
||||
@@ -178,17 +198,22 @@ class QueryParser:
|
||||
category_lines = [f"- {c.name}: {c.description}" for c in self.taxonomy if c.name != UNCATEGORIZED]
|
||||
category_block = "\n".join(category_lines)
|
||||
return (
|
||||
"你是搜索查询分析助手。请分析用户 query,完成三件事:\n"
|
||||
"你是搜索查询分析助手。请分析用户 query,完成以下事项:\n"
|
||||
"1. 判断 query 意图命中以下哪些知识类目,并给出每个类目的置信度(0~1 之间的小数);\n"
|
||||
"2. 将 query 改写为更适合检索的形式;\n"
|
||||
"3. 提取 query 的关键词。\n\n"
|
||||
"3. 提取 query 的关键词;\n"
|
||||
"4. 提取 query 中的实体(人名、产品名、技术名词、组织机构等);\n"
|
||||
"5. 判断 query 的查询意图,从 [查询, 对比, 操作, 定义, 排查] 中选最接近的一个;\n"
|
||||
"6. 提取 query 中的时间范围(如 '2023年'、'上个月'、'Q1'),无则留空字符串。\n\n"
|
||||
f"可选类目:\n{category_block}\n\n"
|
||||
"要求:\n"
|
||||
"- categories 中的 name 只能从上面的类目名中选择,不要输出其他名称;\n"
|
||||
"- 若没有明显命中的类目,categories 返回空列表;\n"
|
||||
"- entities 没有则返回空数组;\n"
|
||||
"- 只输出 JSON,不要输出任何其他内容。\n\n"
|
||||
'输出格式:{"categories": [{"name": "类目名", "confidence": 0.0}], '
|
||||
'"rewrite": "改写后的 query", "keywords": ["关键词"]}\n\n'
|
||||
'"rewrite": "改写后的 query", "keywords": ["关键词"], '
|
||||
'"entities": ["实体"], "intent": "查询", "time_range": ""}\n\n'
|
||||
f"用户 query:{query}"
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
"""检索结果 AI 总结
|
||||
|
||||
对检索返回的 chunk 命中结果,调用 Ollama 本地模型生成一段针对用户 query 的总结回答。
|
||||
仅基于检索结果内容,不编造未提及的信息。
|
||||
"""
|
||||
|
||||
import structlog
|
||||
|
||||
from app.config import settings
|
||||
from app.models.search import SearchHit
|
||||
from app.services.ollama import OllamaClient
|
||||
|
||||
logger = structlog.get_logger()
|
||||
|
||||
|
||||
class ResultSummarizer:
|
||||
"""检索结果总结器"""
|
||||
|
||||
def __init__(self, ollama: OllamaClient | None = None) -> None:
|
||||
self.ollama = ollama or OllamaClient()
|
||||
|
||||
async def summarize(self, query: str, hits: list[SearchHit]) -> str:
|
||||
"""对检索结果生成针对 query 的总结
|
||||
|
||||
取前 settings.result_summary_max_hits 条命中拼接为上下文,
|
||||
无命中时返回空字符串(不调 LLM)。
|
||||
"""
|
||||
if not hits:
|
||||
return ""
|
||||
max_hits = settings.result_summary_max_hits
|
||||
selected = hits[:max_hits]
|
||||
context = self._build_context(selected)
|
||||
prompt = self._build_prompt(query, context)
|
||||
try:
|
||||
summary = await self.ollama.generate(prompt)
|
||||
except Exception:
|
||||
logger.warning("检索结果总结生成失败,返回空字符串", exc_info=True)
|
||||
return ""
|
||||
logger.info("检索结果总结完成", query=query, hits_count=len(selected), summary_len=len(summary))
|
||||
return summary.strip()
|
||||
|
||||
@staticmethod
|
||||
def _build_context(hits: list[SearchHit]) -> str:
|
||||
"""拼接命中结果为带编号的上下文"""
|
||||
blocks: list[str] = []
|
||||
for i, hit in enumerate(hits, start=1):
|
||||
header_parts = [f"[{i}]"]
|
||||
if hit.title:
|
||||
header_parts.append(hit.title)
|
||||
if hit.section_path:
|
||||
header_parts.append(hit.section_path)
|
||||
header = " / ".join(header_parts)
|
||||
blocks.append(f"{header}\n{hit.text}")
|
||||
return "\n\n---\n\n".join(blocks)
|
||||
|
||||
@staticmethod
|
||||
def _build_prompt(query: str, context: str) -> str:
|
||||
return (
|
||||
"请根据以下检索结果,针对用户问题生成一段简洁的总结回答。\n"
|
||||
"要求:\n"
|
||||
"- 综合多条结果信息,不要简单逐条罗列;\n"
|
||||
"- 只基于检索结果内容,不编造未提及的信息;\n"
|
||||
"- 用中文回答,简洁明了。\n\n"
|
||||
f"用户问题:{query}\n\n"
|
||||
f"检索结果:\n{context}"
|
||||
)
|
||||
+39
-14
@@ -15,11 +15,12 @@ from qdrant_client import models
|
||||
|
||||
from app.config import settings
|
||||
from app.core.embeddings import EmbeddingService, create_embedding_service
|
||||
from app.core.query_parser import QueryParser
|
||||
from app.core.query_parser import QueryParser, RouteDecision
|
||||
from app.core.ranker import finalize, rrf_fuse
|
||||
from app.core.result_summarizer import ResultSummarizer
|
||||
from app.core.sparse import SparseEncoder
|
||||
from app.models.knowledge import load_taxonomy
|
||||
from app.models.search import SearchHit, SearchRequest, SearchResponse
|
||||
from app.models.search import ExtractedInfo, SearchHit, SearchRequest, SearchResponse
|
||||
from app.services.ollama import OllamaClient
|
||||
from app.services.qdrant import (
|
||||
COLLECTION_CHUNKS,
|
||||
@@ -47,6 +48,7 @@ class Retriever:
|
||||
query_parser: QueryParser | None = None,
|
||||
embedding: EmbeddingService | None = None,
|
||||
sparse_encoder: SparseEncoder | None = None,
|
||||
result_summarizer: ResultSummarizer | None = None,
|
||||
) -> None:
|
||||
self.qdrant = qdrant or QdrantService()
|
||||
self.query_parser = query_parser or QueryParser(
|
||||
@@ -55,6 +57,7 @@ class Retriever:
|
||||
)
|
||||
self.embedding = embedding or create_embedding_service()
|
||||
self.sparse_encoder = sparse_encoder or SparseEncoder()
|
||||
self.result_summarizer = result_summarizer or ResultSummarizer()
|
||||
|
||||
async def search(self, request: SearchRequest) -> SearchResponse:
|
||||
"""分层检索主流程"""
|
||||
@@ -74,12 +77,8 @@ class Retriever:
|
||||
# L1 无候选文档 → 全库 chunk 兜底
|
||||
chunk_hits = await self._search_collection(COLLECTION_CHUNKS, dense, sparse, settings.retrieval_top_k, None)
|
||||
logger.info("L1 无命中,全库 chunk 兜底", hits=len(chunk_hits))
|
||||
return SearchResponse(
|
||||
query=request.query,
|
||||
hits=self._to_hits(finalize(chunk_hits, self._final_k(request))),
|
||||
routed_categories=route.filter_categories or [],
|
||||
fallback=True,
|
||||
)
|
||||
hits = self._to_hits(finalize(chunk_hits, self._final_k(request)))
|
||||
return await self._build_response(request, route, hits, fallback=True)
|
||||
|
||||
doc_ids = _unique((p.payload or {}).get("doc_id") for p in l1_hits)
|
||||
|
||||
@@ -125,12 +124,8 @@ class Retriever:
|
||||
logger.info("chunk 检索完成", hits=len(chunk_hits))
|
||||
|
||||
final_points = finalize(rrf_fuse([chunk_hits]), self._final_k(request))
|
||||
return SearchResponse(
|
||||
query=request.query,
|
||||
hits=self._to_hits(final_points),
|
||||
routed_categories=route.filter_categories or [],
|
||||
fallback=route.fallback,
|
||||
)
|
||||
hits = self._to_hits(final_points)
|
||||
return await self._build_response(request, route, hits, fallback=route.fallback)
|
||||
|
||||
async def _search_collection(
|
||||
self,
|
||||
@@ -167,3 +162,33 @@ class Retriever:
|
||||
)
|
||||
)
|
||||
return hits
|
||||
|
||||
async def _build_response(
|
||||
self,
|
||||
request: SearchRequest,
|
||||
route: RouteDecision,
|
||||
hits: list[SearchHit],
|
||||
*,
|
||||
fallback: bool,
|
||||
) -> SearchResponse:
|
||||
"""构造最终响应:组装 AI 提取信息,按需生成结果总结"""
|
||||
parsed = route.parsed
|
||||
extracted = ExtractedInfo(
|
||||
rewrite=parsed.rewrite,
|
||||
keywords=parsed.keywords,
|
||||
entities=parsed.entities,
|
||||
intent=parsed.intent,
|
||||
time_range=parsed.time_range,
|
||||
categories=[c.name for c in parsed.categories],
|
||||
)
|
||||
summary: str | None = None
|
||||
if request.summarize:
|
||||
summary = await self.result_summarizer.summarize(request.query, hits)
|
||||
return SearchResponse(
|
||||
query=request.query,
|
||||
hits=hits,
|
||||
routed_categories=route.filter_categories or [],
|
||||
fallback=fallback,
|
||||
extracted_info=extracted,
|
||||
summary=summary,
|
||||
)
|
||||
|
||||
+14
-2
@@ -8,10 +8,12 @@ from fastapi.exceptions import RequestValidationError
|
||||
from fastapi.responses import FileResponse, JSONResponse
|
||||
|
||||
from app.api.response import ApiError, error
|
||||
from app.api.v1.auth import router as auth_router
|
||||
from app.api.v1.document import router as document_router
|
||||
from app.api.v1.knowledge import router as knowledge_router
|
||||
from app.api.v1.search import router as search_router
|
||||
from app.config import settings
|
||||
from app.core.auth import ensure_default_admin
|
||||
from app.services.qdrant import QdrantService
|
||||
|
||||
logger = structlog.get_logger()
|
||||
@@ -19,15 +21,24 @@ logger = structlog.get_logger()
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(_: FastAPI) -> AsyncIterator[None]:
|
||||
"""应用生命周期:启动时初始化 Qdrant 集合
|
||||
"""应用生命周期:启动时初始化 Qdrant 集合与默认管理员
|
||||
|
||||
初始化失败仅记录日志、不阻止启动(本地开发可能无 Qdrant)。
|
||||
初始化失败仅记录日志、不阻止启动(本地开发可能无 Qdrant/Redis)。
|
||||
"""
|
||||
try:
|
||||
await QdrantService().ensure_collections()
|
||||
logger.info("Qdrant 集合初始化完成")
|
||||
except Exception:
|
||||
logger.error("Qdrant 集合初始化失败,跳过初始化继续启动")
|
||||
try:
|
||||
await ensure_default_admin()
|
||||
except Exception:
|
||||
logger.warning("默认管理员初始化失败,跳过", exc_info=True)
|
||||
try:
|
||||
Path(settings.upload_dir).mkdir(parents=True, exist_ok=True)
|
||||
logger.info("上传目录已就绪", upload_dir=settings.upload_dir)
|
||||
except Exception:
|
||||
logger.warning("上传目录初始化失败,文件落盘将按需创建", exc_info=True)
|
||||
yield
|
||||
|
||||
|
||||
@@ -38,6 +49,7 @@ app = FastAPI(
|
||||
lifespan=lifespan,
|
||||
)
|
||||
|
||||
app.include_router(auth_router)
|
||||
app.include_router(search_router)
|
||||
app.include_router(document_router)
|
||||
app.include_router(knowledge_router)
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
"""认证相关数据模型"""
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class LoginRequest(BaseModel):
|
||||
"""登录请求"""
|
||||
|
||||
username: str = Field(description="用户名")
|
||||
password: str = Field(description="明文密码")
|
||||
|
||||
|
||||
class RegisterRequest(BaseModel):
|
||||
"""注册请求"""
|
||||
|
||||
username: str = Field(min_length=3, max_length=32, description="用户名(3-32 字符)")
|
||||
password: str = Field(min_length=6, max_length=128, description="密码(6-128 字符)")
|
||||
|
||||
|
||||
class AuthUser(BaseModel):
|
||||
"""对外暴露的用户信息(不含密码)"""
|
||||
|
||||
username: str = Field(description="用户名")
|
||||
role: str = Field(default="user", description="角色:admin | user")
|
||||
created_at: datetime = Field(description="创建时间")
|
||||
|
||||
|
||||
class StoredUser(AuthUser):
|
||||
"""存储层用户:含密码哈希,仅内部使用,不对外暴露"""
|
||||
|
||||
hashed_password: str = Field(description="bcrypt 密码哈希")
|
||||
|
||||
|
||||
class TokenResponse(BaseModel):
|
||||
"""登录成功返回的 token 信息"""
|
||||
|
||||
access_token: str = Field(description="JWT access token")
|
||||
token_type: str = Field(default="bearer", description="token 类型")
|
||||
expires_in: int = Field(description="token 有效期(秒)")
|
||||
user: AuthUser = Field(description="登录用户信息")
|
||||
@@ -49,3 +49,4 @@ class IngestionResult(BaseModel):
|
||||
chunks_count: int = Field(default=0, description="写入的 chunk 数量")
|
||||
tags: list[str] = Field(default_factory=list, description="附加分类标签")
|
||||
category_confidence: float = Field(default=0.0, description="主类目分类置信度")
|
||||
deduplicated: bool = Field(default=False, description="是否命中去重复用旧文档(True 时 document_id 为既有文档 ID)")
|
||||
|
||||
@@ -8,6 +8,7 @@ class SearchRequest(BaseModel):
|
||||
|
||||
query: str = Field(description="查询文本")
|
||||
top_k: int | None = Field(default=None, description="返回结果数,为空时使用 settings.retrieval_final_k")
|
||||
summarize: bool = Field(default=False, description="是否对检索结果生成 AI 总结")
|
||||
|
||||
|
||||
class SearchHit(BaseModel):
|
||||
@@ -21,6 +22,17 @@ class SearchHit(BaseModel):
|
||||
doc_summary: str = Field(default="", description="L1 文档总结,仅用于上下文标注")
|
||||
|
||||
|
||||
class ExtractedInfo(BaseModel):
|
||||
"""AI 从 query 中提取的关键信息"""
|
||||
|
||||
rewrite: str = Field(default="", description="改写后的 query")
|
||||
keywords: list[str] = Field(default_factory=list, description="关键词")
|
||||
entities: list[str] = Field(default_factory=list, description="实体(人名/产品/技术等)")
|
||||
intent: str = Field(default="", description="查询意图分类")
|
||||
time_range: str = Field(default="", description="时间范围,空表示无")
|
||||
categories: list[str] = Field(default_factory=list, description="命中类目名(含未达阈值的候选)")
|
||||
|
||||
|
||||
class SearchResponse(BaseModel):
|
||||
"""检索响应"""
|
||||
|
||||
@@ -28,3 +40,5 @@ class SearchResponse(BaseModel):
|
||||
hits: list[SearchHit] = Field(default_factory=list, description="命中结果列表")
|
||||
routed_categories: list[str] = Field(default_factory=list, description="query 路由命中的类目")
|
||||
fallback: bool = Field(default=False, description="是否走了全库兜底路径")
|
||||
extracted_info: ExtractedInfo | None = Field(default=None, description="AI 提取的 query 关键信息")
|
||||
summary: str | None = Field(default=None, description="检索结果 AI 总结,仅 summarize=true 时返回")
|
||||
|
||||
+199
-3
@@ -115,11 +115,56 @@
|
||||
.status-running { background: #eff6ff; color: #1d4ed8; border-color: #93c5fd; }
|
||||
.status-done { background: #f0fdf4; color: #15803d; border-color: #86efac; }
|
||||
.status-failed { background: #fef2f2; color: #b91c1c; border-color: #fca5a5; }
|
||||
header { position: relative; }
|
||||
.user-area { position: absolute; top: 14px; right: 24px; display: flex; align-items: center; gap: 8px; }
|
||||
.user-area .action { padding: 4px 10px; }
|
||||
.login-overlay {
|
||||
position: fixed; inset: 0; background: rgba(0,0,0,0.45);
|
||||
display: flex; align-items: center; justify-content: center; z-index: 100;
|
||||
}
|
||||
.login-box {
|
||||
background: #fff; border-radius: 8px; padding: 24px 28px; width: 320px;
|
||||
box-shadow: 0 6px 20px rgba(0,0,0,0.25);
|
||||
}
|
||||
.login-box h2 { font-size: 16px; margin: 0 0 16px; }
|
||||
.extracted-info {
|
||||
background: #f0f9ff; border: 1px solid #bae6fd; border-radius: 6px;
|
||||
padding: 10px 12px; margin-bottom: 12px; font-size: 13px;
|
||||
}
|
||||
.extracted-info .ei-title { font-weight: 600; margin-bottom: 6px; color: #0369a1; }
|
||||
.extracted-info .ei-row { margin-bottom: 3px; }
|
||||
.extracted-info .ei-k { color: #6b7280; margin-right: 4px; }
|
||||
.summary-box {
|
||||
background: #f0fdf4; border: 1px solid #bbf7d0; border-radius: 6px;
|
||||
padding: 12px 14px; margin-bottom: 12px; font-size: 13px; white-space: pre-wrap;
|
||||
}
|
||||
.summary-box .sum-title { font-weight: 600; margin-bottom: 6px; color: #15803d; }
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div id="login-overlay" class="login-overlay">
|
||||
<div class="login-box">
|
||||
<h2>登录知识库后台</h2>
|
||||
<div class="error-bar hidden" id="error-login"></div>
|
||||
<form id="login-form">
|
||||
<div class="field">
|
||||
<label for="login-username">用户名</label>
|
||||
<input type="text" id="login-username" name="username" required autocomplete="username">
|
||||
</div>
|
||||
<div class="field">
|
||||
<label for="login-password">密码</label>
|
||||
<input type="password" id="login-password" name="password" required autocomplete="current-password">
|
||||
</div>
|
||||
<button type="submit" class="primary">登录</button>
|
||||
</form>
|
||||
</div>
|
||||
</div>
|
||||
<header>
|
||||
<h1>知识库管理后台</h1>
|
||||
<div id="user-area" class="user-area hidden">
|
||||
<span class="muted" id="user-name"></span>
|
||||
<button type="button" class="action" id="btn-logout">登出</button>
|
||||
</div>
|
||||
<nav id="nav">
|
||||
<button type="button" data-target="section-overview" class="active">概览</button>
|
||||
<button type="button" data-target="section-docs">文档管理</button>
|
||||
@@ -174,6 +219,22 @@
|
||||
</div>
|
||||
<button type="submit" class="primary" id="btn-ingest-submit">提交入库</button>
|
||||
</form>
|
||||
<h3 style="font-size:14px; margin-top:24px;">或上传文件</h3>
|
||||
<form id="upload-form">
|
||||
<div class="field">
|
||||
<label for="upload-file">选择文件(支持 .txt/.md/.html/.htm/.pdf/.docx)</label>
|
||||
<input type="file" id="upload-file" name="file" accept=".txt,.md,.html,.htm,.pdf,.docx" required>
|
||||
</div>
|
||||
<div class="field">
|
||||
<label for="upload-title">标题(可选,默认取文件名)</label>
|
||||
<input type="text" id="upload-title" name="title" placeholder="留空则使用文件名去扩展">
|
||||
</div>
|
||||
<div class="field">
|
||||
<label for="upload-source">来源(可选,默认 file:原文件名)</label>
|
||||
<input type="text" id="upload-source" name="source" placeholder="例如:manual / web">
|
||||
</div>
|
||||
<button type="submit" class="primary" id="btn-upload-submit">上传入库</button>
|
||||
</form>
|
||||
<div class="result-box hidden" id="ingest-result"></div>
|
||||
</section>
|
||||
|
||||
@@ -189,6 +250,9 @@
|
||||
<label for="search-topk">top_k</label>
|
||||
<input type="number" id="search-topk" name="top_k" value="5" min="1" max="50">
|
||||
</div>
|
||||
<div class="field">
|
||||
<label><input type="checkbox" id="search-summarize" name="summarize"> 对结果生成 AI 总结</label>
|
||||
</div>
|
||||
<button type="submit" class="primary" id="btn-search-submit">检索</button>
|
||||
</form>
|
||||
<div class="result-box hidden" id="search-result"></div>
|
||||
@@ -235,14 +299,76 @@ function hideError(boxId) {
|
||||
document.getElementById(boxId).classList.add("hidden");
|
||||
}
|
||||
|
||||
/* 统一 API 封装:code !== 0 抛错,网络错误同样捕获 */
|
||||
/* ---------- 认证 token 管理 ---------- */
|
||||
function getToken() { return localStorage.getItem("qmd_token") || ""; }
|
||||
function setToken(t) { localStorage.setItem("qmd_token", t); }
|
||||
function clearToken() { localStorage.removeItem("qmd_token"); }
|
||||
|
||||
function showLogin() {
|
||||
document.getElementById("login-overlay").classList.remove("hidden");
|
||||
document.getElementById("user-area").classList.add("hidden");
|
||||
loadedOnce = {};
|
||||
}
|
||||
|
||||
function afterLogin(user) {
|
||||
document.getElementById("login-overlay").classList.add("hidden");
|
||||
document.getElementById("user-area").classList.remove("hidden");
|
||||
document.getElementById("user-name").textContent = user.username + " (" + user.role + ")";
|
||||
loadedOnce = {};
|
||||
activateSection("section-overview");
|
||||
}
|
||||
|
||||
document.getElementById("login-form").addEventListener("submit", function (event) {
|
||||
event.preventDefault();
|
||||
hideError("error-login");
|
||||
var payload = {
|
||||
username: document.getElementById("login-username").value,
|
||||
password: document.getElementById("login-password").value
|
||||
};
|
||||
fetch("/api/v1/auth/login", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify(payload)
|
||||
}).then(function (resp) { return resp.json(); }).then(function (body) {
|
||||
if (body.code !== 0) {
|
||||
showError("error-login", { code: body.code, message: body.message });
|
||||
return;
|
||||
}
|
||||
setToken(body.data.access_token);
|
||||
afterLogin(body.data.user);
|
||||
}).catch(function (err) {
|
||||
showError("error-login", { code: "NETWORK", message: "登录请求失败: " + (err && err.message ? err.message : err) });
|
||||
});
|
||||
});
|
||||
|
||||
document.getElementById("btn-logout").addEventListener("click", function () {
|
||||
clearToken();
|
||||
showLogin();
|
||||
});
|
||||
|
||||
/* 统一 API 封装:自动注入 token,code !== 0 抛错,未认证跳登录 */
|
||||
function api(path, options) {
|
||||
options = options || {};
|
||||
options.headers = Object.assign({}, options.headers || {});
|
||||
var token = getToken();
|
||||
if (token && !options.headers["Authorization"]) {
|
||||
options.headers["Authorization"] = "Bearer " + token;
|
||||
}
|
||||
return fetch(path, options).then(function (resp) {
|
||||
if (resp.status === 401) {
|
||||
clearToken();
|
||||
showLogin();
|
||||
throw { code: 1003, message: "未认证或登录已过期,请重新登录" };
|
||||
}
|
||||
return resp.json().catch(function () {
|
||||
throw { code: "HTTP " + resp.status, message: "响应解析失败" };
|
||||
});
|
||||
}).then(function (body) {
|
||||
if (body.code !== 0) {
|
||||
if (body.code === 1003 || body.code === 1005) {
|
||||
clearToken();
|
||||
showLogin();
|
||||
}
|
||||
throw { code: body.code, message: body.message };
|
||||
}
|
||||
return body.data;
|
||||
@@ -612,6 +738,38 @@ document.getElementById("ingest-form").addEventListener("submit", function (even
|
||||
});
|
||||
});
|
||||
|
||||
document.getElementById("upload-form").addEventListener("submit", function (event) {
|
||||
event.preventDefault();
|
||||
hideError("error-ingest");
|
||||
stopIngestPolling();
|
||||
var fileInput = document.getElementById("upload-file");
|
||||
if (!fileInput.files || fileInput.files.length === 0) {
|
||||
showError("error-ingest", { code: "VALIDATION", message: "请选择文件" });
|
||||
return;
|
||||
}
|
||||
var submitBtn = document.getElementById("btn-upload-submit");
|
||||
submitBtn.disabled = true;
|
||||
submitBtn.textContent = "上传中…";
|
||||
var formData = new FormData();
|
||||
formData.append("file", fileInput.files[0]);
|
||||
var title = document.getElementById("upload-title").value;
|
||||
var source = document.getElementById("upload-source").value;
|
||||
if (title) { formData.append("title", title); }
|
||||
if (source) { formData.append("source", source); }
|
||||
api("/api/v1/documents/upload", {
|
||||
method: "POST",
|
||||
body: formData
|
||||
}).then(function (data) {
|
||||
renderIngestHeader(data.task_id, data.status || "pending");
|
||||
startIngestPolling(data.task_id);
|
||||
}).catch(function (err) {
|
||||
showError("error-ingest", err);
|
||||
}).finally(function () {
|
||||
submitBtn.disabled = false;
|
||||
submitBtn.textContent = "上传入库";
|
||||
});
|
||||
});
|
||||
|
||||
/* ---------- 4. 检索测试台 ---------- */
|
||||
|
||||
document.getElementById("search-form").addEventListener("submit", function (event) {
|
||||
@@ -621,7 +779,8 @@ document.getElementById("search-form").addEventListener("submit", function (even
|
||||
submitBtn.disabled = true;
|
||||
var payload = {
|
||||
query: document.getElementById("search-query").value,
|
||||
top_k: parseInt(document.getElementById("search-topk").value, 10) || 5
|
||||
top_k: parseInt(document.getElementById("search-topk").value, 10) || 5,
|
||||
summarize: document.getElementById("search-summarize").checked
|
||||
};
|
||||
api("/api/v1/search", {
|
||||
method: "POST",
|
||||
@@ -636,6 +795,13 @@ document.getElementById("search-form").addEventListener("submit", function (even
|
||||
});
|
||||
});
|
||||
|
||||
function eiRow(k, v) {
|
||||
var row = el("div", null, "ei-row");
|
||||
row.appendChild(el("span", k + ":", "ei-k"));
|
||||
row.appendChild(el("span", v || "(无)"));
|
||||
return row;
|
||||
}
|
||||
|
||||
function renderSearchResult(data) {
|
||||
var box = document.getElementById("search-result");
|
||||
clearChildren(box);
|
||||
@@ -650,6 +816,26 @@ function renderSearchResult(data) {
|
||||
box.appendChild(el("div", "fallback:本次检索触发了回退策略(路由类目无命中,已降级全局检索)", "fallback-flag"));
|
||||
}
|
||||
|
||||
if (data.summary) {
|
||||
var sumBox = el("div", null, "summary-box");
|
||||
sumBox.appendChild(el("div", "AI 总结", "sum-title"));
|
||||
sumBox.appendChild(el("div", data.summary));
|
||||
box.appendChild(sumBox);
|
||||
}
|
||||
|
||||
var ei = data.extracted_info;
|
||||
if (ei) {
|
||||
var eiBox = el("div", null, "extracted-info");
|
||||
eiBox.appendChild(el("div", "AI 提取的关键信息", "ei-title"));
|
||||
eiBox.appendChild(eiRow("改写", ei.rewrite));
|
||||
eiBox.appendChild(eiRow("关键词", (ei.keywords || []).join(", ")));
|
||||
eiBox.appendChild(eiRow("实体", (ei.entities || []).join(", ")));
|
||||
eiBox.appendChild(eiRow("意图", ei.intent));
|
||||
eiBox.appendChild(eiRow("时间范围", ei.time_range));
|
||||
eiBox.appendChild(eiRow("命中类目", (ei.categories || []).join(", ")));
|
||||
box.appendChild(eiBox);
|
||||
}
|
||||
|
||||
var hits = data.hits || [];
|
||||
box.appendChild(el("div", "命中 " + hits.length + " 条", "muted"));
|
||||
hits.forEach(function (hit) {
|
||||
@@ -686,7 +872,17 @@ function loadCategories() {
|
||||
|
||||
/* ---------- 初始化 ---------- */
|
||||
|
||||
activateSection("section-overview");
|
||||
(function init() {
|
||||
if (!getToken()) {
|
||||
showLogin();
|
||||
return;
|
||||
}
|
||||
api("/api/v1/auth/me").then(function (user) {
|
||||
afterLogin(user);
|
||||
}).catch(function () {
|
||||
showLogin();
|
||||
});
|
||||
})();
|
||||
</script>
|
||||
</body>
|
||||
</html>
|
||||
|
||||
Reference in New Issue
Block a user