feat: 新增多格式文件上传入库与认证体系

- 新增 JWT 认证模块,支持登录/注册/用户管理
- 新增文件上传接口,支持 .txt/.md/.html/.pdf/.docx 等格式解析入库
- 新增检索结果 AI 总结功能
- 新增文本去重缓存机制
- 新增全局认证夹具简化测试
- 新增配置项与环境变量支持
- 完善文档与测试覆盖
This commit is contained in:
2026-07-30 10:30:15 +08:00
parent 51dc8dc4f6
commit dce9e31bde
31 changed files with 3021 additions and 45 deletions
+62
View File
@@ -0,0 +1,62 @@
"""认证 APIPOST /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
View File
@@ -1,13 +1,19 @@
"""文档 APIPOST /api/v1/documents 异步入库 + 任务查询 + 文档管理(列表/详情/删除)"""
"""文档 APIPOST /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_atdone 附 resultfailed 附 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 -3
View File
@@ -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:
+5 -4
View File
@@ -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 缓存:命中直接返回缓存的响应;未命中走检索流程并回写缓存。
+21
View File
@@ -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}
+213
View File
@@ -0,0 +1,213 @@
"""JWT 认证核心:密码哈希、token 签发/验签、用户存储、FastAPI 鉴权依赖
用户持久化在 Rediskey: 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 secretsettings 为空时自动生成,仅开发用)
_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
+204
View File
@@ -0,0 +1,204 @@
"""多格式文件文本提取:按扩展名分发到对应解析器
支持的扩展名:
- .txt / .mdUTF-8 解码(errors="replace" 兜底)
- .html / .htm:标准库 html.parser 剥离标签提取可见文本
- .pdfpypdf 逐页 extract_text 拼接;文本层为空(扫描件/图片型)时
自动降级为 OCRpypdfium2 渲染 + rapidocr-onnxruntime 识别)
- .docxpython-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())
+63 -5
View File
@@ -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)
+28 -3
View File
@@ -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}"
)
+66
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+42
View File
@@ -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="登录用户信息")
+1
View File
@@ -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")
+14
View File
@@ -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
View File
@@ -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>