65 lines
2.1 KiB
Python
65 lines
2.1 KiB
Python
import hashlib
|
|
import json
|
|
from typing import Any, Dict, Optional
|
|
|
|
from config import Config
|
|
from services.cache import NoopCache, RedisCache
|
|
from services.ragflow_client import RagflowClient, extract_table_name
|
|
|
|
|
|
class TemplateMatcher:
|
|
"""模板匹配器:RAGFlow + Redis 缓存"""
|
|
|
|
def __init__(self):
|
|
self._ragflow = RagflowClient()
|
|
self._cache = self._init_cache()
|
|
cfg = Config.get_section("ragflow")
|
|
self._cache_ttl = int(cfg.get("cache_ttl", 600))
|
|
|
|
def _init_cache(self):
|
|
cfg = Config.get_section("redis")
|
|
enabled = str(cfg.get("enabled", "false")).lower() in ("1", "true", "yes")
|
|
if not enabled:
|
|
return NoopCache()
|
|
|
|
url = cfg.get("url")
|
|
db = int(cfg.get("db", 0))
|
|
if not url:
|
|
return NoopCache()
|
|
|
|
try:
|
|
return RedisCache(url=url, db=db)
|
|
except Exception:
|
|
return NoopCache()
|
|
|
|
@staticmethod
|
|
def _cache_key(text: str) -> str:
|
|
return "ragflow:table:" + hashlib.sha256(text.encode("utf-8")).hexdigest()
|
|
|
|
def match(self, normalized_text: str) -> Dict[str, Any]:
|
|
"""返回匹配的表名与原始响应"""
|
|
key = self._cache_key(normalized_text)
|
|
cached = self._cache.get(key)
|
|
if cached:
|
|
return json.loads(cached)
|
|
try:
|
|
response = self._ragflow.retrieve(normalized_text, top_k=3)
|
|
except Exception as e:
|
|
result = {"table_name": None, "raw": {"error": str(e)}}
|
|
self._cache.set(key, json.dumps(result, ensure_ascii=False), self._cache_ttl)
|
|
return result
|
|
|
|
candidates = []
|
|
data = response.get("data") if isinstance(response, dict) else None
|
|
if isinstance(data, list):
|
|
for item in data:
|
|
table_name = extract_table_name(item)
|
|
if table_name:
|
|
candidates.append(table_name)
|
|
|
|
matched = candidates[0] if candidates else None
|
|
result = {"table_name": matched, "raw": response}
|
|
|
|
self._cache.set(key, json.dumps(result, ensure_ascii=False), self._cache_ttl)
|
|
return result
|