Files
more_dots/services/core/template_matcher.py
T
2026-03-24 18:07:22 +08:00

202 lines
7.8 KiB
Python

import json
import re
from pathlib import Path
from typing import Any, Dict, Optional, Set
from config import Config
from services.core.sql_prompt_manager import get_sql_prompt_manager
from services.integrations.ragflow_client import RagflowClient, extract_table_name
IDENTIFIER_RE = re.compile(r"[a-zA-Z_][a-zA-Z0-9_]*")
EXPLICIT_FILTER_FIELD_RE = re.compile(r"([a-zA-Z_][a-zA-Z0-9_]*)\s*=")
def _normalize_term(term: str) -> str:
return str(term or "").strip().lower()
class TemplateMatcher:
"""模板匹配器:RAGFlow 检索"""
KEYWORD_MATCH_BONUS = 50
def __init__(self):
self._ragflow = RagflowClient()
self._sql_prompt_manager = get_sql_prompt_manager()
cfg = Config.get_section("ragflow")
self._dataset_id = (cfg.get("table_retrieval_dataset_id") or "").strip()
self._top_k = int(cfg.get("retrieval_top_k", 3))
self._non_empty_tables, self._table_keywords = self._load_table_config()
self._table_terms_cache: Dict[str, Set[str]] = {}
def _load_table_config(self) -> tuple[Optional[Set[str]], Dict[str, Set[str]]]:
"""从本地 tables.json 读取非空模板表集合和关键词映射。"""
try:
tables_path = Path(__file__).resolve().parents[2] / "config" / "table_retrieval_prompts" / "tables.json"
with open(tables_path, "r", encoding="utf-8") as f:
data = json.load(f)
if not isinstance(data, dict):
return None, {}
non_empty: Set[str] = set()
keywords_map: Dict[str, Set[str]] = {}
for table_name, templates in data.items():
if not isinstance(table_name, str):
continue
if isinstance(templates, list) and len(templates) > 0:
non_empty.add(table_name)
keywords_map[table_name] = {_normalize_term(kw) for kw in templates if isinstance(kw, str)}
return non_empty, keywords_map
except Exception:
return None, {}
def _validate(self) -> None:
if not self._dataset_id:
raise RuntimeError("未配置 ragflow.table_retrieval_dataset_id,无法进行表名检索")
@staticmethod
def _extract_query_terms(normalized_text: str) -> Set[str]:
return {
_normalize_term(match.group(0))
for match in IDENTIFIER_RE.finditer(normalized_text or "")
}
@staticmethod
def _extract_explicit_filter_fields(normalized_text: str) -> Set[str]:
return {
_normalize_term(match.group(1))
for match in EXPLICIT_FILTER_FIELD_RE.finditer(normalized_text or "")
}
def _load_table_terms(self, table_name: str) -> Set[str]:
cached = self._table_terms_cache.get(table_name)
if cached is not None:
return cached
prompt = self._sql_prompt_manager.get_prompt(table_name) or {}
terms: Set[str] = set()
for field in (prompt.get("data_model_specification") or {}).get("fields_list") or []:
if isinstance(field, str):
terms.add(_normalize_term(field))
field_ref = prompt.get("field_mapping_reference") or {}
self._collect_mapping_terms(field_ref, terms)
self._table_terms_cache[table_name] = terms
return terms
def _collect_mapping_terms(self, node: Any, terms: Set[str]) -> None:
if isinstance(node, dict):
for key, value in node.items():
if key == "alias" and isinstance(value, list):
for alias in value:
if isinstance(alias, str):
terms.add(_normalize_term(alias))
continue
if isinstance(value, dict):
if "alias" in value or "type" in value:
terms.add(_normalize_term(key))
self._collect_mapping_terms(value, terms)
elif isinstance(value, list):
# 对字段列表直接入词,增强字段覆盖匹配
if key.endswith("_fields") or key in {"fields_list", "list"}:
for item in value:
if isinstance(item, str):
terms.add(_normalize_term(item))
self._collect_mapping_terms(value, terms)
elif isinstance(node, list):
for item in node:
self._collect_mapping_terms(item, terms)
def _rank_candidates(self, normalized_text: str, candidates: list[Dict[str, Any]]) -> list[Dict[str, Any]]:
query_terms = self._extract_query_terms(normalized_text)
explicit_fields = self._extract_explicit_filter_fields(normalized_text)
if not query_terms or not candidates:
return candidates
ranked: list[Dict[str, Any]] = []
for index, candidate in enumerate(candidates):
table_name = candidate.get("table_name")
if not table_name:
continue
table_terms = self._load_table_terms(table_name)
overlap = len(query_terms & table_terms)
missing_explicit_fields = len([field for field in explicit_fields if field not in table_terms])
keyword_bonus = 0
table_keywords = self._table_keywords.get(table_name, set())
if table_keywords:
matched_keywords = query_terms & table_keywords
keyword_bonus = len(matched_keywords) * self.KEYWORD_MATCH_BONUS
rank_score = overlap + keyword_bonus - (missing_explicit_fields * 5)
ranked.append({
**candidate,
"rank_score": rank_score,
"rank_overlap": overlap,
"rank_keyword_bonus": keyword_bonus,
"rank_missing_explicit_fields": missing_explicit_fields,
"rank_index": index,
})
ranked.sort(key=lambda item: (item["rank_score"], item["rank_overlap"], -item["rank_index"]), reverse=True)
return ranked
def match(self, normalized_text: str) -> Dict[str, Any]:
"""返回匹配的表名与原始响应"""
self._validate()
try:
response = self._ragflow.retrieve(normalized_text, top_k=self._top_k, dataset_id=self._dataset_id)
except Exception as e:
return {"table_name": None, "candidates": [], "raw": {"error": str(e)}}
candidates = []
seen = set()
data = response.get("data") if isinstance(response, dict) else None
records = []
if isinstance(data, list):
records = data
elif isinstance(data, dict):
chunks = data.get("chunks")
if isinstance(chunks, list):
records = chunks
for item in records:
table_name = extract_table_name(item)
if self._non_empty_tables is not None and table_name not in self._non_empty_tables:
continue
if table_name and table_name not in seen:
seen.add(table_name)
candidates.append(
{
"table_name": table_name,
"metadata": item.get("metadata") or {},
"content": item.get("content") or item.get("text") or "",
}
)
candidates = self._rank_candidates(normalized_text, candidates)
matched = candidates[0]["table_name"] if candidates else None
return {"table_name": matched, "candidates": candidates, "raw": response}
_GLOBAL_TEMPLATE_MATCHER: TemplateMatcher | None = None
def get_template_matcher() -> TemplateMatcher:
"""获取全局 TemplateMatcher(单例)"""
global _GLOBAL_TEMPLATE_MATCHER
if _GLOBAL_TEMPLATE_MATCHER is None:
_GLOBAL_TEMPLATE_MATCHER = TemplateMatcher()
return _GLOBAL_TEMPLATE_MATCHER