641e33b834
- 语义检索(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>
265 lines
9.3 KiB
Python
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,
|
|
}
|