Files
law-kb/app/services/retriever.py
T
freedak 641e33b834 feat: QYLAW 法律法规知识库
- 语义检索(FAISS + embedding)+ 精确查找(法规名+条号)
- RAG 问答(SSE 流式,支持 thinking 折叠显示)
- 法规浏览(原文阅读)
- 历史记录(检索+对话持久化到 SQLite)
- 设置页(系统提示词/模板/LLM 参数可配置)
- 检索质量评估脚本

Generated with [Devin](https://devin.ai)

Co-Authored-By: Devin <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-08-07 14:55:25 +08:00

265 lines
9.3 KiB
Python

"""FAISS 检索器 — 加载索引 + SQLite metadata,实现语义检索 + 过滤"""
import sqlite3
import logging
from typing import List, Optional, Dict, Any
import numpy as np
from app.config import (
FAISS_INDEX_PATH,
SQLITE_PATH,
HNSW_EF_SEARCH,
)
logger = logging.getLogger(__name__)
class Retriever:
"""检索器单例"""
_instance: Optional["Retriever"] = None
@classmethod
def get_instance(cls) -> "Retriever":
if cls._instance is None:
cls._instance = cls()
return cls._instance
def __init__(self):
self._faiss = None
self._conn: Optional[sqlite3.Connection] = None
self._dim: Optional[int] = None
self._ready = False
def is_ready(self) -> bool:
return self._ready
def load(self):
"""加载 FAISS 索引和 SQLite"""
if not FAISS_INDEX_PATH.exists():
logger.warning(f"FAISS 索引不存在: {FAISS_INDEX_PATH}")
return
if not SQLITE_PATH.exists():
logger.warning(f"SQLite 不存在: {SQLITE_PATH}")
return
import faiss
logger.info(f"加载 FAISS 索引: {FAISS_INDEX_PATH}")
self._faiss = faiss.read_index(str(FAISS_INDEX_PATH))
self._dim = self._faiss.d
# 设置 HNSW 搜索参数
if hasattr(self._faiss, "hnsw"):
self._faiss.hnsw.efSearch = HNSW_EF_SEARCH
logger.info(f"加载 SQLite: {SQLITE_PATH}")
self._conn = sqlite3.connect(str(SQLITE_PATH), check_same_thread=False)
self._conn.row_factory = sqlite3.Row
self._ready = True
count = self._faiss.ntotal
logger.info(f"索引加载完成: {count} 向量, dim={self._dim}")
def search(
self,
query_vec: List[float],
top_k: int = 20,
category: Optional[str] = None,
province: Optional[str] = None,
city: Optional[str] = None,
) -> List[Dict[str, Any]]:
"""语义检索 + metadata 过滤
Args:
query_vec: 查询向量
top_k: 返回数量
category: 法规类别过滤
province: 省份过滤
city: 市级过滤
Returns:
结果列表,每项含 law_id/law_name/category/chapter/clause_no/content/score/file_path/province
"""
if not self._ready:
return []
# FAISS 检索(取 top_k * 3 用于过滤后仍有足够结果)
vec = np.array([query_vec], dtype=np.float32)
fetch_k = min(top_k * 3 if (category or province or city) else top_k, self._faiss.ntotal)
scores, indices = self._faiss.search(vec, fetch_k)
results = []
for score, faiss_idx in zip(scores[0], indices[0]):
if faiss_idx < 0:
continue
# 从 SQLite 取 metadata
row = self._conn.execute(
"""
SELECT c.id, c.law_id, c.chapter, c.clause_no, c.content, c.faiss_idx,
l.name, l.category, l.publish_date, l.file_path, l.province, l.city
FROM clauses c JOIN laws l ON c.law_id = l.id
WHERE c.faiss_idx = ?
""",
(int(faiss_idx),),
).fetchone()
if row is None:
continue
# metadata 过滤
if category and row["category"] != category:
continue
if province and row["province"] != province:
continue
if city and row["city"] != city:
continue
results.append({
"clause_id": row["id"],
"law_id": row["law_id"],
"law_name": row["name"],
"category": row["category"],
"chapter": row["chapter"],
"clause_no": row["clause_no"],
"content": row["content"],
"score": float(score),
"file_path": row["file_path"],
"province": row["province"],
"city": row["city"],
"publish_date": row["publish_date"],
})
if len(results) >= top_k:
break
return results
def keyword_search(
self,
query: str,
top_k: int = 20,
category: Optional[str] = None,
province: Optional[str] = None,
city: Optional[str] = None,
) -> List[Dict[str, Any]]:
"""关键词精确搜索(按法规名 + 条号 + 条文内容匹配)
支持"法规名 条号"格式(如"郑州市劳动用工条例 第三十二条"),
也支持纯关键词搜索条文内容。
Returns:
结果列表,格式与 search() 一致,score 为匹配相关度(非 FAISS 距离)
"""
if not self._ready or not query.strip():
return []
query = query.strip()
# 尝试解析"法规名 条号"格式
# 中文条号模式:第X条
import re
clause_match = re.search(r'(第[一二三四五六七八九十百千零\d]+条)', query)
clause_no = clause_match.group(1) if clause_match else None
# 法规名 = 去掉条号后的部分
law_name_query = re.sub(r'\s*第[一二三四五六七八九十百千零\d]+条\s*', '', query).strip()
results = []
try:
if clause_no and law_name_query:
# 精确匹配:法规名 LIKE + 条号 =
rows = self._conn.execute(
"""
SELECT c.id, c.law_id, c.chapter, c.clause_no, c.content, c.faiss_idx,
l.name, l.category, l.publish_date, l.file_path, l.province, l.city
FROM clauses c JOIN laws l ON c.law_id = l.id
WHERE l.name LIKE ? AND c.clause_no = ?
""",
(f"%{law_name_query}%", clause_no),
).fetchall()
elif clause_no:
# 只按条号搜索
rows = self._conn.execute(
"""
SELECT c.id, c.law_id, c.chapter, c.clause_no, c.content, c.faiss_idx,
l.name, l.category, l.publish_date, l.file_path, l.province, l.city
FROM clauses c JOIN laws l ON c.law_id = l.id
WHERE c.clause_no = ?
""",
(clause_no,),
).fetchall()
else:
# 纯关键词:搜索法规名或条文内容
rows = self._conn.execute(
"""
SELECT c.id, c.law_id, c.chapter, c.clause_no, c.content, c.faiss_idx,
l.name, l.category, l.publish_date, l.file_path, l.province, l.city
FROM clauses c JOIN laws l ON c.law_id = l.id
WHERE l.name LIKE ? OR c.content LIKE ?
""",
(f"%{query}%", f"%{query}%"),
).fetchall()
for row in rows:
# metadata 过滤
if category and row["category"] != category:
continue
if province and row["province"] != province:
continue
if city and row["city"] != city:
continue
results.append({
"clause_id": row["id"],
"law_id": row["law_id"],
"law_name": row["name"],
"category": row["category"],
"chapter": row["chapter"],
"clause_no": row["clause_no"],
"content": row["content"],
"score": 0.0, # 关键词搜索无相似度分数
"file_path": row["file_path"],
"province": row["province"],
"city": row["city"],
"publish_date": row["publish_date"],
})
if len(results) >= top_k:
break
except Exception as e:
logger.error(f"关键词搜索失败: {e}")
return results
def get_stats(self) -> Dict[str, Any]:
"""返回索引统计"""
if not self._ready:
return {
"total_laws": 0,
"total_clauses": 0,
"category_stats": {},
"index_built_at": None,
"faiss_dim": None,
"ready": False,
}
total_laws = self._conn.execute("SELECT COUNT(*) FROM laws").fetchone()[0]
total_clauses = self._conn.execute("SELECT COUNT(*) FROM clauses").fetchone()[0]
cat_rows = self._conn.execute(
"SELECT category, COUNT(*) as cnt FROM laws GROUP BY category"
).fetchall()
category_stats = {row["category"]: row["cnt"] for row in cat_rows}
# 索引构建时间(取最新一条 law 的 indexed_at)
built_at_row = self._conn.execute(
"SELECT indexed_at FROM laws ORDER BY indexed_at DESC LIMIT 1"
).fetchone()
index_built_at = built_at_row["indexed_at"] if built_at_row else None
return {
"total_laws": total_laws,
"total_clauses": total_clauses,
"category_stats": category_stats,
"index_built_at": index_built_at,
"faiss_dim": self._dim,
"ready": True,
}