Initial commit: QMDSearch 分层信息检索服务

- FastAPI + Qdrant + Redis + Ollama 技术栈
- L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合)
- 文档三级总结与 2.5 级回退
- query 解析路由与分类
- /admin 管理页面
This commit is contained in:
2026-07-29 21:24:40 +08:00
commit 51dc8dc4f6
83 changed files with 10794 additions and 0 deletions
View File
View File
+26
View File
@@ -0,0 +1,26 @@
"""统一 API 响应包装与业务异常
响应格式:{"code": 0, "data": ..., "message": "ok"}
错误码约定:0 成功;1001 请求参数校验失败;2000 服务器内部错误。
"""
from typing import Any
def ok(data: Any) -> dict[str, Any]:
"""成功响应包装"""
return {"code": 0, "data": data, "message": "ok"}
def error(code: int, message: str) -> dict[str, Any]:
"""错误响应包装"""
return {"code": code, "data": None, "message": message}
class ApiError(Exception):
"""业务异常:携带错误码与消息,由全局异常处理器转为统一错误响应"""
def __init__(self, code: int, message: str) -> None:
self.code = code
self.message = message
super().__init__(message)
View File
+111
View File
@@ -0,0 +1,111 @@
"""文档 APIPOST /api/v1/documents 异步入库 + 任务查询 + 文档管理(列表/详情/删除)"""
from typing import Any
import structlog
from fastapi import APIRouter, Query
from fastapi.responses import JSONResponse
from app.api.response import ApiError, ok
from app.config import Settings
from app.core.ingest_tasks import IngestTaskManager
from app.core.ingestion import Ingester
from app.models.document import DocumentInput
from app.services.qdrant import QdrantService
from app.services.redis import RedisCache, get_cache
logger = structlog.get_logger()
router = APIRouter(prefix="/api/v1", tags=["document"])
# 模块级懒加载单例,避免每请求重建 Ingester 及其下游依赖
_ingester: Ingester | None = None
_qdrant: QdrantService | None = None
_task_manager: IngestTaskManager | None = None
def _get_ingester() -> Ingester:
global _ingester
if _ingester is None:
_ingester = Ingester()
return _ingester
def _get_qdrant() -> QdrantService:
global _qdrant
if _qdrant is None:
_qdrant = QdrantService()
return _qdrant
def _get_task_manager() -> IngestTaskManager:
"""入库任务管理器懒加载单例
Redis 沿用全局缓存单例(RedisCache 读写全容错,不可用时镜像写失败仅告警);
获取缓存实例异常时传 None,退化为纯内存模式。
"""
global _task_manager
if _task_manager is None:
try:
redis: RedisCache | None = get_cache()
except Exception:
logger.warning("Redis 缓存不可用,入库任务状态仅保留在内存", exc_info=True)
redis = None
_task_manager = IngestTaskManager(ingester=_get_ingester(), redis=redis, settings=Settings())
return _task_manager
@router.post("/documents")
async def ingest_document(doc: DocumentInput) -> JSONResponse:
"""文档入库入口:登记异步任务并返回 202 + task_id,入库结果经任务查询端点获取"""
if not doc.text.strip():
raise ApiError(1001, "文档内容不能为空")
task_id = await _get_task_manager().submit(doc)
return JSONResponse(status_code=202, content=ok({"task_id": task_id, "status": "pending"}))
@router.get("/documents/tasks/{task_id}")
async def get_ingest_task(task_id: str) -> dict[str, Any]:
"""查询入库任务状态:含 task_id/status/created_at/updated_atdone 附 resultfailed 附 error"""
task = await _get_task_manager().get(task_id)
if task is None:
raise ApiError(1004, "任务不存在")
return ok(task)
@router.get("/documents")
async def list_documents(
limit: int = Query(default=20, ge=1, le=100),
offset: str | None = None,
) -> dict[str, Any]:
"""分页列出文档(L1 摘要),返回 items 与下一页游标 next_offset"""
try:
items, next_offset = await _get_qdrant().scroll_l1(limit=limit, offset=offset)
except Exception as exc:
logger.error("文档列表查询失败", error=str(exc))
raise ApiError(2000, f"文档列表查询失败: {exc}") from exc
return ok({"items": items, "next_offset": next_offset})
@router.get("/documents/{doc_id}")
async def get_document(doc_id: str) -> dict[str, Any]:
"""获取文档详情:L1 记录 + L2/L3 节点 + chunks 数量"""
try:
detail = await _get_qdrant().get_doc_detail(doc_id)
except Exception as exc:
logger.error("文档详情查询失败", doc_id=doc_id, error=str(exc))
raise ApiError(2000, f"文档详情查询失败: {exc}") from exc
if detail is None:
raise ApiError(1004, "文档不存在")
return ok(detail)
@router.delete("/documents/{doc_id}")
async def delete_document(doc_id: str) -> dict[str, Any]:
"""删除文档:四层集合中该 doc_id 的所有点;幂等,不存在也返回成功(删除数全 0)"""
try:
deleted = await _get_qdrant().delete_by_doc_id(doc_id)
except Exception as exc:
logger.error("文档删除失败", doc_id=doc_id, error=str(exc))
raise ApiError(2000, f"文档删除失败: {exc}") from exc
return ok({"doc_id": doc_id, "deleted": deleted, "deleted_total": sum(deleted.values())})
+75
View File
@@ -0,0 +1,75 @@
"""知识分类 APIGET /api/v1/knowledge/categories、GET /api/v1/knowledge/stats"""
from functools import lru_cache
from typing import Any
import structlog
from fastapi import APIRouter
from app.api.response import ApiError, ok
from app.config import settings
from app.models.knowledge import UNCATEGORIZED, TaxonomyCategory, load_taxonomy
from app.services.qdrant import ALL_COLLECTIONS, COLLECTION_L1, QdrantService
logger = structlog.get_logger()
router = APIRouter(prefix="/api/v1", tags=["knowledge"])
# 类目分布统计的分页大小
_STATS_SCROLL_PAGE_SIZE = 100
# 模块级懒加载单例,避免每请求重建 QdrantService
_qdrant: QdrantService | None = None
def _get_qdrant() -> QdrantService:
global _qdrant
if _qdrant is None:
_qdrant = QdrantService()
return _qdrant
@lru_cache(maxsize=1)
def _get_taxonomy() -> list[TaxonomyCategory]:
"""加载并缓存 taxonomy 类目集(进程内只加载一次)"""
return load_taxonomy(settings.taxonomy_path)
@router.get("/knowledge/categories")
async def list_categories() -> 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]:
"""返回四层集合规模与 L1 类目分布统计"""
service = _get_qdrant()
try:
collections = {collection: await service.count(collection) for collection in ALL_COLLECTIONS}
categories = await _aggregate_l1_categories(service)
except Exception as exc:
logger.error("获取知识库统计失败", error=str(exc))
raise ApiError(2000, f"获取知识库统计失败: {exc}") from exc
return ok(
{
"collections": collections,
"categories": categories,
"uncategorized_count": categories.get(UNCATEGORIZED, 0),
"documents_total": collections[COLLECTION_L1],
}
)
async def _aggregate_l1_categories(service: QdrantService) -> dict[str, int]:
"""分页遍历 L1 文档,按 category 聚合文档数(空类目归入 uncategorized"""
categories: dict[str, int] = {}
offset: str | None = None
while True:
items, offset = await service.scroll_l1(limit=_STATS_SCROLL_PAGE_SIZE, offset=offset)
for item in items:
category = item.get("category") or UNCATEGORIZED
categories[category] = categories.get(category, 0) + 1
if offset is None:
return categories
+58
View File
@@ -0,0 +1,58 @@
"""检索 APIPOST /api/v1/search"""
from hashlib import sha256
from typing import Any
import structlog
from fastapi import APIRouter
from app.api.response import ApiError, ok
from app.core.retriever import Retriever
from app.models.search import SearchRequest
from app.services.redis import get_cache
logger = structlog.get_logger()
router = APIRouter(prefix="/api/v1", tags=["search"])
# 模块级懒加载单例,避免每请求重建 Qdrant/Embedding client
_retriever: Retriever | None = None
def _get_retriever() -> Retriever:
global _retriever
if _retriever is None:
_retriever = Retriever()
return _retriever
def _cache_key(request: SearchRequest) -> str:
"""检索缓存键:query + top_k 的短哈希(top_k 影响结果集,需参与键计算)"""
digest = sha256((request.query + "|" + str(request.top_k)).encode()).hexdigest()[:16]
return f"search:{digest}"
@router.post("/search")
async def search(request: SearchRequest) -> dict[str, Any]:
"""分层检索入口,返回统一包装的 SearchResponse
先查 Redis 缓存:命中直接返回缓存的响应;未命中走检索流程并回写缓存。
缓存读写失败均降级为无缓存行为,不影响检索。
"""
cache_key = _cache_key(request)
cached = await get_cache().get_json(cache_key)
if cached is not None:
logger.info("检索缓存命中", query=request.query, cache_key=cache_key)
return cached
try:
response = await _get_retriever().search(request)
except ApiError:
raise
except Exception as exc:
logger.exception("检索失败", query=request.query)
raise ApiError(2000, f"检索失败: {exc}") from exc
result = ok(response.model_dump())
await get_cache().set_json(cache_key, result)
return result
+60
View File
@@ -0,0 +1,60 @@
from pydantic_settings import BaseSettings
class Settings(BaseSettings):
"""应用配置,通过环境变量注入"""
# 应用
app_name: str = "QMDSearch"
log_level: str = "info"
# 嵌入模型
embedding_provider: str = "openai" # openai | local
openai_api_key: str = ""
openai_base_url: str = "https://api.openai.com/v1"
embedding_model: str = "text-embedding-3-small"
embedding_dimension: int = 1536
# Ollama 本地模型(用于文档三级总结)
ollama_base_url: str = "http://localhost:11434"
ollama_model: str = "qwen2.5:1.5b" # 备选: qwen2.5:3b
ollama_embedding_model: str = "bge-m3" # embedding_provider=local 时使用的嵌入模型
# Qdrant
qdrant_host: str = "localhost"
qdrant_port: int = 6333
# Redis
redis_url: str = "redis://localhost:6379/0"
# 检索参数
retrieval_top_k: int = 20 # L2 语义检索召回数
retrieval_final_k: int = 5 # L3 重排后返回数
# 文档入库参数
summary_min_text_length: int = 500 # 低于此字符数触发 2.5 级回退
# 入库异步任务
ingest_max_concurrency: int = 2 # 入库后台任务并发上限
ingest_task_ttl_done: int = 86400 # 任务状态 Redis 保留秒数(进行中与已完成,24h)
ingest_task_ttl_failed: int = 604800 # 失败任务状态 Redis 保留秒数(7 天)
# 知识分类(taxonomy
taxonomy_path: str = "" # taxonomy JSON 文件路径,为空用内置默认
classify_confidence_threshold: float = 0.6 # 低于此值归 uncategorized
classify_max_categories: int = 3 # query 路由命中类目数上限,超过走全库兜底
# 分层检索参数
l1_doc_top_n: int = 10 # L1 层候选文档数
l2_section_top_n: int = 5 # L2 层候选 section 数
l3_top_n: int = 10 # L3 层定位数
sparse_enabled: bool = True # 是否启用稀疏检索
cache_ttl: int = 300 # Redis 缓存秒数
# 分块参数
chunk_max_chars: int = 800 # chunk 超长二次切分阈值
model_config = {"env_prefix": "", "case_sensitive": False}
settings = Settings()
View File
+121
View File
@@ -0,0 +1,121 @@
"""文档 chunk 切分器
按文档原生标题树切分 chunk
- 结构化文本:每个标题起点切分 section(标题行到下一个标题前),
超长 section 按空行段落二次切分,单段落仍超长则硬切
- 无结构文本:直接按空行段落累加切分
每个 chunk 记录 section_path(祖先标题链," / " 连接),与 L2 大纲节点互相定位。
"""
import re
import structlog
from app.config import settings
from app.core.headings import Heading, parse_headings
from app.models.document import ChunkModel
logger = structlog.get_logger()
# 段落分隔:一个或多个空行
_PARAGRAPH_SPLIT_PATTERN = re.compile(r"\n\s*\n")
class Chunker:
"""按标题树切分文档 chunk"""
def __init__(self, max_chars: int = settings.chunk_max_chars) -> None:
self.max_chars = max_chars
def chunk(self, text: str, doc_id: str) -> list[ChunkModel]:
"""将文档文本切分为 chunk 列表
Args:
text: 文档纯文本内容
doc_id: 文档 ID
Returns:
list[ChunkModel]: 切分结果,chunk_index 从 0 递增
"""
stripped = text.strip()
if not stripped:
return []
# 全文不超长:整篇单 chunk
if len(stripped) <= self.max_chars:
return [ChunkModel(doc_id=doc_id, chunk_index=0, text=stripped)]
headings = parse_headings(stripped)
chunks: list[ChunkModel] = []
if headings:
# 结构化:先按标题切分 section,再按长度二次切分
for section_text, section_path in self._split_sections(stripped, headings):
for piece in self._split_by_length(section_text):
chunks.append(
ChunkModel(doc_id=doc_id, chunk_index=len(chunks), text=piece, section_path=section_path)
)
else:
# 无结构:直接按段落累加切分
for piece in self._split_by_length(stripped):
chunks.append(ChunkModel(doc_id=doc_id, chunk_index=len(chunks), text=piece))
logger.info("文档切分完成", doc_id=doc_id, chunks_count=len(chunks), has_headings=bool(headings))
return chunks
def _split_sections(self, text: str, headings: list[Heading]) -> list[tuple[str, str]]:
"""按标题树切分 section,返回 (section 文本, section_path) 列表
每个标题起点切分一个 section,section 文本含标题行本身;
section_path 为祖先标题链(含自身标题),用 " / " 连接;
首个标题前的引导正文归入无前缀 sectionsection_path 为空)。
"""
lines = text.splitlines()
sections: list[tuple[str, str]] = []
# 首个标题前的引导内容
preamble = "\n".join(lines[: headings[0].line_index]).strip()
if preamble:
sections.append((preamble, ""))
# 维护祖先标题栈:遇到同级或更高级标题时弹栈
stack: list[Heading] = []
for i, heading in enumerate(headings):
while stack and stack[-1].level >= heading.level:
stack.pop()
stack.append(heading)
end = headings[i + 1].line_index if i + 1 < len(headings) else len(lines)
section_text = "\n".join(lines[heading.line_index : end]).strip()
section_path = " / ".join(h.title for h in stack)
sections.append((section_text, section_path))
return sections
def _split_by_length(self, text: str) -> list[str]:
"""按 max_chars 切分文本:先按空行段落累加,单段落超长则硬切"""
if len(text) <= self.max_chars:
return [text]
pieces: list[str] = []
current = ""
for paragraph in _PARAGRAPH_SPLIT_PATTERN.split(text):
paragraph = paragraph.strip()
if not paragraph:
continue
candidate = f"{current}\n\n{paragraph}" if current else paragraph
if len(candidate) <= self.max_chars:
current = candidate
continue
if current:
pieces.append(current)
current = ""
# 单段落仍超长:按 max_chars 硬切
if len(paragraph) > self.max_chars:
pieces.extend(paragraph[i : i + self.max_chars] for i in range(0, len(paragraph), self.max_chars))
else:
current = paragraph
if current:
pieces.append(current)
return pieces
+139
View File
@@ -0,0 +1,139 @@
"""文档分类器
入库链路第二步:基于 L1 总结,用 Ollama 小模型将文档判定为 taxonomy 中的
主类目 + 附加标签:
- LLM 输出解析失败 / 类目名不在 taxonomy → 归 uncategorizedconfidence=0.0
- 置信度低于阈值 → 主类目归 uncategorized,候选类目名保留进 tags(软召回用)
"""
import json
import re
import structlog
from app.config import settings
from app.models.knowledge import UNCATEGORIZED, CategoryResult, TaxonomyCategory, load_taxonomy
from app.services.ollama import OllamaClient
logger = structlog.get_logger()
# 从 LLM 输出中提取第一个 {...} JSON 块(贪婪匹配到最后的 },兼容嵌套对象)
_JSON_BLOCK_RE = re.compile(r"\{.*\}", re.DOTALL)
def _extract_json(raw: str) -> dict | None:
"""从 LLM 输出中提取 JSON 对象
先尝试直接解析;失败则用正则提取第一个 {...} 块再解析。
返回 None 表示无法提取出合法的 JSON 对象。
"""
text = raw.strip()
try:
data = json.loads(text)
return data if isinstance(data, dict) else None
except json.JSONDecodeError:
pass
match = _JSON_BLOCK_RE.search(text)
if not match:
return None
try:
data = json.loads(match.group(0))
return data if isinstance(data, dict) else None
except json.JSONDecodeError:
return None
class Classifier:
"""文档分类器:将 L1 总结判定为 taxonomy 主类目 + 附加标签"""
def __init__(self, ollama: OllamaClient | None = None, taxonomy: list[TaxonomyCategory] | None = None) -> None:
self.ollama = ollama or OllamaClient()
self.taxonomy = taxonomy if taxonomy is not None else load_taxonomy()
# 合法类目名集合(含 uncategorized
self._valid_names = {c.name for c in self.taxonomy}
async def classify(self, l1_summary: str, title: str = "") -> CategoryResult:
"""对文档进行分类判定
Args:
l1_summary: 文档 L1 总结
title: 文档标题(可选,辅助判定)
Returns:
CategoryResult: 主类目 / 附加标签 / 置信度
"""
prompt = self._build_prompt(l1_summary, title)
raw = await self.ollama.generate(prompt, json_mode=True)
data = _extract_json(raw)
if data is None:
logger.warning("分类失败:LLM 输出非合法 JSON", title=title, output=raw[:200])
return CategoryResult(main_category=UNCATEGORIZED, tags=[], confidence=0.0)
main_category = data.get("main_category")
confidence = data.get("confidence")
if (
not isinstance(main_category, str)
or main_category not in self._valid_names
or not isinstance(confidence, (int, float))
or not 0 <= confidence <= 1
):
logger.warning(
"分类失败:类目名不在 taxonomy 或 confidence 非法",
title=title,
main_category=main_category,
confidence=confidence,
)
return CategoryResult(main_category=UNCATEGORIZED, tags=[], confidence=0.0)
tags = self._clean_tags(data.get("tags"), main_category)
confidence = float(confidence)
# 低置信度软召回:主类目归 uncategorized,候选类目名保留进 tags
if confidence < settings.classify_confidence_threshold:
logger.info(
"分类置信度低于阈值,归入 uncategorized",
title=title,
candidate=main_category,
confidence=confidence,
threshold=settings.classify_confidence_threshold,
)
if main_category != UNCATEGORIZED and main_category not in tags:
tags.insert(0, main_category)
return CategoryResult(main_category=UNCATEGORIZED, tags=tags, confidence=confidence)
return CategoryResult(main_category=main_category, tags=tags, confidence=confidence)
def _build_prompt(self, l1_summary: str, title: str) -> str:
"""构造分类 prompt:列出全部 taxonomy 类目(含 uncategorized 及其用途说明)"""
category_lines = []
for c in self.taxonomy:
if c.name == UNCATEGORIZED:
category_lines.append(f"- {c.name}: 当文档跨多个类目或无法明确归入其他类目时选择此类目")
else:
category_lines.append(f"- {c.name}: {c.description}")
category_block = "\n".join(category_lines)
prompt = (
"你是知识库分类助手。请根据文档标题和总结,判断文档最适合归入以下哪个类目。\n\n"
f"可选类目:\n{category_block}\n\n"
"要求:\n"
"- main_category 只能从上面的类目名中选择,不要输出其他名称;\n"
"- tags 为 0~3 个附加标签(词或短语),且不能包含 main_category 本身;\n"
"- confidence 为 0~1 之间的小数,表示对主类目判断的置信度;\n"
"- 只输出 JSON,不要输出任何其他内容。\n\n"
'输出格式:{"main_category": "类目名", "tags": ["标签"], "confidence": 0.0}\n\n'
)
if title:
prompt += f"文档标题:{title}\n"
prompt += f"文档总结:{l1_summary}"
return prompt
@staticmethod
def _clean_tags(raw_tags: object, main_category: str) -> list[str]:
"""清洗 LLM 输出的 tags:仅保留非空字符串、剔除主类目名、最多 3 个"""
if not isinstance(raw_tags, list):
return []
tags = [t for t in raw_tags if isinstance(t, str) and t and t != main_category]
return tags[:3]
+99
View File
@@ -0,0 +1,99 @@
"""Dense 向量嵌入服务
提供统一的 EmbeddingService 接口,支持两种 provider
- openaiOpenAI 兼容 APIAsyncOpenAI
- local:本地 Ollama /api/embed 接口
"""
from typing import Protocol, runtime_checkable
import httpx
import structlog
from openai import AsyncOpenAI
from app.config import settings
logger = structlog.get_logger()
@runtime_checkable
class EmbeddingService(Protocol):
"""统一嵌入服务接口"""
async def embed(self, texts: list[str]) -> list[list[float]]:
"""批量生成文本向量
Args:
texts: 待嵌入文本列表,空列表时直接返回空列表
Returns:
与输入等长的向量列表
"""
...
def _check_dimension(vectors: list[list[float]], provider: str) -> None:
"""返回维度与配置不一致时告警(不抛错),每次调用最多提示一次"""
for vec in vectors:
if len(vec) != settings.embedding_dimension:
logger.warning(
"嵌入向量维度与配置不一致",
provider=provider,
actual=len(vec),
expected=settings.embedding_dimension,
)
return
class OpenAIEmbeddingService:
"""OpenAI 兼容 API 嵌入服务"""
def __init__(self, api_key: str, base_url: str, model: str) -> None:
self._client = AsyncOpenAI(api_key=api_key, base_url=base_url)
self._model = model
async def embed(self, texts: list[str]) -> list[list[float]]:
if not texts:
return []
resp = await self._client.embeddings.create(model=self._model, input=texts)
vectors = [item.embedding for item in resp.data]
_check_dimension(vectors, provider="openai")
logger.debug("OpenAI 嵌入完成", model=self._model, count=len(vectors))
return vectors
class LocalEmbeddingService:
"""本地 Ollama 嵌入服务(/api/embed"""
def __init__(self, base_url: str, model: str, timeout: float = 60.0) -> None:
self.base_url = base_url.rstrip("/")
self.model = model
self.timeout = timeout
async def embed(self, texts: list[str]) -> list[list[float]]:
if not texts:
return []
url = f"{self.base_url}/api/embed"
payload = {"model": self.model, "input": texts}
async with httpx.AsyncClient(timeout=self.timeout) as client:
resp = await client.post(url, json=payload)
resp.raise_for_status()
data = resp.json()
vectors: list[list[float]] = data.get("embeddings", [])
_check_dimension(vectors, provider="local")
logger.debug("Ollama 嵌入完成", model=self.model, count=len(vectors))
return vectors
def create_embedding_service() -> EmbeddingService:
"""按 settings.embedding_provider 创建嵌入服务实例"""
if settings.embedding_provider == "local":
return LocalEmbeddingService(
base_url=settings.ollama_base_url,
model=settings.ollama_embedding_model,
)
return OpenAIEmbeddingService(
api_key=settings.openai_api_key,
base_url=settings.openai_base_url,
model=settings.embedding_model,
)
+94
View File
@@ -0,0 +1,94 @@
"""标题树解析器
从纯文本中解析文档原生标题结构,支持两类模式:
- Markdown ATX 标题(# ~ ####### 数量即层级)
- 中文编号标题(第X章/节/篇、一、1.1 等编号,行长度不超过 60 字符)
解析结果用于 L2 大纲生成与 chunk 的 section 切分。
"""
import re
from pydantic import BaseModel, Field
class Heading(BaseModel):
"""文档标题节点"""
title: str = Field(description="标题文本")
level: int = Field(description="标题层级,1 为最顶层")
line_index: int = Field(description="标题所在行号(从 0 开始)")
# Markdown ATX 标题:1~6 个 # 后跟空白
_ATX_PATTERN = re.compile(r"^(#{1,6})\s+(.+?)\s*$")
# 中文篇章节编号:第X章/第X节/第X篇(章/篇=1 级,节=2 级)
_CN_CHAPTER_PATTERN = re.compile(r"^第[一二三四五六七八九十百\d]+([章节篇])")
# 中文序号:一、二、……,固定 1 级
_CN_ENUM_PATTERN = re.compile(r"^[一二三四五六七八九十]+、")
# 数字编号:1. / 1、/ 1.1 / 1.1.1 等,按点分段数定层级
_NUM_PATTERN = re.compile(r"^(\d+(?:\.\d+)*)[、.\s]")
# 编号类标题行的最大长度,超过则视为正文
_MAX_HEADING_LINE_LENGTH = 60
# 第X[章节篇] 后缀对应的层级
_CN_CHAPTER_LEVELS = {"": 1, "": 2, "": 1}
def parse_headings(text: str) -> list[Heading]:
"""解析文本中的标题,按行号升序返回
Args:
text: 文档纯文本内容
Returns:
list[Heading]: 标题列表,无标题时返回空列表
"""
headings: list[Heading] = []
for line_index, line in enumerate(text.splitlines()):
heading = _match_heading(line, line_index)
if heading is not None:
headings.append(heading)
return headings
def render_outline(headings: list[Heading]) -> str:
"""将标题树渲染为大纲文本,每行一个节点,按层级缩进"""
return "\n".join(f"{' ' * (h.level - 1)}- {h.title}" for h in headings)
def _match_heading(line: str, line_index: int) -> Heading | None:
"""匹配单行是否为标题,是则返回 Heading,否则返回 None"""
stripped = line.strip()
if not stripped:
return None
# Markdown ATX 标题(无长度限制)
match = _ATX_PATTERN.match(stripped)
if match:
return Heading(title=match.group(2), level=len(match.group(1)), line_index=line_index)
# 编号类标题有行长度限制,过长视为正文
if len(stripped) > _MAX_HEADING_LINE_LENGTH:
return None
# 第X章/节/篇
match = _CN_CHAPTER_PATTERN.match(stripped)
if match:
return Heading(title=stripped, level=_CN_CHAPTER_LEVELS[match.group(1)], line_index=line_index)
# 一、二、……
if _CN_ENUM_PATTERN.match(stripped):
return Heading(title=stripped, level=1, line_index=line_index)
# 数字编号,层级 = 点分段数(1.=1、1.1=2、1.1.1=3
match = _NUM_PATTERN.match(stripped)
if match:
level = match.group(1).count(".") + 1
return Heading(title=stripped, level=level, line_index=line_index)
return None
+172
View File
@@ -0,0 +1,172 @@
"""入库异步任务管理器
将文档入库包装为后台异步任务:submit 登记任务并立即返回 task_id,
后台受并发上限控制执行 Ingester.ingest,并按阶段推进任务状态。
内存注册表为主(记录 status/created_at/updated_at/result/error),
Redis 为持久镜像(key: ingest_task:{task_id}),每次状态迁移同步写入;
Redis 不可用或写入失败仅记录 warning,不影响任务执行。
"""
import asyncio
import uuid
from datetime import UTC, datetime
from enum import StrEnum
from typing import Any
import structlog
from app.config import Settings
from app.core.ingestion import Ingester, IngestionError
from app.models.document import DocumentInput
from app.services.redis import RedisCache
logger = structlog.get_logger()
# Redis 任务状态 key 前缀
REDIS_KEY_PREFIX = "ingest_task:"
class IngestTaskStatus(StrEnum):
"""入库任务状态"""
PENDING = "pending" # 已登记,排队等待执行
SUMMARIZING = "summarizing" # 三级总结中
CLASSIFYING = "classifying" # 分类判定中
EMBEDDING = "embedding" # 向量化中
WRITING = "writing" # 写入 Qdrant 中
DONE = "done" # 入库完成
FAILED = "failed" # 入库失败
# 终态集合
TERMINAL_STATUSES: frozenset[str] = frozenset({IngestTaskStatus.DONE, IngestTaskStatus.FAILED})
def _utc_now_iso() -> str:
"""当前 UTC 时间的 ISO8601 字符串"""
return datetime.now(UTC).isoformat()
class IngestTaskManager:
"""入库异步任务管理器:登记、后台执行、状态查询与 Redis 持久镜像"""
def __init__(self, ingester: Ingester, redis: RedisCache | None, settings: Settings) -> None:
self._ingester = ingester
self._redis = redis
self._settings = settings
self._tasks: dict[str, dict[str, Any]] = {}
self._semaphore = asyncio.Semaphore(settings.ingest_max_concurrency)
# 持有后台任务与镜像任务引用,避免被 GC 提前回收
self._background_tasks: set[asyncio.Task[None]] = set()
self._mirror_tasks: set[asyncio.Task[None]] = set()
async def submit(self, doc: DocumentInput) -> str:
"""登记入库任务并后台执行,立即返回 task_id"""
task_id = uuid.uuid4().hex
now = _utc_now_iso()
self._tasks[task_id] = {
"task_id": task_id,
"status": IngestTaskStatus.PENDING,
"created_at": now,
"updated_at": now,
"result": None,
"error": None,
}
self._schedule_mirror(task_id)
background = asyncio.create_task(self._run(task_id, doc))
self._background_tasks.add(background)
background.add_done_callback(self._background_tasks.discard)
logger.info("入库任务已登记", task_id=task_id, title=doc.title)
return task_id
async def get(self, task_id: str) -> dict[str, Any] | None:
"""查询任务状态:先查内存注册表,miss 再查 Redis 镜像,都没有返回 None"""
record = self._tasks.get(task_id)
if record is not None:
return record
if self._redis is None:
return None
try:
return await self._redis.get_json(f"{REDIS_KEY_PREFIX}{task_id}")
except Exception:
logger.warning("入库任务状态读取 Redis 失败,降级为未命中", task_id=task_id, exc_info=True)
return None
async def wait_done(self, task_id: str, timeout: float = 30.0) -> dict[str, Any]:
"""轮询内存注册表直到任务进入终态(done/failed)或超时
进入终态后会等待已调度的 Redis 镜像写完再返回;超时抛 TimeoutError。
"""
loop = asyncio.get_running_loop()
deadline = loop.time() + timeout
while True:
record = self._tasks.get(task_id)
if record is not None and record["status"] in TERMINAL_STATUSES:
if self._mirror_tasks:
await asyncio.gather(*self._mirror_tasks, return_exceptions=True)
return record
if loop.time() >= deadline:
raise TimeoutError(f"入库任务 {task_id}{timeout}s 内未进入终态")
await asyncio.sleep(0.01)
async def _run(self, task_id: str, doc: DocumentInput) -> None:
"""后台执行入库:并发限流 + 阶段状态推进 + 结果/错误落账"""
async with self._semaphore:
try:
result = await self._ingester.ingest(doc, progress_cb=lambda stage: self._on_progress(task_id, stage))
except IngestionError as exc:
# 入库已知失败:透传阶段与已产出的部分总结
self._finish_failed(
task_id,
{
"stage": exc.stage,
"message": str(exc),
"partial_summary": exc.summary.model_dump(mode="json") if exc.summary is not None else None,
},
)
except Exception as exc:
self._finish_failed(task_id, {"stage": "unknown", "message": str(exc), "partial_summary": None})
else:
self._tasks[task_id].update(
status=IngestTaskStatus.DONE,
updated_at=_utc_now_iso(),
result=result.model_dump(mode="json"),
)
self._schedule_mirror(task_id)
logger.info("入库任务完成", task_id=task_id)
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)
self._schedule_mirror(task_id)
logger.error("入库任务失败", task_id=task_id, stage=error["stage"], error=error["message"])
def _on_progress(self, task_id: str, stage: str) -> None:
"""Ingester 阶段回调:推进任务状态并同步镜像(同步函数,供 progress_cb 使用)"""
record = self._tasks.get(task_id)
if record is None:
return
record["status"] = stage
record["updated_at"] = _utc_now_iso()
self._schedule_mirror(task_id)
def _schedule_mirror(self, task_id: str) -> None:
"""将当前任务状态快照异步镜像到 Redis(同步上下文也可调用)"""
if self._redis is None:
return
mirror = asyncio.create_task(self._mirror_to_redis(task_id, dict(self._tasks[task_id])))
self._mirror_tasks.add(mirror)
mirror.add_done_callback(self._mirror_tasks.discard)
async def _mirror_to_redis(self, task_id: str, snapshot: dict[str, Any]) -> None:
"""写入 Redis 镜像:进行中与 done 用 ttl_donefailed 用 ttl_failed;写失败仅告警"""
ttl = (
self._settings.ingest_task_ttl_failed
if snapshot["status"] == IngestTaskStatus.FAILED
else self._settings.ingest_task_ttl_done
)
try:
await self._redis.set_json(f"{REDIS_KEY_PREFIX}{task_id}", snapshot, ttl=ttl)
except Exception:
logger.warning("入库任务状态镜像 Redis 失败", task_id=task_id, exc_info=True)
+306
View File
@@ -0,0 +1,306 @@
"""文档入库模块
入库流程:文档输入 → 三级总结(Ollama) → 分类判定(L1总结) → 切分 chunk
→ 构建 L2/L3 大纲节点 → 批量向量化(dense + sparse)→ 写入 Qdrant 四层集合
L2/L3 大纲节点的构建策略见 _build_l2_nodes / _build_l3_nodes。
"""
import uuid
from collections.abc import Callable
from typing import Any
import structlog
from app.config import settings
from app.core.chunker import Chunker
from app.core.classifier import Classifier
from app.core.embeddings import EmbeddingService, create_embedding_service
from app.core.headings import Heading, parse_headings
from app.core.sparse import SparseEncoder
from app.core.summarizer import Summarizer
from app.models.document import ChunkModel, DocumentInput, DocumentSummary, IngestionResult, SummaryLevel
from app.models.knowledge import CategoryResult
from app.services.qdrant import COLLECTION_L2, COLLECTION_L3, QdrantService, SparseVectorTuple
logger = structlog.get_logger()
# IngestionError 阶段标识
STAGE_SUMMARIZE = "summarize"
STAGE_CLASSIFY = "classify"
STAGE_EMBED = "embed"
STAGE_QDRANT = "qdrant"
class IngestionError(Exception):
"""入库失败异常
携带失败阶段(stage)与已产出的总结(summary,如有),
Qdrant 写入失败时总结不丢,上层可按阶段重试。
"""
def __init__(self, stage: str, message: str, summary: DocumentSummary | None = None) -> None:
super().__init__(message)
self.stage = stage
self.summary = summary
def _heading_paths(headings: list[Heading]) -> list[tuple[str, str]]:
"""按文档顺序计算每个标题的 (标题文本, 祖先标题链含自身),链用 " / " 连接"""
paths: list[tuple[str, str]] = []
stack: list[Heading] = []
for heading in headings:
# 遇到同级或更高级标题时弹栈,维护当前祖先链
while stack and stack[-1].level >= heading.level:
stack.pop()
stack.append(heading)
paths.append((heading.title, " / ".join(h.title for h in stack)))
return paths
def _split_l3_blocks(outline: str) -> list[tuple[str, str]]:
"""将内容大纲按 "## " 行分块,返回 (块标题, 块文本) 列表
"## " 行时整块作为一个节点(块标题为空);
首个 "## " 之前的引导内容直接忽略。
"""
stripped = outline.strip()
if not stripped:
return []
blocks: list[tuple[str, list[str]]] = []
for line in stripped.splitlines():
if line.startswith("## "):
blocks.append((line[3:].strip(), [line]))
elif blocks:
blocks[-1][1].append(line)
if not blocks:
return [("", stripped)]
return [(title, "\n".join(lines).strip()) for title, lines in blocks]
class Ingester:
"""文档入库器:编排总结、分类、切分、向量化与 Qdrant 写入全链路"""
def __init__(
self,
summarizer: Summarizer | None = None,
classifier: Classifier | None = None,
chunker: Chunker | None = None,
embedding: EmbeddingService | None = None,
sparse: SparseEncoder | None = None,
qdrant: QdrantService | None = None,
) -> None:
self.summarizer = summarizer or Summarizer()
self.classifier = classifier or Classifier()
self.chunker = chunker or Chunker()
self.embedding = embedding or create_embedding_service()
self.sparse = sparse or SparseEncoder()
self.qdrant = qdrant or QdrantService()
async def ingest(self, doc: DocumentInput, progress_cb: Callable[[str], None] | None = None) -> IngestionResult:
"""执行文档入库
Args:
doc: 文档输入(文本内容 + 元数据)
progress_cb: 可选的阶段进度回调(同步函数),在各阶段边界以
"summarizing" / "classifying" / "embedding" / "writing" 调用
Returns:
IngestionResult: 入库结果
Raises:
IngestionError: 任一阶段失败时抛出,携带 stage 与已产出总结
"""
def _report(stage: str) -> None:
if progress_cb is not None:
progress_cb(stage)
logger.info("开始文档入库", title=doc.title, text_length=len(doc.text))
doc_id = uuid.uuid4().hex
# 1. 三级总结
_report("summarizing")
try:
summary = await self.summarizer.summarize(doc.text, title=doc.title)
except Exception as exc:
logger.error("入库失败:三级总结", stage=STAGE_SUMMARIZE, error=str(exc))
raise IngestionError(STAGE_SUMMARIZE, f"三级总结失败: {exc}") from exc
logger.info("三级总结完成", doc_id=doc_id, level=summary.level.value)
# 2. 分类判定(基于 L1 总结)
_report("classifying")
try:
category = await self.classifier.classify(summary.l1_summary, title=doc.title)
except Exception as exc:
logger.error("入库失败:分类判定", stage=STAGE_CLASSIFY, error=str(exc))
raise IngestionError(STAGE_CLASSIFY, f"分类判定失败: {exc}", summary=summary) from exc
logger.info("分类判定完成", doc_id=doc_id, category=category.main_category, confidence=category.confidence)
# 3. 切分 chunk 并构建 L2/L3 大纲节点((text, section_path) 列表)
chunks = self.chunker.chunk(doc.text, doc_id)
l2_nodes = self._build_l2_nodes(doc.text, summary)
l3_nodes = self._build_l3_nodes(doc.text, summary)
# 4. 批量 embeddingL1 + L2 + L3 + chunks 一次调用,按序切片取向量
texts = [
summary.l1_summary,
*(node_text for node_text, _ in l2_nodes),
*(node_text for node_text, _ in l3_nodes),
*(c.text for c in chunks),
]
try:
_report("embedding")
vectors = await self.embedding.embed(texts)
except Exception as exc:
logger.error("入库失败:向量化", stage=STAGE_EMBED, error=str(exc))
raise IngestionError(STAGE_EMBED, f"向量化失败: {exc}", summary=summary) from exc
l1_vector = vectors[0]
l2_vectors = vectors[1 : 1 + len(l2_nodes)]
l3_vectors = vectors[1 + len(l2_nodes) : 1 + len(l2_nodes) + len(l3_nodes)]
chunk_vectors = vectors[1 + len(l2_nodes) + len(l3_nodes) :]
# 5. sparse 向量(仅 L1 与 chunks 需要)
l1_sparse: SparseVectorTuple | None = None
chunk_sparses: list[SparseVectorTuple | None] = [None] * len(chunks)
if settings.sparse_enabled:
l1_sparse = self.sparse.encode(summary.l1_summary)
chunk_sparses = [self.sparse.encode(c.text) for c in chunks]
# 6. 写入 Qdrant 四层集合
_report("writing")
try:
await self._write_qdrant(
doc_id,
doc,
summary,
category,
chunks,
l2_nodes,
l3_nodes,
l1_vector,
l2_vectors,
l3_vectors,
chunk_vectors,
l1_sparse,
chunk_sparses,
)
except Exception as exc:
logger.error("入库失败:Qdrant 写入", stage=STAGE_QDRANT, doc_id=doc_id, error=str(exc))
raise IngestionError(STAGE_QDRANT, f"Qdrant 写入失败: {exc}", summary=summary) from exc
logger.info(
"文档入库完成",
doc_id=doc_id,
chunks_count=len(chunks),
l2_nodes=len(l2_nodes),
l3_nodes=len(l3_nodes),
category=category.main_category,
)
return IngestionResult(
document_id=doc_id,
summary=summary,
category=category.main_category,
tags=category.tags,
category_confidence=category.confidence,
collection="四层集合",
chunks_count=len(chunks),
)
def _build_l2_nodes(self, text: str, summary: DocumentSummary) -> list[tuple[str, str]]:
"""构建 L2 大纲节点,返回 (text, section_path) 列表
- 有标题结构(L3 级且标题数 >= 2):每个标题一个节点,
text 与 section_path 均为该节点的祖先标题链
- 否则若 l2_outline 非空(LLM 生成的大纲):按非空行拆节点,section_path 为空
- 2.5 级文档(l2_outline 为 None):无 L2 节点
"""
headings = parse_headings(text)
if summary.level == SummaryLevel.L3 and len(headings) >= 2:
return [(path, path) for _, path in _heading_paths(headings)]
if summary.l2_outline:
return [(line.strip(), "") for line in summary.l2_outline.splitlines() if line.strip()]
return []
def _build_l3_nodes(self, text: str, summary: DocumentSummary) -> list[tuple[str, str]]:
"""构建 L3 内容大纲节点,返回 (text, section_path) 列表
"## " 分块(无 "## " 则整块一个节点);
section_path 尽力匹配文档标题链(块标题与文档标题文本精确匹配),匹配不到用 ""
"""
path_by_title: dict[str, str] = {}
for title, path in _heading_paths(parse_headings(text)):
path_by_title.setdefault(title, path)
return [
(block_text, path_by_title.get(block_title, ""))
for block_title, block_text in _split_l3_blocks(summary.l3_content_outline)
]
async def _write_qdrant(
self,
doc_id: str,
doc: DocumentInput,
summary: DocumentSummary,
category: CategoryResult,
chunks: list[ChunkModel],
l2_nodes: list[tuple[str, str]],
l3_nodes: list[tuple[str, str]],
l1_vector: list[float],
l2_vectors: list[list[float]],
l3_vectors: list[list[float]],
chunk_vectors: list[list[float]],
l1_sparse: SparseVectorTuple | None,
chunk_sparses: list[SparseVectorTuple | None],
) -> None:
"""将 L1/L2/L3/chunks 四层数据写入 Qdrant(任一失败向上抛出)"""
await self.qdrant.upsert_l1(
doc_id=doc_id,
title=doc.title,
summary=summary.l1_summary,
category=category.main_category,
tags=category.tags,
dense_vector=l1_vector,
sparse_vector=l1_sparse,
)
# L2/L3 大纲节点(为空时跳过对应集合的 upsert)
for collection, nodes, vectors in (
(COLLECTION_L2, l2_nodes, l2_vectors),
(COLLECTION_L3, l3_nodes, l3_vectors),
):
if not nodes:
continue
await self.qdrant.upsert_nodes(
collection,
[
{
"doc_id": doc_id,
"section_path": section_path,
"text": node_text,
"category": category.main_category,
"tags": category.tags,
"dense_vector": vector,
}
for (node_text, section_path), vector in zip(nodes, vectors, strict=True)
],
)
if chunks:
# chunk dict 额外携带 doc_summary(= L1 总结),检索侧直接取用,不用回查 L1
chunk_dicts: list[dict[str, Any]] = [
{
"doc_id": doc_id,
"chunk_index": chunk.chunk_index,
"text": chunk.text,
"section_path": chunk.section_path,
"title": doc.title,
"category": category.main_category,
"tags": category.tags,
"dense_vector": vector,
"sparse_vector": sparse,
"doc_summary": summary.l1_summary,
}
for chunk, vector, sparse in zip(chunks, chunk_vectors, chunk_sparses, strict=True)
]
await self.qdrant.upsert_chunks(chunk_dicts)
+214
View File
@@ -0,0 +1,214 @@
"""query 解析与分类路由模块
分层 RAG 在线侧第一步:用 Ollama 小模型将用户 query 解析为结构化 JSON
(命中类目+置信度、rewrite 后 query、关键词),再由纯函数做路由决策:
- 高置信且命中类目数 <= 上限 → 按类目过滤检索
- 低置信 / 解析失败 / 命中类目过多 → 全库兜底(不丢召回)
解析(LLM 调用)与决策(纯函数)分离,便于单元测试。
"""
import json
import re
from hashlib import sha256
import structlog
from pydantic import BaseModel, Field, ValidationError
from app.config import settings
from app.models.knowledge import UNCATEGORIZED, TaxonomyCategory
from app.services.ollama import OllamaClient
from app.services.redis import RedisCache, get_cache
logger = structlog.get_logger()
# 从 LLM 输出中提取第一个 {...} JSON 块(贪婪匹配到最后的 },兼容嵌套对象)
_JSON_BLOCK_RE = re.compile(r"\{.*\}", re.DOTALL)
class CategoryHit(BaseModel):
"""query 命中的类目及置信度"""
name: str = Field(description="类目名称,必须来自 taxonomy")
confidence: float = Field(ge=0, le=1, description="命中置信度,范围 [0, 1]")
class ParsedQuery(BaseModel):
"""query 解析结果"""
raw_query: str = Field(description="原始 query")
rewrite: str = Field(description="rewrite 后的 query")
keywords: list[str] = Field(default_factory=list, description="提取的关键词")
categories: list[CategoryHit] = Field(default_factory=list, description="命中类目列表")
parse_failed: bool = Field(default=False, description="LLM 输出解析是否失败")
class RouteDecision(BaseModel):
"""路由决策结果"""
fallback: bool = Field(description="是否走全库兜底")
filter_categories: list[str] | None = Field(default=None, description="过滤类目名列表,None 表示全库不过滤")
reason: str = Field(description="决策原因:parse_failed | low_confidence | too_many_categories | routed")
parsed: ParsedQuery = Field(description="对应的 query 解析结果")
def _extract_json(raw: str) -> dict | None:
"""从 LLM 输出中提取 JSON 对象
先尝试直接解析;失败则用正则提取第一个 {...} 块再解析。
返回 None 表示无法提取出合法的 JSON 对象。
"""
text = raw.strip()
try:
data = json.loads(text)
return data if isinstance(data, dict) else None
except json.JSONDecodeError:
pass
match = _JSON_BLOCK_RE.search(text)
if not match:
return None
try:
data = json.loads(match.group(0))
return data if isinstance(data, dict) else None
except json.JSONDecodeError:
return None
def decide_route(parsed: ParsedQuery, threshold: float, max_categories: int) -> RouteDecision:
"""路由决策(纯函数)
- 解析失败 → 全库兜底(parse_failed
- 无 confidence >= threshold 的类目 → 全库兜底(low_confidence
- 命中类目数 > max_categories → 全库兜底(too_many_categories
- 否则按类目过滤检索,类目按 confidence 降序排列(routed
"""
if parsed.parse_failed:
return RouteDecision(fallback=True, filter_categories=None, reason="parse_failed", parsed=parsed)
hits = [c for c in parsed.categories if c.confidence >= threshold]
if not hits:
return RouteDecision(fallback=True, filter_categories=None, reason="low_confidence", parsed=parsed)
if len(hits) > max_categories:
return RouteDecision(fallback=True, filter_categories=None, reason="too_many_categories", parsed=parsed)
hits.sort(key=lambda c: c.confidence, reverse=True)
return RouteDecision(
fallback=False,
filter_categories=[c.name for c in hits],
reason="routed",
parsed=parsed,
)
class QueryParser:
"""query 解析器:调用 Ollama 小模型将 query 解析为结构化 JSON"""
def __init__(
self, ollama: OllamaClient, taxonomy: list[TaxonomyCategory], cache: RedisCache | None = None
) -> None:
self.ollama = ollama
self.taxonomy = taxonomy
# 解析结果缓存,缺省用全局单例;RedisCache 全操作容错,缓存不可用时退化为无缓存行为
self.cache = cache if cache is not None else get_cache()
# 可作为路由命中类目的名字集合(uncategorized 不可作为路由命中类目)
self._routable_names = {c.name for c in taxonomy if c.name != UNCATEGORIZED}
async def parse(self, query: str) -> ParsedQuery:
"""调用 LLM 将 query 解析为结构化结果
解析失败(输出非 JSON / 必填字段缺失或类型错误)时,
返回 parse_failed=True 的兜底结果(rewrite 为原 query)。
"""
prompt = self._build_prompt(query)
raw = await self.ollama.generate(prompt, json_mode=True)
data = _extract_json(raw)
if data is None:
logger.warning("query 解析失败:LLM 输出非合法 JSON", query=query, output=raw[:200])
return ParsedQuery(raw_query=query, rewrite=query, parse_failed=True)
rewrite = data.get("rewrite")
keywords = data.get("keywords")
categories = data.get("categories")
if not isinstance(rewrite, str) or not isinstance(keywords, list) or not isinstance(categories, list):
logger.warning("query 解析失败:JSON 必填字段缺失或类型错误", query=query, output=raw[:200])
return ParsedQuery(raw_query=query, rewrite=query, parse_failed=True)
return ParsedQuery(
raw_query=query,
rewrite=rewrite,
keywords=[str(k) for k in keywords],
categories=self._validate_categories(categories, query),
)
async def parse_and_route(self, query: str) -> RouteDecision:
"""解析 query 并做路由决策(threshold / max_categories 取自 settings
parse() 的 LLM 解析结果按 query 哈希缓存;路由决策每次现算(纯函数,
阈值取最新 settings),缓存数据损坏时回退为重新走 LLM 解析。
"""
cache_key = f"qparse:{sha256(query.encode()).hexdigest()[:16]}"
parsed = await self._get_cached_parsed(cache_key, query)
if parsed is None:
parsed = await self.parse(query)
await self.cache.set_json(cache_key, parsed.model_dump())
return decide_route(
parsed,
threshold=settings.classify_confidence_threshold,
max_categories=settings.classify_max_categories,
)
async def _get_cached_parsed(self, cache_key: str, query: str) -> ParsedQuery | None:
"""读取缓存的解析结果;未命中或缓存数据无法重建时返回 None"""
cached = await self.cache.get_json(cache_key)
if cached is None:
return None
try:
parsed = ParsedQuery.model_validate(cached)
except ValidationError:
logger.warning("query 解析缓存数据损坏,按未命中处理", query=query, cache_key=cache_key)
return None
logger.info("query 解析缓存命中", query=query, cache_key=cache_key)
return parsed
def _build_prompt(self, query: str) -> str:
"""构造解析 prompt:列出 taxonomy 类目名+描述,要求模型只输出 JSON"""
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"
"1. 判断 query 意图命中以下哪些知识类目,并给出每个类目的置信度(0~1 之间的小数);\n"
"2. 将 query 改写为更适合检索的形式;\n"
"3. 提取 query 的关键词。\n\n"
f"可选类目:\n{category_block}\n\n"
"要求:\n"
"- categories 中的 name 只能从上面的类目名中选择,不要输出其他名称;\n"
"- 若没有明显命中的类目,categories 返回空列表;\n"
"- 只输出 JSON,不要输出任何其他内容。\n\n"
'输出格式:{"categories": [{"name": "类目名", "confidence": 0.0}], '
'"rewrite": "改写后的 query", "keywords": ["关键词"]}\n\n'
f"用户 query{query}"
)
def _validate_categories(self, categories: list, query: str) -> list[CategoryHit]:
"""校验并过滤 LLM 输出的类目条目
丢弃:类目名不在 taxonomy 可路由类目中的条目(含 uncategorized)、
confidence 缺失或越界(不在 [0, 1])的条目。
"""
hits: list[CategoryHit] = []
for item in categories:
if not isinstance(item, dict):
continue
name = item.get("name")
confidence = item.get("confidence")
if not isinstance(name, str) or name not in self._routable_names:
logger.warning("丢弃不在 taxonomy 中的类目", query=query, name=name)
continue
if not isinstance(confidence, (int, float)) or not 0 <= confidence <= 1:
logger.warning("丢弃 confidence 越界的类目", query=query, name=name, confidence=confidence)
continue
hits.append(CategoryHit(name=name, confidence=float(confidence)))
return hits
+24
View File
@@ -0,0 +1,24 @@
"""检索结果重排:RRF 融合与最终截断"""
from qdrant_client import models
def rrf_fuse(result_lists: list[list[models.ScoredPoint]], k: int = 60) -> list[models.ScoredPoint]:
"""标准 RRFReciprocal Rank Fusion)融合
融合分 = Σ 1/(k + rank)rank 从 1 开始),按 point id 去重合并,
返回按融合分降序的列表,score 字段写回融合分。空输入返回 []。
"""
scores: dict[str | int, float] = {}
points: dict[str | int, models.ScoredPoint] = {}
for results in result_lists:
for rank, point in enumerate(results, start=1):
scores[point.id] = scores.get(point.id, 0.0) + 1.0 / (k + rank)
points.setdefault(point.id, point)
ordered = sorted(points, key=lambda pid: scores[pid], reverse=True)
return [points[pid].model_copy(update={"score": scores[pid]}) for pid in ordered]
def finalize(points: list[models.ScoredPoint], final_k: int) -> list[models.ScoredPoint]:
"""截断为最终返回的 top final_k"""
return points[:final_k]
+169
View File
@@ -0,0 +1,169 @@
"""分层检索引擎
L1(文档总结)→ L2(章节大纲)→ L3(小节定位)→ chunks(原文)逐层收窄:
- L1 无候选文档:直接全库 chunk 兜底,fallback=True
- L2 无命中:全部候选文档回退到 L3 的 doc 级查询(b 路)
- 2.5 级文档无 L2 节点,天然落入 L3 b 路(仅按 doc 过滤)
- L3 两路(aL2 命中文档按 section 过滤;b:其余文档仅按 doc 过滤)RRF 融合;
L3 无命中时 chunk 层回退为 L1 候选文档级检索
"""
from collections.abc import Iterable
import structlog
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.ranker import finalize, rrf_fuse
from app.core.sparse import SparseEncoder
from app.models.knowledge import load_taxonomy
from app.models.search import SearchHit, SearchRequest, SearchResponse
from app.services.ollama import OllamaClient
from app.services.qdrant import (
COLLECTION_CHUNKS,
COLLECTION_L1,
COLLECTION_L2,
COLLECTION_L3,
SPARSE_COLLECTIONS,
QdrantService,
)
logger = structlog.get_logger()
def _unique(values: Iterable[str | None]) -> list[str]:
"""去重保序并丢弃空值"""
return list(dict.fromkeys(v for v in values if v))
class Retriever:
"""分层检索引擎,依赖均可注入(默认自建,便于测试替换)"""
def __init__(
self,
qdrant: QdrantService | None = None,
query_parser: QueryParser | None = None,
embedding: EmbeddingService | None = None,
sparse_encoder: SparseEncoder | None = None,
) -> None:
self.qdrant = qdrant or QdrantService()
self.query_parser = query_parser or QueryParser(
ollama=OllamaClient(),
taxonomy=load_taxonomy(settings.taxonomy_path),
)
self.embedding = embedding or create_embedding_service()
self.sparse_encoder = sparse_encoder or SparseEncoder()
async def search(self, request: SearchRequest) -> SearchResponse:
"""分层检索主流程"""
route = await self.query_parser.parse_and_route(request.query)
query_text = route.parsed.rewrite or request.query
dense = (await self.embedding.embed([query_text]))[0]
sparse = self.sparse_encoder.encode(query_text) if settings.sparse_enabled else None
# 路由兜底时不做类目过滤,避免丢召回
categories = None if route.fallback else route.filter_categories
# L1:文档级检索,产出候选文档
l1_filter = QdrantService.build_filter(categories=categories)
l1_hits = await self._search_collection(COLLECTION_L1, dense, sparse, settings.l1_doc_top_n, l1_filter)
logger.info("L1 检索完成", hits=len(l1_hits), categories=categories)
if not l1_hits:
# 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,
)
doc_ids = _unique((p.payload or {}).get("doc_id") for p in l1_hits)
# L2:候选文档内检索章节大纲
l2_filter = QdrantService.build_filter(categories=categories, doc_ids=doc_ids)
l2_hits = await self._search_collection(
COLLECTION_L2, dense, sparse, settings.l2_section_top_n * len(doc_ids), l2_filter
)
logger.info("L2 检索完成", hits=len(l2_hits))
l2_doc_ids = _unique((p.payload or {}).get("doc_id") for p in l2_hits)
l2_section_paths = _unique((p.payload or {}).get("section_path") for p in l2_hits)
# 无 L2 命中的文档(2.5 级文档或该 doc 的 L2 未命中)走 L3 b 路
remaining_doc_ids = [d for d in doc_ids if d not in l2_doc_ids]
# L3:两路查询后 RRF 融合
l3_lists: list[list[models.ScoredPoint]] = []
if l2_doc_ids:
l3_filter_a = QdrantService.build_filter(
categories=categories, doc_ids=l2_doc_ids, section_paths=l2_section_paths
)
l3_lists.append(await self._search_collection(COLLECTION_L3, dense, sparse, settings.l3_top_n, l3_filter_a))
if remaining_doc_ids:
l3_filter_b = QdrantService.build_filter(categories=categories, doc_ids=remaining_doc_ids)
l3_lists.append(await self._search_collection(COLLECTION_L3, dense, sparse, settings.l3_top_n, l3_filter_b))
l3_hits = rrf_fuse(l3_lists)
logger.info("L3 检索完成", hits=len(l3_hits))
# chunk 层:L3 有命中按 section 收窄;无命中回退为 L1 候选文档级检索
if l3_hits:
l3_doc_ids = _unique((p.payload or {}).get("doc_id") for p in l3_hits)
# L3 命中 section_path 全为空(如 2.5 级文档)时仅按 doc 过滤
l3_section_paths = _unique((p.payload or {}).get("section_path") for p in l3_hits)
chunk_filter = QdrantService.build_filter(
categories=categories, doc_ids=l3_doc_ids, section_paths=l3_section_paths
)
else:
logger.info("L3 无命中,回退到 L1 候选文档级 chunk 检索")
chunk_filter = QdrantService.build_filter(categories=categories, doc_ids=doc_ids)
chunk_hits = await self._search_collection(
COLLECTION_CHUNKS, dense, sparse, settings.retrieval_top_k, chunk_filter
)
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,
)
async def _search_collection(
self,
collection: str,
dense: list[float],
sparse: tuple[list[int], list[float]] | None,
limit: int,
query_filter: models.Filter | None,
) -> list[models.ScoredPoint]:
"""按集合是否支持 sparse 选择 hybrid 或 dense 检索"""
if sparse is not None and collection in SPARSE_COLLECTIONS:
return await self.qdrant.search_hybrid(collection, dense, sparse, limit, query_filter=query_filter)
return await self.qdrant.search_dense(collection, dense, limit, query_filter=query_filter)
@staticmethod
def _final_k(request: SearchRequest) -> int:
"""最终返回数:请求指定优先,否则用配置默认值"""
return request.top_k or settings.retrieval_final_k
@staticmethod
def _to_hits(points: list[models.ScoredPoint]) -> list[SearchHit]:
"""chunk 点组装为 SearchHit,字段取自 chunk payloaddoc_summary 仅上下文标注)"""
hits: list[SearchHit] = []
for p in points:
payload = p.payload or {}
hits.append(
SearchHit(
text=payload.get("text", ""),
doc_id=payload.get("doc_id", ""),
title=payload.get("title", ""),
section_path=payload.get("section_path") or "",
score=p.score,
doc_summary=payload.get("doc_summary") or "",
)
)
return hits
+68
View File
@@ -0,0 +1,68 @@
"""BM25 轻量近似稀疏向量编码器
零第三方依赖的稀疏编码实现:无全局 IDF 的 BM25 近似(仅用词频 tf 加权),
输出 Qdrant SparseVector 所需的 indices/values 格式。
后续可替换为 SPLADE / BM42 等更强的稀疏编码器。
"""
import hashlib
import math
import re
# 稀疏向量维度(哈希空间大小)
SPARSE_DIM = 2**18
# 连续 CJK 字符段(一-鿿)
_CJK_RE = re.compile(r"[一-鿿]+")
# CJK 段或英文/数字连续段
_TOKEN_RE = re.compile(r"[一-鿿]+|[a-zA-Z0-9]+")
def _tokenize(text: str) -> list[str]:
"""分词(纯正则实现)
- 连续 CJK 字符段:长度 1 时保留 unigram,长度 >= 2 时生成字符 bigram
- 英文/数字连续段:小写化后作为整词
"""
tokens: list[str] = []
for match in _TOKEN_RE.finditer(text):
seg = match.group()
if _CJK_RE.fullmatch(seg):
if len(seg) == 1:
tokens.append(seg)
else:
tokens.extend(seg[i : i + 2] for i in range(len(seg) - 1))
else:
tokens.append(seg.lower())
return tokens
def _hash(token: str) -> int:
"""将词哈希到 [0, SPARSE_DIM) 的索引空间"""
digest = hashlib.blake2b(token.encode("utf-8"), digest_size=8).digest()
return int.from_bytes(digest, "big") % SPARSE_DIM
class SparseEncoder:
"""稀疏向量编码器,输出 Qdrant SparseVector 的 indices/values"""
def encode(self, text: str) -> tuple[list[int], list[float]]:
"""编码单条文本
权重 = 1 + log(tf),词经哈希到 [0, SPARSE_DIM),哈希冲突时权重累加。
Returns:
(indices, values)indices 升序且无重复,与 values 等长
"""
# 按哈希索引累计词频,天然处理哈希冲突
tf: dict[int, float] = {}
for token in _tokenize(text):
idx = _hash(token)
tf[idx] = tf.get(idx, 0.0) + 1.0
indices = sorted(tf)
values = [1.0 + math.log(tf[idx]) for idx in indices]
return indices, values
def encode_batch(self, texts: list[str]) -> list[tuple[list[int], list[float]]]:
"""批量编码,与逐条 encode 结果一致"""
return [self.encode(text) for text in texts]
+138
View File
@@ -0,0 +1,138 @@
"""文档三级总结模块
通过 Ollama 本地小模型对文档进行分级总结:
- L1: 总结(一句话高度概括)
- L2: 大纲(主要章节和关键主题)
- L3: 内容大纲(每个章节的详细内容摘要)
当文档内容不足以支撑三级总结时,自动降级为 2.5 级(L1 + L2.5 内容大纲)。
"""
import structlog
from app.core.headings import parse_headings, render_outline
from app.models.document import DocumentSummary, SummaryLevel
from app.services.ollama import OllamaClient
logger = structlog.get_logger()
# 文本长度阈值(字符数),低于此值触发 2.5 级回退
MIN_TEXT_LENGTH_FOR_L3 = 500
class Summarizer:
"""文档三级总结器"""
def __init__(self, ollama: OllamaClient | None = None) -> None:
self.ollama = ollama or OllamaClient()
async def summarize(self, text: str, *, title: str = "") -> DocumentSummary:
"""对文档文本进行三级总结
Args:
text: 文档纯文本内容
title: 文档标题(可选,辅助总结)
Returns:
DocumentSummary: 包含各级总结的结果
"""
text_length = len(text.strip())
# 生成 L1 总结
l1_summary = await self._generate_l1(text, title)
logger.info("L1 总结完成", title=title, l1=l1_summary[:100])
# 判断是否需要 2.5 级回退
use_fallback = text_length < MIN_TEXT_LENGTH_FOR_L3
if use_fallback:
logger.info("文档内容不足,使用 2.5 级回退", text_length=text_length)
# L2.5: 直接生成内容大纲(跳过大纲层)
l2_half = await self._generate_l2_half(text, title, l1_summary)
return DocumentSummary(
l1_summary=l1_summary,
l2_outline=None,
l3_content_outline=l2_half,
level=SummaryLevel.L2_HALF,
)
# 生成 L2 大纲:优先使用文档原生标题树(结构导航,不调 LLM),
# 无结构文本(标题数 < 2)回退 LLM 生成
headings = parse_headings(text)
if len(headings) >= 2:
l2_outline = render_outline(headings)
logger.info("L2 大纲完成", title=title, outline_source="headings", heading_count=len(headings))
else:
l2_outline = await self._generate_l2(text, title, l1_summary)
logger.info("L2 大纲完成", title=title, outline_source="llm")
# 生成 L3 内容大纲
l3_content_outline = await self._generate_l3(text, title, l1_summary, l2_outline)
logger.info("L3 内容大纲完成", title=title)
return DocumentSummary(
l1_summary=l1_summary,
l2_outline=l2_outline,
l3_content_outline=l3_content_outline,
level=SummaryLevel.L3,
)
async def _generate_l1(self, text: str, title: str) -> str:
"""生成 L1 总结:一句话高度概括"""
prompt = (
"请用一句话对以下文档内容进行高度概括,要求简洁精炼,"
"突出文档的核心主题和关键信息。\n\n"
)
if title:
prompt += f"文档标题:{title}\n\n"
prompt += f"文档内容:\n{text}"
return await self.ollama.generate(prompt)
async def _generate_l2(self, text: str, title: str, l1_summary: str) -> str:
"""生成 L2 大纲:主要章节和关键主题"""
prompt = (
"请提取以下文档的主要章节结构和关键主题,"
"以大纲形式呈现。每个主题用一行表示,"
"格式为:序号. 主题名称\n\n"
f"文档总结:{l1_summary}\n\n"
)
if title:
prompt += f"文档标题:{title}\n\n"
prompt += f"文档内容:\n{text}"
return await self.ollama.generate(prompt)
async def _generate_l3(
self, text: str, title: str, l1_summary: str, l2_outline: str
) -> str:
"""生成 L3 内容大纲:每个章节的详细内容摘要"""
prompt = (
"请对以下文档的每个章节/主题进行详细的内容摘要,"
"格式为:\n"
"## 章节名称\n"
"详细摘要内容...\n\n"
f"文档总结:{l1_summary}\n\n"
f"文档大纲:\n{l2_outline}\n\n"
)
if title:
prompt += f"文档标题:{title}\n\n"
prompt += f"文档内容:\n{text}"
return await self.ollama.generate(prompt)
async def _generate_l2_half(self, text: str, title: str, l1_summary: str) -> str:
"""生成 L2.5 内容大纲(2.5 级回退)
跳过大纲层,直接对短文本生成详细摘要。
"""
prompt = (
"请对以下文档内容进行详细摘要,"
"涵盖所有关键信息点。以要点形式呈现。\n\n"
f"文档总结:{l1_summary}\n\n"
)
if title:
prompt += f"文档标题:{title}\n\n"
prompt += f"文档内容:\n{text}"
return await self.ollama.generate(prompt)
+79
View File
@@ -0,0 +1,79 @@
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from pathlib import Path
import structlog
from fastapi import FastAPI, Request
from fastapi.exceptions import RequestValidationError
from fastapi.responses import FileResponse, JSONResponse
from app.api.response import ApiError, error
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.services.qdrant import QdrantService
logger = structlog.get_logger()
@asynccontextmanager
async def lifespan(_: FastAPI) -> AsyncIterator[None]:
"""应用生命周期:启动时初始化 Qdrant 集合
初始化失败仅记录日志、不阻止启动(本地开发可能无 Qdrant)。
"""
try:
await QdrantService().ensure_collections()
logger.info("Qdrant 集合初始化完成")
except Exception:
logger.error("Qdrant 集合初始化失败,跳过初始化继续启动")
yield
app = FastAPI(
title=settings.app_name,
description="AI Agent 分层信息检索服务",
version="0.1.0",
lifespan=lifespan,
)
app.include_router(search_router)
app.include_router(document_router)
app.include_router(knowledge_router)
@app.exception_handler(ApiError)
async def api_error_handler(request: Request, exc: ApiError) -> JSONResponse:
"""业务异常 → 统一错误响应(code 取自异常)"""
return JSONResponse(content=error(exc.code, exc.message))
@app.exception_handler(RequestValidationError)
async def validation_error_handler(request: Request, exc: RequestValidationError) -> JSONResponse:
"""请求参数校验失败 → code 1001"""
errors = exc.errors()
message = f"请求参数校验失败: {errors[0]['msg']}" if errors else "请求参数校验失败"
return JSONResponse(content=error(1001, message))
@app.exception_handler(Exception)
async def unhandled_error_handler(request: Request, exc: Exception) -> JSONResponse:
"""未捕获异常 → code 2000"""
logger.exception("未捕获异常", path=request.url.path)
return JSONResponse(content=error(2000, "服务器内部错误"))
@app.get("/api/v1/health")
async def health() -> dict[str, str]:
"""健康检查"""
return {"status": "ok"}
_ADMIN_HTML = Path(__file__).resolve().parent / "static" / "admin.html"
@app.get("/admin", include_in_schema=False)
async def admin_page() -> FileResponse:
"""管理后台单页(单文件静态 HTML,零外部依赖)"""
return FileResponse(_ADMIN_HTML, media_type="text/html")
View File
+51
View File
@@ -0,0 +1,51 @@
"""文档和总结相关的数据模型"""
from enum import StrEnum
from pydantic import BaseModel, Field
class SummaryLevel(StrEnum):
"""总结层级"""
L3 = "L3" # 完整三级:总结 → 大纲 → 内容大纲
L2_HALF = "L2.5" # 2.5 级回退:总结 → 内容大纲(跳过大纲)
class DocumentSummary(BaseModel):
"""文档三级总结结果"""
l1_summary: str = Field(description="L1 总结:一句话高度概括")
l2_outline: str | None = Field(default=None, description="L2 大纲:主要章节和关键主题")
l3_content_outline: str = Field(description="L3/L2.5 内容大纲:详细内容摘要")
level: SummaryLevel = Field(description="实际使用的总结层级")
class DocumentInput(BaseModel):
"""文档入库输入"""
text: str = Field(description="文档纯文本内容")
title: str = Field(default="", description="文档标题")
source: str = Field(default="", description="来源标识(文件路径/URL等)")
metadata: dict[str, str] = Field(default_factory=dict, description="附加元数据")
class ChunkModel(BaseModel):
"""文档分块结果"""
doc_id: str = Field(description="所属文档 ID")
chunk_index: int = Field(description="chunk 在文档内的序号")
text: str = Field(description="chunk 文本内容")
section_path: str = Field(default="", description="chunk 所在的章节路径")
class IngestionResult(BaseModel):
"""文档入库结果"""
document_id: str = Field(description="写入后的文档 ID")
summary: DocumentSummary = Field(description="三级总结结果")
category: str = Field(description="分类标签(主类目)")
collection: str = Field(description="写入的 Qdrant 集合名")
chunks_count: int = Field(default=0, description="写入的 chunk 数量")
tags: list[str] = Field(default_factory=list, description="附加分类标签")
category_confidence: float = Field(default=0.0, description="主类目分类置信度")
+68
View File
@@ -0,0 +1,68 @@
"""知识分类(taxonomy)相关的数据模型与加载逻辑"""
import json
from pathlib import Path
import structlog
from pydantic import BaseModel, Field
logger = structlog.get_logger()
# 未分类常量:分类置信度不足或无法归类时使用
UNCATEGORIZED = "uncategorized"
class TaxonomyCategory(BaseModel):
"""知识分类类目定义"""
name: str = Field(description="类目名称,全局唯一")
description: str = Field(default="", description="类目描述,用于辅助分类判断")
class CategoryResult(BaseModel):
"""文档/查询的分类结果"""
main_category: str = Field(description="主类目名称")
tags: list[str] = Field(default_factory=list, description="附加标签列表")
confidence: float = Field(ge=0, le=1, description="分类置信度,范围 [0, 1]")
def _default_taxonomy() -> list[TaxonomyCategory]:
"""内置默认类目集(通用企业知识库场景)"""
return [
TaxonomyCategory(name="技术文档", description="架构设计、API 文档、开发规范、运维手册等技术资料"),
TaxonomyCategory(name="产品手册", description="产品功能介绍、使用说明、版本发布说明"),
TaxonomyCategory(name="运营规范", description="运营流程、活动方案、内容规范、客服话术"),
TaxonomyCategory(name="财务行政", description="财务制度、报销流程、行政通知、办公管理"),
TaxonomyCategory(name="市场资料", description="市场分析、竞品调研、营销素材、品牌规范"),
TaxonomyCategory(name="人事制度", description="招聘、考勤、绩效、培训、员工手册等 HR 制度"),
TaxonomyCategory(name="法律法规", description="合同模板、合规要求、法律条文、知识产权"),
TaxonomyCategory(name=UNCATEGORIZED, description="无法归入其他类目的文档"),
]
def load_taxonomy(path: str = "") -> list[TaxonomyCategory]:
"""加载 taxonomy 类目集
path 为空时使用内置默认类目集;非空时从 JSON 文件加载,
文件格式为 [{"name": ..., "description": ...}]。
校验类目 name 唯一;若缺少 uncategorized 类目则自动追加。
"""
if not path:
return _default_taxonomy()
raw = json.loads(Path(path).read_text(encoding="utf-8"))
categories = [TaxonomyCategory.model_validate(item) for item in raw]
# 校验 name 唯一
names = [c.name for c in categories]
if len(names) != len(set(names)):
duplicates = sorted({n for n in names if names.count(n) > 1})
raise ValueError(f"taxonomy 类目 name 重复: {duplicates}")
# 必含 uncategorized,缺失则自动追加
if UNCATEGORIZED not in names:
logger.warning("taxonomy 缺少 uncategorized 类目,已自动追加", path=path)
categories.append(TaxonomyCategory(name=UNCATEGORIZED, description="无法归入其他类目的文档"))
return categories
+30
View File
@@ -0,0 +1,30 @@
"""检索请求与响应的数据模型"""
from pydantic import BaseModel, Field
class SearchRequest(BaseModel):
"""检索请求"""
query: str = Field(description="查询文本")
top_k: int | None = Field(default=None, description="返回结果数,为空时使用 settings.retrieval_final_k")
class SearchHit(BaseModel):
"""单条检索命中结果"""
text: str = Field(description="原文 chunk 内容")
doc_id: str = Field(description="所属文档 ID")
title: str = Field(default="", description="文档标题")
section_path: str = Field(default="", description="chunk 所在的章节路径")
score: float = Field(description="相关性得分")
doc_summary: str = Field(default="", description="L1 文档总结,仅用于上下文标注")
class SearchResponse(BaseModel):
"""检索响应"""
query: str = Field(description="原始查询文本")
hits: list[SearchHit] = Field(default_factory=list, description="命中结果列表")
routed_categories: list[str] = Field(default_factory=list, description="query 路由命中的类目")
fallback: bool = Field(default=False, description="是否走了全库兜底路径")
View File
+62
View File
@@ -0,0 +1,62 @@
"""Ollama 客户端
通过 Ollama HTTP API 调用本地模型进行文本生成。
"""
import httpx
import structlog
from app.config import settings
logger = structlog.get_logger()
class OllamaClient:
"""Ollama HTTP API 客户端"""
def __init__(
self,
base_url: str | None = None,
model: str | None = None,
timeout: float = 120.0,
) -> None:
self.base_url = (base_url or settings.ollama_base_url).rstrip("/")
self.model = model or settings.ollama_model
self.timeout = timeout
async def generate(self, prompt: str, json_mode: bool = False) -> str:
"""调用 Ollama 生成文本
Args:
prompt: 输入提示词
json_mode: 为 True 时启用 Ollama 原生 JSON 约束输出(payload 加 format: json
Returns:
生成的文本内容
"""
url = f"{self.base_url}/api/generate"
payload = {
"model": self.model,
"prompt": prompt,
"stream": False,
}
if json_mode:
payload["format"] = "json"
async with httpx.AsyncClient(timeout=self.timeout) as client:
resp = await client.post(url, json=payload)
resp.raise_for_status()
data = resp.json()
result = data.get("response", "")
logger.debug("Ollama 生成完成", model=self.model, output_length=len(result))
return result
async def is_available(self) -> bool:
"""检查 Ollama 服务是否可用"""
try:
async with httpx.AsyncClient(timeout=5.0) as client:
resp = await client.get(f"{self.base_url}/api/tags")
return resp.status_code == 200
except httpx.HTTPError:
return False
+351
View File
@@ -0,0 +1,351 @@
"""Qdrant 向量数据库服务封装
分层摘要索引 RAG 的存储层:
- doc_l1:文档级总结(dense + sparse
- doc_l2 / doc_l3:大纲节点(dense
- chunks:原文分块(dense + sparse
提供集合初始化(幂等)、按层 upsert、dense / hybridRRF 融合)检索接口。
"""
import uuid
from typing import Any
import structlog
from qdrant_client import AsyncQdrantClient, models
from app.config import settings
logger = structlog.get_logger()
# 集合名
COLLECTION_L1 = "doc_l1"
COLLECTION_L2 = "doc_l2"
COLLECTION_L3 = "doc_l3"
COLLECTION_CHUNKS = "chunks"
ALL_COLLECTIONS = (COLLECTION_L1, COLLECTION_L2, COLLECTION_L3, COLLECTION_CHUNKS)
# 命名向量
VECTOR_DENSE = "dense"
VECTOR_SPARSE = "sparse"
# 需要 sparse 向量的集合
SPARSE_COLLECTIONS = (COLLECTION_L1, COLLECTION_CHUNKS)
# 需要建立 KEYWORD payload 索引的字段
PAYLOAD_INDEX_FIELDS = ("doc_id", "category", "tags", "section_path")
# upsert_nodes 集合名 -> point id 层级前缀
_NODE_ID_PREFIX = {COLLECTION_L2: "l2", COLLECTION_L3: "l3"}
# 稀疏向量统一表示:(indices, values)
SparseVectorTuple = tuple[list[int], list[float]]
def _point_id(key: str) -> str:
"""由确定性 key 生成 UUID5 point id,保证重复入库幂等覆盖"""
return str(uuid.uuid5(uuid.NAMESPACE_URL, key))
class QdrantService:
"""Qdrant 异步客户端封装"""
def __init__(self, client: AsyncQdrantClient | None = None) -> None:
# 允许注入自定义 client(测试可用 location=":memory:" 的本地模式)
self._client = client or AsyncQdrantClient(host=settings.qdrant_host, port=settings.qdrant_port)
@property
def client(self) -> AsyncQdrantClient:
return self._client
async def ensure_collections(self) -> None:
"""初始化 4 个集合与 payload 索引(幂等,已存在则跳过)"""
existing = await self._client.get_collections()
existing_names = {c.name for c in existing.collections}
for collection in ALL_COLLECTIONS:
if collection in existing_names:
continue
sparse_config = None
if settings.sparse_enabled and collection in SPARSE_COLLECTIONS:
sparse_config = {VECTOR_SPARSE: models.SparseVectorParams(modifier=models.Modifier.IDF)}
dense_params = models.VectorParams(size=settings.embedding_dimension, distance=models.Distance.COSINE)
await self._client.create_collection(
collection_name=collection,
vectors_config={VECTOR_DENSE: dense_params},
sparse_vectors_config=sparse_config,
)
logger.info("创建 Qdrant 集合", collection=collection, sparse=sparse_config is not None)
# payload 索引:先查已有索引,缺失才创建,保证幂等
for collection in ALL_COLLECTIONS:
info = await self._client.get_collection(collection)
indexed = set(info.payload_schema.keys()) if info.payload_schema else set()
for field in PAYLOAD_INDEX_FIELDS:
if field in indexed:
continue
await self._client.create_payload_index(
collection_name=collection,
field_name=field,
field_schema=models.PayloadSchemaType.KEYWORD,
)
logger.debug("payload 索引就绪", collection=collection)
# ---------- upsert ----------
async def upsert_l1(
self,
doc_id: str,
title: str,
summary: str,
category: str,
tags: list[str],
dense_vector: list[float],
sparse_vector: SparseVectorTuple | None = None,
) -> None:
"""写入 L1 文档总结,payload 含 doc_id/title/category/tags/text(=summary)"""
vector: dict[str, Any] = {VECTOR_DENSE: dense_vector}
if sparse_vector is not None:
vector[VECTOR_SPARSE] = models.SparseVector(indices=sparse_vector[0], values=sparse_vector[1])
point = models.PointStruct(
id=_point_id(f"{doc_id}:l1"),
vector=vector,
payload={"doc_id": doc_id, "title": title, "category": category, "tags": tags, "text": summary},
)
await self._client.upsert(collection_name=COLLECTION_L1, points=[point])
async def upsert_nodes(self, collection: str, nodes: list[dict[str, Any]]) -> None:
"""批量写入 L2/L3 大纲节点
每个 node 含 doc_id/section_path/text/category/tags/dense_vector。
"""
if collection not in _NODE_ID_PREFIX:
raise ValueError(f"upsert_nodes 仅支持 {COLLECTION_L2}/{COLLECTION_L3},收到: {collection}")
prefix = _NODE_ID_PREFIX[collection]
points = [
models.PointStruct(
id=_point_id(f"{node['doc_id']}:{prefix}:{i}"),
vector={VECTOR_DENSE: node["dense_vector"]},
payload={
"doc_id": node["doc_id"],
"section_path": node["section_path"],
"text": node["text"],
"category": node["category"],
"tags": node["tags"],
},
)
for i, node in enumerate(nodes)
]
if points:
await self._client.upsert(collection_name=collection, points=points)
async def upsert_chunks(self, chunks: list[dict[str, Any]]) -> None:
"""批量写入原文 chunk
每个 chunk 含 doc_id/chunk_index/text/section_path/title/category/tags/dense_vector
sparse_vector 可选((indices, values) 形式)。
"""
points = []
for chunk in chunks:
vector: dict[str, Any] = {VECTOR_DENSE: chunk["dense_vector"]}
sparse = chunk.get("sparse_vector")
if sparse is not None:
vector[VECTOR_SPARSE] = models.SparseVector(indices=sparse[0], values=sparse[1])
points.append(
models.PointStruct(
id=_point_id(f"{chunk['doc_id']}:chunk:{chunk['chunk_index']}"),
vector=vector,
payload={
"doc_id": chunk["doc_id"],
"chunk_index": chunk["chunk_index"],
"text": chunk["text"],
"section_path": chunk["section_path"],
"title": chunk["title"],
"category": chunk["category"],
"tags": chunk["tags"],
# 所属文档 L1 总结,仅作检索结果上下文标注
"doc_summary": chunk.get("doc_summary", ""),
},
)
)
if points:
await self._client.upsert(collection_name=COLLECTION_CHUNKS, points=points)
# ---------- 查询 ----------
async def search_dense(
self,
collection: str,
vector: list[float],
limit: int,
query_filter: models.Filter | None = None,
) -> list[models.ScoredPoint]:
"""dense 命名向量检索"""
resp = await self._client.query_points(
collection_name=collection,
query=vector,
using=VECTOR_DENSE,
limit=limit,
query_filter=query_filter,
)
return resp.points
async def search_hybrid(
self,
collection: str,
dense_vector: list[float],
sparse: SparseVectorTuple,
limit: int,
query_filter: models.Filter | None = None,
) -> list[models.ScoredPoint]:
"""dense + sparse 两路 prefetch,服务端 RRF 融合"""
resp = await self._client.query_points(
collection_name=collection,
prefetch=[
models.Prefetch(query=dense_vector, using=VECTOR_DENSE, limit=limit, filter=query_filter),
models.Prefetch(
query=models.SparseVector(indices=sparse[0], values=sparse[1]),
using=VECTOR_SPARSE,
limit=limit,
filter=query_filter,
),
],
query=models.FusionQuery(fusion=models.Fusion.RRF),
limit=limit,
query_filter=query_filter,
)
return resp.points
# ---------- 管理操作 ----------
async def count(self, collection: str) -> int:
"""精确统计集合中点的总数"""
result = await self._client.count(collection_name=collection, exact=True)
return result.count
async def scroll_l1(self, limit: int = 20, offset: str | None = None) -> tuple[list[dict[str, Any]], str | None]:
"""分页浏览 L1 文档列表(不取向量)
返回 (items, next_offset)item 含 doc_id/title/category/tags/summary(=payload text)
next_offset 为下一页游标,无更多数据时为 None。
"""
records, next_offset = await self._client.scroll(
collection_name=COLLECTION_L1,
limit=limit,
offset=offset,
with_payload=["doc_id", "title", "category", "tags", "text"],
with_vectors=False,
)
items = []
for record in records:
payload = record.payload or {}
items.append(
{
"doc_id": payload.get("doc_id", ""),
"title": payload.get("title", ""),
"category": payload.get("category", ""),
"tags": payload.get("tags", []),
"summary": payload.get("text", ""),
}
)
return items, str(next_offset) if next_offset is not None else None
async def get_doc_detail(self, doc_id: str) -> dict[str, Any] | None:
"""获取文档详情:L1 记录 + L2/L3 全部节点 + chunks 数量,文档不存在返回 None"""
doc_filter = self.build_filter(doc_ids=[doc_id])
l1_records, _ = await self._client.scroll(
collection_name=COLLECTION_L1,
scroll_filter=doc_filter,
limit=1,
with_payload=True,
with_vectors=False,
)
if not l1_records:
return None
l2_nodes = await self._scroll_payloads(COLLECTION_L2, doc_filter)
l3_nodes = await self._scroll_payloads(COLLECTION_L3, doc_filter)
chunks_count = await self._client.count(
collection_name=COLLECTION_CHUNKS,
count_filter=doc_filter,
exact=True,
)
return {
"l1": l1_records[0].payload or {},
"l2_nodes": l2_nodes,
"l3_nodes": l3_nodes,
"chunks_count": chunks_count.count,
}
async def delete_by_doc_id(self, doc_id: str) -> dict[str, int]:
"""删除四层集合中该 doc_id 的所有点,返回各集合删除数量(不存在的 doc_id 全 0,幂等)"""
doc_filter = self.build_filter(doc_ids=[doc_id])
deleted: dict[str, int] = {}
for collection in ALL_COLLECTIONS:
# 先记录删除前数量,再按过滤器删除
before = await self._client.count(collection_name=collection, count_filter=doc_filter, exact=True)
await self._client.delete(
collection_name=collection,
points_selector=models.FilterSelector(filter=doc_filter),
)
deleted[collection] = before.count
logger.info("按 doc_id 删除文档数据", doc_id=doc_id, deleted=deleted)
return deleted
async def _scroll_payloads(
self,
collection: str,
scroll_filter: models.Filter | None,
page_size: int = 256,
) -> list[dict[str, Any]]:
"""循环 scroll 取出所有匹配点的 payload(含 section_path/text 等全字段)"""
payloads: list[dict[str, Any]] = []
offset = None
while True:
records, next_offset = await self._client.scroll(
collection_name=collection,
scroll_filter=scroll_filter,
limit=page_size,
offset=offset,
with_payload=True,
with_vectors=False,
)
payloads.extend(record.payload or {} for record in records)
if next_offset is None:
return payloads
offset = next_offset
# ---------- 过滤构造 ----------
@staticmethod
def build_filter(
categories: list[str] | None = None,
doc_ids: list[str] | None = None,
section_paths: list[str] | None = None,
) -> models.Filter | None:
"""构造 payload 过滤器
- categories:主类硬过滤 + 多标签软召回(category 或 tags 命中其一即召回,min_should=1
- doc_ids / section_pathsmust 条件 MatchAny
- 全空返回 None
"""
must: list[models.FieldCondition] = []
if doc_ids:
must.append(models.FieldCondition(key="doc_id", match=models.MatchAny(any=doc_ids)))
if section_paths:
must.append(models.FieldCondition(key="section_path", match=models.MatchAny(any=section_paths)))
min_should = None
if categories:
min_should = models.MinShould(
conditions=[
models.FieldCondition(key="category", match=models.MatchAny(any=categories)),
models.FieldCondition(key="tags", match=models.MatchAny(any=categories)),
],
min_count=1,
)
if not must and min_should is None:
return None
return models.Filter(must=must or None, min_should=min_should)
+75
View File
@@ -0,0 +1,75 @@
"""Redis 缓存服务
为检索结果与 query 解析结果提供 JSON 缓存。所有操作容错:
Redis 不可用或缓存数据异常时降级为未命中 / 写入失败,绝不影响主流程。
"""
import json
from typing import Any
import structlog
from redis import asyncio as redis_async
from app.config import settings
logger = structlog.get_logger()
class RedisCache:
"""Redis JSON 缓存客户端(懒连接,全操作容错)"""
def __init__(self) -> None:
self._client: redis_async.Redis | None = None
def _get_client(self) -> redis_async.Redis:
"""懒创建 Redis 客户端(TCP 连接在首次执行命令时才建立)"""
if self._client is None:
self._client = redis_async.from_url(settings.redis_url, decode_responses=True)
return self._client
async def get_json(self, key: str) -> dict[str, Any] | None:
"""读取缓存并反序列化为 dict
任何异常(连接失败 / JSON 解析失败)或值非 dict 时都降级为未命中,返回 None。
"""
try:
raw = await self._get_client().get(key)
if raw is None:
return None
data = json.loads(raw)
return data if isinstance(data, dict) else None
except Exception:
logger.warning("Redis 读取缓存失败,降级为未命中", key=key, exc_info=True)
return None
async def set_json(self, key: str, value: dict[str, Any], ttl: int | None = None) -> bool:
"""写入缓存(SETEX),ttl 缺省取 settings.cache_ttl
异常时记录日志并返回 False,不影响主流程。
"""
try:
await self._get_client().setex(
key, ttl if ttl is not None else settings.cache_ttl, json.dumps(value, ensure_ascii=False)
)
return True
except Exception:
logger.warning("Redis 写入缓存失败", key=key, exc_info=True)
return False
async def close(self) -> None:
"""关闭底层连接"""
if self._client is not None:
await self._client.aclose()
self._client = None
# 模块级懒加载单例,避免每请求重建客户端
_cache: RedisCache | None = None
def get_cache() -> RedisCache:
"""获取全局 RedisCache 单例"""
global _cache
if _cache is None:
_cache = RedisCache()
return _cache
+692
View File
@@ -0,0 +1,692 @@
<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>知识库管理后台</title>
<style>
* { box-sizing: border-box; }
body {
margin: 0;
font-family: -apple-system, "PingFang SC", "Microsoft YaHei", sans-serif;
background: #f5f6f8;
color: #24292f;
}
header {
background: #1f2937;
color: #fff;
padding: 12px 24px;
}
header h1 { font-size: 18px; margin: 0 0 10px; }
nav button {
background: transparent;
border: 1px solid #4b5563;
color: #d1d5db;
padding: 6px 14px;
margin-right: 8px;
border-radius: 4px;
cursor: pointer;
font-size: 14px;
}
nav button.active { background: #3b82f6; border-color: #3b82f6; color: #fff; }
main { padding: 20px 24px; max-width: 1080px; margin: 0 auto; }
section { background: #fff; border: 1px solid #e5e7eb; border-radius: 6px; padding: 16px 20px; }
section h2 { font-size: 16px; margin: 0 0 14px; border-bottom: 1px solid #e5e7eb; padding-bottom: 8px; }
.hidden { display: none; }
.error-bar {
background: #fef2f2;
border: 1px solid #fca5a5;
color: #b91c1c;
border-radius: 4px;
padding: 8px 12px;
margin-bottom: 12px;
font-size: 13px;
white-space: pre-wrap;
}
.cards { display: flex; flex-wrap: wrap; gap: 12px; margin-bottom: 16px; }
.card {
flex: 1 1 140px;
border: 1px solid #e5e7eb;
border-radius: 6px;
padding: 12px;
background: #fafafa;
}
.card .num { font-size: 24px; font-weight: 600; }
.card .label { font-size: 12px; color: #6b7280; margin-top: 4px; }
.bar-row { display: flex; align-items: center; margin-bottom: 6px; font-size: 13px; }
.bar-row .name { width: 160px; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; }
.bar-row .bar-track { flex: 1; background: #f3f4f6; border-radius: 3px; height: 16px; margin: 0 8px; }
.bar-row .bar-fill { background: #3b82f6; height: 100%; border-radius: 3px; min-width: 2px; }
.bar-row .count { width: 48px; text-align: right; color: #374151; }
table { width: 100%; border-collapse: collapse; font-size: 13px; }
th, td { border-bottom: 1px solid #e5e7eb; padding: 8px; text-align: left; vertical-align: top; }
th { background: #f9fafb; color: #374151; }
button.action {
border: 1px solid #d1d5db;
background: #fff;
border-radius: 4px;
padding: 4px 10px;
margin-right: 6px;
cursor: pointer;
font-size: 12px;
}
button.action:hover { background: #f3f4f6; }
button.danger { color: #b91c1c; border-color: #fca5a5; }
button.primary {
background: #3b82f6; color: #fff; border: none;
border-radius: 4px; padding: 8px 18px; cursor: pointer; font-size: 14px;
}
button.primary:disabled { background: #9ca3af; cursor: not-allowed; }
form .field { margin-bottom: 12px; }
form label { display: block; font-size: 13px; margin-bottom: 4px; color: #374151; }
input[type="text"], input[type="number"], textarea {
width: 100%; border: 1px solid #d1d5db; border-radius: 4px;
padding: 8px; font-size: 13px; font-family: inherit;
}
textarea { min-height: 160px; resize: vertical; }
.result-box {
margin-top: 14px; border: 1px solid #e5e7eb; border-radius: 6px;
background: #fafafa; padding: 12px; font-size: 13px;
}
.result-box .kv { margin-bottom: 6px; }
.result-box .k { color: #6b7280; display: inline-block; min-width: 110px; }
.detail-panel {
margin-top: 14px; border: 1px solid #d1d5db; border-radius: 6px; padding: 12px;
}
.detail-panel h3 { margin: 0 0 8px; font-size: 14px; }
.detail-panel pre {
background: #f9fafb; border: 1px solid #e5e7eb; border-radius: 4px;
padding: 8px; white-space: pre-wrap; word-break: break-word;
font-size: 12px; max-height: 260px; overflow: auto;
}
.node-item { border-left: 3px solid #93c5fd; padding: 4px 8px; margin-bottom: 6px; background: #f8fafc; }
.node-item .path { font-size: 12px; color: #6b7280; }
.hit-item { border: 1px solid #e5e7eb; border-radius: 6px; padding: 10px; margin-bottom: 10px; }
.hit-item .meta { font-size: 12px; color: #6b7280; margin-bottom: 6px; }
.hit-item .snippet { font-size: 13px; white-space: pre-wrap; word-break: break-word; }
.tag { display: inline-block; background: #eff6ff; color: #1d4ed8; border-radius: 3px; padding: 1px 6px; font-size: 12px; margin-right: 4px; }
.fallback-flag { color: #b45309; background: #fffbeb; border: 1px solid #fcd34d; border-radius: 4px; padding: 6px 10px; font-size: 13px; margin-bottom: 10px; }
.toolbar { margin: 12px 0; }
.muted { color: #6b7280; font-size: 12px; }
.status-badge {
display: inline-block; border-radius: 3px; padding: 1px 8px;
font-size: 12px; margin-left: 6px; border: 1px solid transparent;
}
.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; }
</style>
</head>
<body>
<header>
<h1>知识库管理后台</h1>
<nav id="nav">
<button type="button" data-target="section-overview" class="active">概览</button>
<button type="button" data-target="section-docs">文档管理</button>
<button type="button" data-target="section-ingest">文档入库</button>
<button type="button" data-target="section-search">检索测试台</button>
<button type="button" data-target="section-categories">类目列表</button>
</nav>
</header>
<main>
<section id="section-overview">
<h2>概览</h2>
<div class="error-bar hidden" id="error-overview"></div>
<div class="toolbar">
<button type="button" class="action" id="btn-refresh-overview">刷新</button>
</div>
<div class="cards" id="overview-cards"></div>
<h3 style="font-size:14px;">类目分布</h3>
<div id="overview-categories"></div>
</section>
<section id="section-docs" class="hidden">
<h2>文档管理</h2>
<div class="error-bar hidden" id="error-docs"></div>
<table>
<thead>
<tr><th>标题</th><th>类目</th><th>标签</th><th>L1 摘要</th><th>操作</th></tr>
</thead>
<tbody id="docs-tbody"></tbody>
</table>
<div class="toolbar">
<button type="button" class="action" id="btn-load-more">加载更多</button>
<span class="muted" id="docs-status"></span>
</div>
<div class="detail-panel hidden" id="doc-detail"></div>
</section>
<section id="section-ingest" class="hidden">
<h2>文档入库</h2>
<div class="error-bar hidden" id="error-ingest"></div>
<form id="ingest-form">
<div class="field">
<label for="ingest-title">标题</label>
<input type="text" id="ingest-title" name="title" required>
</div>
<div class="field">
<label for="ingest-source">来源</label>
<input type="text" id="ingest-source" name="source" placeholder="例如:manual / web / file">
</div>
<div class="field">
<label for="ingest-text">正文</label>
<textarea id="ingest-text" name="text" required></textarea>
</div>
<button type="submit" class="primary" id="btn-ingest-submit">提交入库</button>
</form>
<div class="result-box hidden" id="ingest-result"></div>
</section>
<section id="section-search" class="hidden">
<h2>检索测试台</h2>
<div class="error-bar hidden" id="error-search"></div>
<form id="search-form">
<div class="field">
<label for="search-query">查询语句</label>
<input type="text" id="search-query" name="query" required>
</div>
<div class="field">
<label for="search-topk">top_k</label>
<input type="number" id="search-topk" name="top_k" value="5" min="1" max="50">
</div>
<button type="submit" class="primary" id="btn-search-submit">检索</button>
</form>
<div class="result-box hidden" id="search-result"></div>
</section>
<section id="section-categories" class="hidden">
<h2>类目列表</h2>
<div class="error-bar hidden" id="error-categories"></div>
<table>
<thead>
<tr><th>名称</th><th>描述</th></tr>
</thead>
<tbody id="categories-tbody"></tbody>
</table>
</section>
</main>
<script>
"use strict";
/* ---------- 通用工具 ---------- */
function el(tag, text, className) {
var node = document.createElement(tag);
if (className) { node.className = className; }
if (text !== undefined && text !== null) { node.textContent = String(text); }
return node;
}
function clearChildren(node) {
while (node.firstChild) { node.removeChild(node.firstChild); }
}
function showError(boxId, err) {
var box = document.getElementById(boxId);
clearChildren(box);
var code = (err && err.code !== undefined) ? err.code : "网络错误";
var message = (err && err.message) ? err.message : String(err);
box.textContent = "错误 [" + code + "] " + message;
box.classList.remove("hidden");
}
function hideError(boxId) {
document.getElementById(boxId).classList.add("hidden");
}
/* 统一 API 封装:code !== 0 抛错,网络错误同样捕获 */
function api(path, options) {
return fetch(path, options).then(function (resp) {
return resp.json().catch(function () {
throw { code: "HTTP " + resp.status, message: "响应解析失败" };
});
}).then(function (body) {
if (body.code !== 0) {
throw { code: body.code, message: body.message };
}
return body.data;
}).catch(function (err) {
if (err && err.code !== undefined) { throw err; }
throw { code: "NETWORK", message: "网络请求失败: " + (err && err.message ? err.message : err) };
});
}
function truncate(text, maxLen) {
if (!text) { return ""; }
var s = String(text);
return s.length > maxLen ? s.slice(0, maxLen) + "…" : s;
}
/* ---------- 导航切换 ---------- */
var loadedOnce = {};
function activateSection(targetId) {
var sections = document.querySelectorAll("main section");
sections.forEach(function (sec) {
sec.classList.toggle("hidden", sec.id !== targetId);
});
var buttons = document.querySelectorAll("#nav button");
buttons.forEach(function (btn) {
btn.classList.toggle("active", btn.getAttribute("data-target") === targetId);
});
if (!loadedOnce[targetId]) {
loadedOnce[targetId] = true;
if (targetId === "section-overview") { loadOverview(); }
if (targetId === "section-docs") { loadDocuments(true); }
if (targetId === "section-categories") { loadCategories(); }
}
}
document.getElementById("nav").addEventListener("click", function (event) {
var btn = event.target.closest("button[data-target]");
if (btn) { activateSection(btn.getAttribute("data-target")); }
});
/* ---------- 1. 概览 ---------- */
function loadOverview() {
hideError("error-overview");
api("/api/v1/knowledge/stats").then(function (data) {
renderOverviewCards(data);
renderOverviewCategories(data.categories || {});
}).catch(function (err) { showError("error-overview", err); });
}
function renderOverviewCards(data) {
var container = document.getElementById("overview-cards");
clearChildren(container);
var collections = data.collections || {};
var cards = [
["L1 文档总结", collections.doc_l1],
["L2 大纲节点", collections.doc_l2],
["L3 内容大纲", collections.doc_l3],
["Chunks", collections.chunks],
["文档总数", data.documents_total],
["未分类文档", data.uncategorized_count]
];
cards.forEach(function (pair) {
var card = el("div", null, "card");
card.appendChild(el("div", pair[1] === undefined ? 0 : pair[1], "num"));
card.appendChild(el("div", pair[0], "label"));
container.appendChild(card);
});
}
function renderOverviewCategories(categories) {
var container = document.getElementById("overview-categories");
clearChildren(container);
var names = Object.keys(categories);
if (names.length === 0) {
container.appendChild(el("div", "暂无数据", "muted"));
return;
}
var max = 1;
names.forEach(function (name) { max = Math.max(max, categories[name]); });
names.sort(function (a, b) { return categories[b] - categories[a]; });
names.forEach(function (name) {
var count = categories[name];
var row = el("div", null, "bar-row");
row.appendChild(el("span", name, "name"));
var track = el("div", null, "bar-track");
var fill = el("div", null, "bar-fill");
fill.style.width = Math.round((count / max) * 100) + "%";
track.appendChild(fill);
row.appendChild(track);
row.appendChild(el("span", count, "count"));
container.appendChild(row);
});
}
document.getElementById("btn-refresh-overview").addEventListener("click", loadOverview);
/* ---------- 2. 文档管理 ---------- */
var docsOffset = null;
function loadDocuments(reset) {
hideError("error-docs");
if (reset) {
docsOffset = null;
clearChildren(document.getElementById("docs-tbody"));
}
var path = "/api/v1/documents?limit=20";
if (docsOffset !== null) { path += "&offset=" + encodeURIComponent(docsOffset); }
api(path).then(function (data) {
renderDocuments(data.items || []);
docsOffset = data.next_offset;
document.getElementById("btn-load-more").disabled = docsOffset === null;
document.getElementById("docs-status").textContent =
docsOffset === null ? "已加载全部" : "还有更多";
}).catch(function (err) { showError("error-docs", err); });
}
function renderDocuments(items) {
var tbody = document.getElementById("docs-tbody");
items.forEach(function (item) {
var tr = el("tr");
tr.appendChild(el("td", item.title || "(无标题)"));
tr.appendChild(el("td", item.category || ""));
var tagsTd = el("td");
(item.tags || []).forEach(function (tag) { tagsTd.appendChild(el("span", tag, "tag")); });
tr.appendChild(tagsTd);
tr.appendChild(el("td", truncate(item.summary, 80)));
var opsTd = el("td");
var detailBtn = el("button", "详情", "action");
detailBtn.type = "button";
detailBtn.addEventListener("click", function () { loadDocDetail(item.doc_id); });
var deleteBtn = el("button", "删除", "action danger");
deleteBtn.type = "button";
deleteBtn.addEventListener("click", function () { deleteDocument(item.doc_id, item.title); });
opsTd.appendChild(detailBtn);
opsTd.appendChild(deleteBtn);
tr.appendChild(opsTd);
tbody.appendChild(tr);
});
}
document.getElementById("btn-load-more").addEventListener("click", function () {
loadDocuments(false);
});
function loadDocDetail(docId) {
hideError("error-docs");
api("/api/v1/documents/" + encodeURIComponent(docId)).then(function (data) {
renderDocDetail(data);
}).catch(function (err) { showError("error-docs", err); });
}
function renderDocDetail(data) {
var panel = document.getElementById("doc-detail");
clearChildren(panel);
panel.classList.remove("hidden");
var l1 = data.l1 || {};
panel.appendChild(el("h3", "文档详情:" + (l1.title || data.l1 && l1.doc_id || "")));
var meta = el("div", null, "kv");
meta.appendChild(el("span", "doc_id", "k"));
meta.appendChild(el("span", l1.doc_id || ""));
panel.appendChild(meta);
var meta2 = el("div", null, "kv");
meta2.appendChild(el("span", "类目 / 标签", "k"));
meta2.appendChild(el("span", (l1.category || "") + " / " + (l1.tags || []).join(", ")));
panel.appendChild(meta2);
var meta3 = el("div", null, "kv");
meta3.appendChild(el("span", "chunks_count", "k"));
meta3.appendChild(el("span", data.chunks_count));
panel.appendChild(meta3);
panel.appendChild(el("h3", "L1 全文"));
panel.appendChild(el("pre", l1.text || ""));
panel.appendChild(el("h3", "L2 节点(" + (data.l2_nodes || []).length + ""));
renderNodes(panel, data.l2_nodes || []);
panel.appendChild(el("h3", "L3 节点(" + (data.l3_nodes || []).length + ""));
renderNodes(panel, data.l3_nodes || []);
}
function renderNodes(panel, nodes) {
if (nodes.length === 0) {
panel.appendChild(el("div", "无", "muted"));
return;
}
nodes.forEach(function (node) {
var item = el("div", null, "node-item");
item.appendChild(el("div", node.section_path || "", "path"));
item.appendChild(el("div", node.text || ""));
panel.appendChild(item);
});
}
function deleteDocument(docId, title) {
if (!confirm("确定删除文档「" + (title || docId) + "」(" + docId + ") 吗?该操作将删除四层集合中的全部数据,不可恢复。")) {
return;
}
hideError("error-docs");
api("/api/v1/documents/" + encodeURIComponent(docId), { method: "DELETE" }).then(function (data) {
document.getElementById("doc-detail").classList.add("hidden");
loadDocuments(true);
loadOverview();
var status = document.getElementById("docs-status");
status.textContent = "已删除 " + (data.deleted_total || 0) + " 条数据";
}).catch(function (err) { showError("error-docs", err); });
}
/* ---------- 3. 文档入库(异步任务轮询) ---------- */
var INGEST_STATUS_TEXT = {
pending: "排队中",
summarizing: "总结中",
classifying: "分类中",
embedding: "向量化中",
writing: "写入中",
done: "完成",
failed: "失败"
};
var INGEST_POLL_INTERVAL_MS = 2000;
var INGEST_POLL_MAX_ATTEMPTS = 150; /* 150 次 × 2s = 5 分钟超时 */
var ingestPollTimer = null;
function stopIngestPolling() {
if (ingestPollTimer !== null) {
clearInterval(ingestPollTimer);
ingestPollTimer = null;
}
}
function restoreIngestButton() {
var submitBtn = document.getElementById("btn-ingest-submit");
submitBtn.disabled = false;
submitBtn.textContent = "提交入库";
}
function makeStatusBadge(status) {
var cls = "status-running";
if (status === "done") { cls = "status-done"; }
if (status === "failed") { cls = "status-failed"; }
var badge = el("span", INGEST_STATUS_TEXT[status] || status, "status-badge " + cls);
badge.id = "ingest-status-badge";
return badge;
}
function updateIngestStatus(status) {
var old = document.getElementById("ingest-status-badge");
if (old && old.parentNode) {
old.parentNode.replaceChild(makeStatusBadge(status), old);
}
}
function ingestKv(box, label, value) {
var row = el("div", null, "kv");
row.appendChild(el("span", label, "k"));
row.appendChild(el("span", value));
box.appendChild(row);
}
function renderIngestHeader(taskId, status) {
var box = document.getElementById("ingest-result");
clearChildren(box);
box.classList.remove("hidden");
var row = el("div", null, "kv");
row.appendChild(el("span", "任务已提交:" + taskId));
row.appendChild(makeStatusBadge(status));
box.appendChild(row);
return box;
}
function renderIngestDone(task) {
var box = renderIngestHeader(task.task_id, "done");
var result = task.result || {};
var summary = result.summary || {};
ingestKv(box, "document_id", result.document_id);
ingestKv(box, "类目", (result.category || "") + "(置信度 " + result.category_confidence + "");
ingestKv(box, "标签", (result.tags || []).join(", "));
ingestKv(box, "总结层级 level", summary.level);
ingestKv(box, "写入集合", result.collection);
ingestKv(box, "chunks_count", result.chunks_count);
box.appendChild(el("h3", "L1 总结"));
box.appendChild(el("pre", summary.l1_summary || ""));
}
function renderIngestFailed(task) {
var box = renderIngestHeader(task.task_id, "failed");
var error = task.error || {};
ingestKv(box, "失败阶段", error.stage || "");
ingestKv(box, "错误信息", error.message || "");
var partial = error.partial_summary;
if (partial && partial.l1_summary) {
box.appendChild(el("div", "已产出总结保留:任务失败前已生成 L1 摘要", "fallback-flag"));
box.appendChild(el("h3", "L1 总结"));
box.appendChild(el("pre", partial.l1_summary));
}
}
function renderIngestTimeout(taskId) {
var box = document.getElementById("ingest-result");
box.appendChild(el("div", "任务仍在进行,可稍后凭 task_id 查询:" + taskId, "fallback-flag"));
}
function startIngestPolling(taskId) {
var attempts = 0;
document.getElementById("btn-ingest-submit").textContent = "入库中…";
ingestPollTimer = setInterval(function () {
attempts += 1;
if (attempts > INGEST_POLL_MAX_ATTEMPTS) {
stopIngestPolling();
renderIngestTimeout(taskId);
restoreIngestButton();
return;
}
api("/api/v1/documents/tasks/" + encodeURIComponent(taskId)).then(function (task) {
updateIngestStatus(task.status);
if (task.status === "done") {
stopIngestPolling();
renderIngestDone(task);
restoreIngestButton();
} else if (task.status === "failed") {
stopIngestPolling();
renderIngestFailed(task);
restoreIngestButton();
}
}).catch(function (err) {
stopIngestPolling();
showError("error-ingest", err);
restoreIngestButton();
});
}, INGEST_POLL_INTERVAL_MS);
}
document.getElementById("ingest-form").addEventListener("submit", function (event) {
event.preventDefault();
hideError("error-ingest");
stopIngestPolling();
var submitBtn = document.getElementById("btn-ingest-submit");
submitBtn.disabled = true;
submitBtn.textContent = "提交中…";
var payload = {
title: document.getElementById("ingest-title").value,
source: document.getElementById("ingest-source").value,
text: document.getElementById("ingest-text").value
};
api("/api/v1/documents", {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify(payload)
}).then(function (data) {
renderIngestHeader(data.task_id, data.status || "pending");
startIngestPolling(data.task_id);
}).catch(function (err) {
showError("error-ingest", err);
restoreIngestButton();
});
});
/* ---------- 4. 检索测试台 ---------- */
document.getElementById("search-form").addEventListener("submit", function (event) {
event.preventDefault();
hideError("error-search");
var submitBtn = document.getElementById("btn-search-submit");
submitBtn.disabled = true;
var payload = {
query: document.getElementById("search-query").value,
top_k: parseInt(document.getElementById("search-topk").value, 10) || 5
};
api("/api/v1/search", {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify(payload)
}).then(function (data) {
renderSearchResult(data);
}).catch(function (err) {
showError("error-search", err);
}).finally(function () {
submitBtn.disabled = false;
});
});
function renderSearchResult(data) {
var box = document.getElementById("search-result");
clearChildren(box);
box.classList.remove("hidden");
var routed = el("div", null, "kv");
routed.appendChild(el("span", "routed_categories", "k"));
routed.appendChild(el("span", (data.routed_categories || []).join(", ") || "(无)"));
box.appendChild(routed);
if (data.fallback) {
box.appendChild(el("div", "fallback:本次检索触发了回退策略(路由类目无命中,已降级全局检索)", "fallback-flag"));
}
var hits = data.hits || [];
box.appendChild(el("div", "命中 " + hits.length + " 条", "muted"));
hits.forEach(function (hit) {
var item = el("div", null, "hit-item");
var meta = el("div", null, "meta");
meta.textContent = "score=" + hit.score + " | doc_id=" + hit.doc_id +
" | 标题=" + (hit.title || "") + " | section=" + (hit.section_path || "");
item.appendChild(meta);
item.appendChild(el("div", hit.text || "", "snippet"));
if (hit.doc_summary) {
var summaryRow = el("div", null, "meta");
summaryRow.textContent = "文档摘要:" + hit.doc_summary;
item.appendChild(summaryRow);
}
box.appendChild(item);
});
}
/* ---------- 5. 类目列表 ---------- */
function loadCategories() {
hideError("error-categories");
api("/api/v1/knowledge/categories").then(function (data) {
var tbody = document.getElementById("categories-tbody");
clearChildren(tbody);
(data.categories || []).forEach(function (cat) {
var tr = el("tr");
tr.appendChild(el("td", cat.name));
tr.appendChild(el("td", cat.description || ""));
tbody.appendChild(tr);
});
}).catch(function (err) { showError("error-categories", err); });
}
/* ---------- 初始化 ---------- */
activateSection("section-overview");
</script>
</body>
</html>
View File