feat: 新增多格式文件上传入库与认证体系
- 新增 JWT 认证模块,支持登录/注册/用户管理 - 新增文件上传接口,支持 .txt/.md/.html/.pdf/.docx 等格式解析入库 - 新增检索结果 AI 总结功能 - 新增文本去重缓存机制 - 新增全局认证夹具简化测试 - 新增配置项与环境变量支持 - 完善文档与测试覆盖
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user