"""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, }