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
+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,
)