init
This commit is contained in:
@@ -2,11 +2,6 @@ from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
try:
|
||||
import redis
|
||||
except Exception:
|
||||
redis = None
|
||||
|
||||
|
||||
class CacheBase:
|
||||
"""缓存接口"""
|
||||
@@ -28,16 +23,3 @@ class NoopCache(CacheBase):
|
||||
return None
|
||||
|
||||
|
||||
class RedisCache(CacheBase):
|
||||
"""Redis 缓存"""
|
||||
|
||||
def __init__(self, url: str, db: int = 0):
|
||||
if redis is None:
|
||||
raise ImportError("未安装 redis 依赖")
|
||||
self._client = redis.Redis.from_url(url, db=db, decode_responses=True)
|
||||
|
||||
def get(self, key: str) -> Optional[str]:
|
||||
return self._client.get(key)
|
||||
|
||||
def set(self, key: str, value: str, ttl: int) -> None:
|
||||
self._client.set(key, value, ex=ttl)
|
||||
|
||||
@@ -3,12 +3,7 @@ import logging
|
||||
import socket
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
try:
|
||||
import nacos
|
||||
except Exception:
|
||||
nacos = None
|
||||
|
||||
import nacos
|
||||
from config import Config
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -29,3 +29,14 @@ class PromptManager:
|
||||
def list_prompts(self, group: str) -> list[str]:
|
||||
"""列出分组内提示词"""
|
||||
return list(self._data.get(group, {}).keys())
|
||||
|
||||
|
||||
_GLOBAL_PROMPT_MANAGER: Optional[PromptManager] = None
|
||||
|
||||
|
||||
def get_prompt_manager(config_path: Optional[str] = None) -> PromptManager:
|
||||
"""获取全局 PromptManager(单例)"""
|
||||
global _GLOBAL_PROMPT_MANAGER
|
||||
if _GLOBAL_PROMPT_MANAGER is None:
|
||||
_GLOBAL_PROMPT_MANAGER = PromptManager(config_path=config_path)
|
||||
return _GLOBAL_PROMPT_MANAGER
|
||||
|
||||
@@ -14,22 +14,27 @@ class RagflowClient:
|
||||
self._base_url = cfg.get("url", "")
|
||||
self._api_key = cfg.get("api_key", "")
|
||||
self._retrieval_path = cfg.get("retrieval", "/api/v1/retrieval")
|
||||
self._dataset_ids = cfg.get("dataset_ids", "")
|
||||
self._table_retrieval_dataset_id = cfg.get("table_retrieval_dataset_id", "")
|
||||
self._sql_gen_dataset_id = cfg.get("sql_gen_dataset_id", "")
|
||||
|
||||
def _build_url(self) -> str:
|
||||
return self._base_url.rstrip("/") + "/" + self._retrieval_path.lstrip("/")
|
||||
|
||||
def retrieve(self, query: str, top_k: int = 3) -> Dict[str, Any]:
|
||||
def retrieve(self, query: str, top_k: int = 3, dataset_id: Optional[str] = None, document_ids: Optional[str] = None) -> Dict[str, Any]:
|
||||
"""检索匹配文档"""
|
||||
if not self._base_url or not self._retrieval_path:
|
||||
raise RuntimeError("未配置 ragflow.url 或 ragflow.retrieval")
|
||||
if not dataset_id:
|
||||
raise ValueError("未提供 ragflow.dataset_id,无法进行检索")
|
||||
url = self._build_url()
|
||||
headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
|
||||
payload = {
|
||||
"dataset_ids": self._dataset_ids,
|
||||
"dataset_ids": dataset_id or "",
|
||||
"query": query,
|
||||
"top_k": top_k,
|
||||
}
|
||||
if document_ids:
|
||||
payload["document_ids"] = document_ids
|
||||
|
||||
with httpx.Client(timeout=30) as client:
|
||||
response = client.post(url, json=payload, headers=headers)
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
import json
|
||||
import os
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from config import Config
|
||||
|
||||
|
||||
def _build_document_for_table(table: str, templates: List[str]) -> str:
|
||||
lines = [f"table: {table}"]
|
||||
for t in templates:
|
||||
lines.append(f"- {t}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _build_sql_gen_document(table: str, prompt: Dict[str, any]) -> str:
|
||||
system_prompt = prompt.get("system_prompt", "")
|
||||
business_prompt = prompt.get("business_prompt", "")
|
||||
constraints = prompt.get("constraints", [])
|
||||
lines = [f"table: {table}"]
|
||||
if system_prompt:
|
||||
lines.append("[system] " + system_prompt)
|
||||
if business_prompt:
|
||||
lines.append("[business] " + business_prompt)
|
||||
if constraints:
|
||||
lines.append("[constraints]")
|
||||
for c in constraints:
|
||||
lines.append(f"- {c}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
class RagflowSync:
|
||||
"""RAGFlow 同步工具"""
|
||||
|
||||
def __init__(self):
|
||||
cfg = Config.get_section("ragflow")
|
||||
self._base_url = cfg.get("url", "").rstrip("/")
|
||||
self._api_key = cfg.get("api_key", "")
|
||||
self._table_retrieval_dataset_id = (cfg.get("table_retrieval_dataset_id") or "").strip()
|
||||
self._sql_gen_dataset_id = (cfg.get("sql_gen_dataset_id") or "").strip()
|
||||
self._upload_path = (cfg.get("upload") or "").strip()
|
||||
self._upload_mode = (cfg.get("upload_mode") or "overwrite").strip().lower()
|
||||
|
||||
def _validate_common(self) -> None:
|
||||
if not self._base_url:
|
||||
raise RuntimeError("未配置 ragflow.url")
|
||||
if not self._upload_path:
|
||||
raise RuntimeError("未配置 ragflow.upload 上传接口,请在 config/config.ini 中设置")
|
||||
if self._upload_mode not in ("overwrite", "append"):
|
||||
raise RuntimeError("ragflow.upload_mode 仅支持 overwrite 或 append")
|
||||
|
||||
def _post(self, documents: List[Dict[str, any]]):
|
||||
self._validate_common()
|
||||
url = self._base_url + "/" + self._upload_path.lstrip("/")
|
||||
headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
|
||||
payload: Dict[str, any] = {"documents": documents}
|
||||
if self._upload_mode == "overwrite":
|
||||
payload["mode"] = "overwrite"
|
||||
with httpx.Client(timeout=60) as client:
|
||||
response = client.post(url, json=payload, headers=headers)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
def sync_table_retrieval(self) -> Dict[str, any]:
|
||||
if not self._table_retrieval_dataset_id:
|
||||
raise RuntimeError("未配置 ragflow.table_retrieval_dataset_id,无法同步表名检索模板")
|
||||
root = os.path.dirname(os.path.dirname(__file__))
|
||||
tables_file = os.path.join(root, "config", "table_retrieval_prompts", "tables.json")
|
||||
with open(tables_file, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
|
||||
tables = data.get("tables", {})
|
||||
documents = []
|
||||
for table, templates in tables.items():
|
||||
doc = {
|
||||
"dataset_ids": self._table_retrieval_dataset_id,
|
||||
"content": _build_document_for_table(table, templates),
|
||||
"metadata": {"table": table},
|
||||
}
|
||||
documents.append(doc)
|
||||
|
||||
return self._post(documents)
|
||||
|
||||
def sync_sql_gen_prompts(self) -> Dict[str, any]:
|
||||
if not self._sql_gen_dataset_id:
|
||||
raise RuntimeError("未配置 ragflow.sql_gen_dataset_id,无法同步 SQL 生成提示词")
|
||||
root = os.path.dirname(os.path.dirname(__file__))
|
||||
prompts_dir = os.path.join(root, "config", "sql_gen_prompts")
|
||||
documents = []
|
||||
|
||||
for name in os.listdir(prompts_dir):
|
||||
if not name.endswith(".json"):
|
||||
continue
|
||||
path = os.path.join(prompts_dir, name)
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
prompt = json.load(f)
|
||||
table = prompt.get("table") or os.path.splitext(name)[0]
|
||||
doc = {
|
||||
"dataset_ids": self._sql_gen_dataset_id,
|
||||
"content": _build_sql_gen_document(table, prompt),
|
||||
"metadata": {"table": table},
|
||||
}
|
||||
documents.append(doc)
|
||||
|
||||
return self._post(documents)
|
||||
@@ -0,0 +1,26 @@
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
|
||||
class SqlPromptManager:
|
||||
"""按表名读取 SQL 提示词"""
|
||||
|
||||
def __init__(self, base_dir: Optional[str] = None):
|
||||
root_dir = os.path.dirname(os.path.dirname(__file__))
|
||||
self._base_dir = base_dir or os.path.join(root_dir, "config", "sql_gen_prompts")
|
||||
|
||||
@staticmethod
|
||||
def _safe_filename(name: str) -> str:
|
||||
return name.replace("..", "").replace("/", "_").replace("\\", "_")
|
||||
|
||||
def get_prompt(self, table_name: str) -> Optional[Dict[str, Any]]:
|
||||
"""读取指定表的提示词 JSON"""
|
||||
if not table_name:
|
||||
return None
|
||||
filename = self._safe_filename(table_name) + ".json"
|
||||
path = os.path.join(self._base_dir, filename)
|
||||
if not os.path.exists(path):
|
||||
return None
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
@@ -1,53 +1,28 @@
|
||||
import hashlib
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any, Dict
|
||||
|
||||
from config import Config
|
||||
from services.cache import NoopCache, RedisCache
|
||||
from services.ragflow_client import RagflowClient, extract_table_name
|
||||
|
||||
|
||||
class TemplateMatcher:
|
||||
"""模板匹配器:RAGFlow + Redis 缓存"""
|
||||
"""模板匹配器:RAGFlow 检索"""
|
||||
|
||||
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))
|
||||
self._dataset_id = (cfg.get("table_retrieval_dataset_id") or "").strip()
|
||||
|
||||
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 _validate(self) -> None:
|
||||
if not self._dataset_id:
|
||||
raise RuntimeError("未配置 ragflow.table_retrieval_dataset_id,无法进行表名检索")
|
||||
|
||||
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)
|
||||
self._validate()
|
||||
try:
|
||||
response = self._ragflow.retrieve(normalized_text, top_k=3)
|
||||
response = self._ragflow.retrieve(normalized_text, top_k=3, dataset_id=self._dataset_id)
|
||||
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
|
||||
return {"table_name": None, "raw": {"error": str(e)}}
|
||||
|
||||
candidates = []
|
||||
data = response.get("data") if isinstance(response, dict) else None
|
||||
@@ -58,7 +33,15 @@ class TemplateMatcher:
|
||||
candidates.append(table_name)
|
||||
|
||||
matched = candidates[0] if candidates else None
|
||||
result = {"table_name": matched, "raw": response}
|
||||
return {"table_name": matched, "raw": response}
|
||||
|
||||
self._cache.set(key, json.dumps(result, ensure_ascii=False), self._cache_ttl)
|
||||
return result
|
||||
|
||||
_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
|
||||
|
||||
Reference in New Issue
Block a user