init
This commit is contained in:
@@ -0,0 +1,29 @@
|
||||
# Services 模块
|
||||
|
||||
## 目录说明
|
||||
|
||||
`services` 负责底层能力封装,包括模型、检索、存储、工具路由和异常体系。
|
||||
|
||||
```
|
||||
services/
|
||||
├── core/ # LLM / Prompt / 表匹配 / SQL Prompt 管理
|
||||
├── integrations/ # RAGFlow / Nacos 等外部集成
|
||||
├── storage/ # 消息存储、日志、缓存
|
||||
├── tools/ # 工具路由
|
||||
├── common/ # 通用错误定义
|
||||
└── README.md
|
||||
```
|
||||
|
||||
## 核心职责
|
||||
|
||||
- 为 Agent 提供可复用的基础服务
|
||||
- 屏蔽外部系统交互细节
|
||||
- 提供统一的数据持久化与缓存能力
|
||||
|
||||
## 文件
|
||||
|
||||
- `core/`:模型构建、提示词加载、表匹配、SQL Prompt 管理
|
||||
- `integrations/`:RAGFlow 与 Nacos 外部接入
|
||||
- `storage/`:缓存、消息存储、结构化日志
|
||||
- `tools/`:工具路由
|
||||
- `common/`:错误码与应用异常
|
||||
@@ -0,0 +1,74 @@
|
||||
"""服务层模块 - 核心与扩展分离架构"""
|
||||
|
||||
# 核心服务
|
||||
from services.core import (
|
||||
create_chat_model,
|
||||
get_prompt_manager,
|
||||
get_sql_prompt_manager,
|
||||
get_template_matcher,
|
||||
)
|
||||
|
||||
# 外部集成
|
||||
from services.integrations import (
|
||||
RagflowClient,
|
||||
extract_table_name,
|
||||
RagflowSync,
|
||||
)
|
||||
|
||||
# 可选的 Nacos 导入
|
||||
try:
|
||||
from services.integrations import (
|
||||
NacosManager,
|
||||
NacosConfig,
|
||||
ServiceConfig,
|
||||
load_nacos_config,
|
||||
load_service_config,
|
||||
)
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# 数据存储
|
||||
from services.storage import (
|
||||
get_message_storage,
|
||||
MessageStorage,
|
||||
get_structured_logger,
|
||||
CacheBase,
|
||||
NoopCache,
|
||||
RedisCache,
|
||||
)
|
||||
|
||||
# 工具服务
|
||||
from services.tools import ToolRouter
|
||||
|
||||
# 基础设施
|
||||
from services.common import AppError, ErrorCode
|
||||
|
||||
__all__ = [
|
||||
# 核心服务
|
||||
"create_chat_model",
|
||||
"get_prompt_manager",
|
||||
"get_sql_prompt_manager",
|
||||
"get_template_matcher",
|
||||
# 外部集成
|
||||
"RagflowClient",
|
||||
"extract_table_name",
|
||||
"RagflowSync",
|
||||
# Nacos (可选)
|
||||
"NacosManager",
|
||||
"NacosConfig",
|
||||
"ServiceConfig",
|
||||
"load_nacos_config",
|
||||
"load_service_config",
|
||||
# 数据存储
|
||||
"get_message_storage",
|
||||
"MessageStorage",
|
||||
"get_structured_logger",
|
||||
"CacheBase",
|
||||
"NoopCache",
|
||||
"RedisCache",
|
||||
# 工具服务
|
||||
"ToolRouter",
|
||||
# 基础设施
|
||||
"AppError",
|
||||
"ErrorCode",
|
||||
]
|
||||
@@ -1,23 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class CacheBase:
|
||||
"""缓存接口"""
|
||||
|
||||
def get(self, key: str) -> Optional[str]:
|
||||
raise NotImplementedError
|
||||
|
||||
def set(self, key: str, value: str, ttl: int) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class NoopCache(CacheBase):
|
||||
"""空实现缓存"""
|
||||
|
||||
def get(self, key: str) -> Optional[str]:
|
||||
return None
|
||||
|
||||
def set(self, key: str, value: str, ttl: int) -> None:
|
||||
return None
|
||||
@@ -0,0 +1,6 @@
|
||||
"""基础设施模块"""
|
||||
|
||||
from .app_errors import AppError, ErrorCode
|
||||
from .datetime_utils import DateTimeBundle, DateTimeGenerator
|
||||
|
||||
__all__ = ["AppError", "ErrorCode", "DateTimeBundle", "DateTimeGenerator"]
|
||||
@@ -6,11 +6,17 @@ from typing import Any, Dict, Optional
|
||||
|
||||
|
||||
class ErrorCode(str, Enum):
|
||||
INVALID_REQUEST = "INVALID_REQUEST"
|
||||
INVALID_WORKFLOW_TYPE = "INVALID_WORKFLOW_TYPE"
|
||||
INVALID_RESPONSE_MODE = "INVALID_RESPONSE_MODE"
|
||||
SQL_GENERATION_FAILED = "SQL_GENERATION_FAILED"
|
||||
TABLE_MATCH_FAILED = "TABLE_MATCH_FAILED"
|
||||
SQL_EXECUTION_FAILED = "SQL_EXECUTION_FAILED"
|
||||
RAGFLOW_RETRIEVE_FAILED = "RAGFLOW_RETRIEVE_FAILED"
|
||||
CONVERSATION_NOT_FOUND = "CONVERSATION_NOT_FOUND"
|
||||
CONVERSATION_CREATE_FAILED = "CONVERSATION_CREATE_FAILED"
|
||||
CONVERSATION_UPDATE_FAILED = "CONVERSATION_UPDATE_FAILED"
|
||||
MESSAGE_SAVE_FAILED = "MESSAGE_SAVE_FAILED"
|
||||
CONFIG_INVALID = "CONFIG_INVALID"
|
||||
INTERNAL_ERROR = "INTERNAL_ERROR"
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime, time as dt_time
|
||||
from typing import Any
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DateTimeBundle:
|
||||
"""Unified datetime payload for storage and API usage."""
|
||||
|
||||
dt: datetime
|
||||
db_datetime: datetime
|
||||
epoch_seconds: int
|
||||
epoch_millis: int
|
||||
yyyymmdd: str
|
||||
date_str: str
|
||||
datetime_str: str
|
||||
iso_str: str
|
||||
|
||||
|
||||
class DateTimeGenerator:
|
||||
"""Parse and generate datetime values in multiple common formats."""
|
||||
|
||||
DEFAULT_TZ = ZoneInfo("Asia/Shanghai")
|
||||
SUPPORTED_FORMATS = (
|
||||
"%Y%m%d",
|
||||
"%Y%m%d%H%M%S",
|
||||
"%Y-%m-%d",
|
||||
"%Y/%m/%d",
|
||||
"%Y-%m-%d %H:%M",
|
||||
"%Y/%m/%d %H:%M",
|
||||
"%Y-%m-%d %H:%M:%S",
|
||||
"%Y/%m/%d %H:%M:%S",
|
||||
"%Y-%m-%d %H:%M:%S.%f",
|
||||
"%Y/%m/%d %H:%M:%S.%f",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def now(cls) -> DateTimeBundle:
|
||||
return cls.bundle()
|
||||
|
||||
@classmethod
|
||||
def bundle(cls, value: Any = None, *, default_to_now: bool = True) -> DateTimeBundle:
|
||||
dt = cls.parse(value, default_to_now=default_to_now)
|
||||
epoch_seconds = int(dt.timestamp())
|
||||
epoch_millis = int(dt.timestamp() * 1000)
|
||||
return DateTimeBundle(
|
||||
dt=dt,
|
||||
db_datetime=dt.replace(tzinfo=None),
|
||||
epoch_seconds=epoch_seconds,
|
||||
epoch_millis=epoch_millis,
|
||||
yyyymmdd=dt.strftime("%Y%m%d"),
|
||||
date_str=dt.strftime("%Y-%m-%d"),
|
||||
datetime_str=dt.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
iso_str=dt.isoformat(),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def pair(cls, created_value: Any = None, updated_value: Any = None) -> tuple[DateTimeBundle, DateTimeBundle]:
|
||||
created = cls.bundle(created_value, default_to_now=True)
|
||||
updated = cls.bundle(updated_value if updated_value is not None else created.epoch_millis, default_to_now=True)
|
||||
return created, updated
|
||||
|
||||
@classmethod
|
||||
def parse(cls, value: Any = None, *, default_to_now: bool = True) -> datetime:
|
||||
if value is None:
|
||||
if default_to_now:
|
||||
return datetime.now(cls.DEFAULT_TZ)
|
||||
raise ValueError("datetime value is None")
|
||||
|
||||
if isinstance(value, datetime):
|
||||
return cls._normalize_datetime(value)
|
||||
|
||||
if isinstance(value, date):
|
||||
return datetime.combine(value, dt_time.min).replace(tzinfo=cls.DEFAULT_TZ)
|
||||
|
||||
if isinstance(value, (int, float)):
|
||||
return cls._parse_numeric(str(int(value)), default_to_now=default_to_now)
|
||||
|
||||
if isinstance(value, str):
|
||||
text = value.strip()
|
||||
if not text:
|
||||
if default_to_now:
|
||||
return datetime.now(cls.DEFAULT_TZ)
|
||||
raise ValueError("datetime value is blank")
|
||||
|
||||
if text.isdigit():
|
||||
return cls._parse_numeric(text, default_to_now=default_to_now)
|
||||
|
||||
iso_candidate = text.replace("Z", "+00:00")
|
||||
try:
|
||||
return cls._normalize_datetime(datetime.fromisoformat(iso_candidate))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for fmt in cls.SUPPORTED_FORMATS:
|
||||
try:
|
||||
parsed = datetime.strptime(text, fmt)
|
||||
return parsed.replace(tzinfo=cls.DEFAULT_TZ)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
if default_to_now:
|
||||
return datetime.now(cls.DEFAULT_TZ)
|
||||
raise ValueError(f"unsupported datetime value: {value!r}")
|
||||
|
||||
@classmethod
|
||||
def _parse_numeric(cls, text: str, *, default_to_now: bool) -> datetime:
|
||||
if len(text) == 8:
|
||||
try:
|
||||
return datetime.strptime(text, "%Y%m%d").replace(tzinfo=cls.DEFAULT_TZ)
|
||||
except Exception:
|
||||
if default_to_now:
|
||||
return datetime.now(cls.DEFAULT_TZ)
|
||||
raise
|
||||
|
||||
if len(text) == 14:
|
||||
try:
|
||||
return datetime.strptime(text, "%Y%m%d%H%M%S").replace(tzinfo=cls.DEFAULT_TZ)
|
||||
except Exception:
|
||||
if default_to_now:
|
||||
return datetime.now(cls.DEFAULT_TZ)
|
||||
raise
|
||||
|
||||
if len(text) == 10:
|
||||
try:
|
||||
return datetime.fromtimestamp(int(text), tz=cls.DEFAULT_TZ)
|
||||
except Exception:
|
||||
if default_to_now:
|
||||
return datetime.now(cls.DEFAULT_TZ)
|
||||
raise
|
||||
|
||||
if len(text) == 13:
|
||||
try:
|
||||
return datetime.fromtimestamp(int(text) / 1000, tz=cls.DEFAULT_TZ)
|
||||
except Exception:
|
||||
if default_to_now:
|
||||
return datetime.now(cls.DEFAULT_TZ)
|
||||
raise
|
||||
|
||||
if default_to_now:
|
||||
return datetime.now(cls.DEFAULT_TZ)
|
||||
raise ValueError(f"unsupported numeric datetime value: {text!r}")
|
||||
|
||||
@classmethod
|
||||
def _normalize_datetime(cls, value: datetime) -> datetime:
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=cls.DEFAULT_TZ)
|
||||
return value.astimezone(cls.DEFAULT_TZ)
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
"""核心服务模块"""
|
||||
|
||||
from .llm_factory import create_chat_model
|
||||
from .prompt_manager import get_prompt_manager
|
||||
from .sql_prompt_manager import get_sql_prompt_manager
|
||||
from .template_matcher import get_template_matcher
|
||||
|
||||
__all__ = [
|
||||
"create_chat_model",
|
||||
"get_prompt_manager",
|
||||
"get_sql_prompt_manager",
|
||||
"get_template_matcher",
|
||||
]
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Optional
|
||||
from langchain_openai import ChatOpenAI
|
||||
from config import Config
|
||||
from config import Config, MAX_RETRIES, TIMEOUT
|
||||
|
||||
|
||||
def create_chat_model(model_section: Optional[str] = None) -> ChatOpenAI:
|
||||
@@ -11,6 +11,6 @@ def create_chat_model(model_section: Optional[str] = None) -> ChatOpenAI:
|
||||
api_key=model_config['api_key'],
|
||||
base_url=model_config.get('base_url'),
|
||||
temperature=0.1,
|
||||
max_retries=Config.MAX_RETRIES,
|
||||
timeout=Config.TIMEOUT
|
||||
max_retries=MAX_RETRIES,
|
||||
timeout=TIMEOUT
|
||||
)
|
||||
@@ -8,7 +8,7 @@ class PromptManager:
|
||||
"""提示词配置管理器"""
|
||||
|
||||
def __init__(self, config_path: Optional[str] = None):
|
||||
root_dir = os.path.dirname(os.path.dirname(__file__))
|
||||
root_dir = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
self._config_path = config_path or os.path.join(root_dir, "config", "prompts.yaml")
|
||||
self._data: Dict[str, Any] = {}
|
||||
self.reload()
|
||||
@@ -0,0 +1,238 @@
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from config import Config
|
||||
from services.storage.cache import CacheBase, NoopCache, RedisCache
|
||||
|
||||
|
||||
class SqlPromptManager:
|
||||
"""按表名读取 SQL 提示词,Redis 主存储 + 本地文件回退"""
|
||||
|
||||
KEY_PREFIX = "sql_prompt"
|
||||
TABLE_LIST_KEY = "sql_prompt:table_list"
|
||||
SOURCE_REDIS = "redis"
|
||||
SOURCE_FILE = "file"
|
||||
|
||||
def __init__(self, base_dir: Optional[str] = None):
|
||||
root_dir = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
self._fallback_dir = base_dir or os.path.join(root_dir, "config", "sql_gen_prompts")
|
||||
self._cache = self._init_cache()
|
||||
self._cache_ttl = self._get_cache_ttl()
|
||||
self._use_redis_primary = self._get_use_redis_primary()
|
||||
|
||||
@staticmethod
|
||||
def _first_line(value: Any) -> str:
|
||||
if value is None:
|
||||
return ""
|
||||
text = str(value)
|
||||
return text.splitlines()[0].strip() if text else ""
|
||||
|
||||
@classmethod
|
||||
def _to_bool(cls, value: Any, default: bool = False) -> bool:
|
||||
text = cls._first_line(value).lower()
|
||||
if not text:
|
||||
return default
|
||||
return text in ("1", "true", "yes", "y", "on")
|
||||
|
||||
@classmethod
|
||||
def _to_int(cls, value: Any, default: int = 0) -> int:
|
||||
text = cls._first_line(value)
|
||||
if not text:
|
||||
return default
|
||||
try:
|
||||
return int(text)
|
||||
except Exception:
|
||||
return default
|
||||
|
||||
@classmethod
|
||||
def _extract_key_from_multiline_values(cls, redis_cfg: Dict[str, Any], target_key: str) -> str:
|
||||
token = f"{target_key}="
|
||||
for raw in redis_cfg.values():
|
||||
text = str(raw or "")
|
||||
for line in text.splitlines()[1:]:
|
||||
cleaned = line.strip()
|
||||
normalized = cleaned.replace(" ", "")
|
||||
if normalized.lower().startswith(token.lower()):
|
||||
return cleaned.split("=", 1)[1].strip()
|
||||
return ""
|
||||
|
||||
@classmethod
|
||||
def _get_cfg_value(cls, redis_cfg: Dict[str, Any], key: str, default: Any = "") -> Any:
|
||||
if key in redis_cfg:
|
||||
return redis_cfg.get(key, default)
|
||||
recovered = cls._extract_key_from_multiline_values(redis_cfg, key)
|
||||
return recovered if recovered else default
|
||||
|
||||
@classmethod
|
||||
def _get_cache_ttl(cls) -> Optional[int]:
|
||||
redis_cfg = Config.get_section("redis")
|
||||
ttl = cls._to_int(cls._get_cfg_value(redis_cfg, "sql_prompt_ttl", 0), default=0)
|
||||
if ttl:
|
||||
return ttl
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def _get_use_redis_primary(cls) -> bool:
|
||||
redis_cfg = Config.get_section("redis")
|
||||
enabled = cls._to_bool(cls._get_cfg_value(redis_cfg, "enabled", "false"), default=False)
|
||||
primary = cls._to_bool(cls._get_cfg_value(redis_cfg, "sql_prompt_redis_primary", "false"), default=False)
|
||||
return enabled and primary
|
||||
|
||||
@classmethod
|
||||
def _init_cache(cls):
|
||||
redis_cfg = Config.get_section("redis")
|
||||
enabled = cls._to_bool(cls._get_cfg_value(redis_cfg, "enabled", "false"), default=False)
|
||||
if not enabled:
|
||||
return NoopCache()
|
||||
|
||||
url = cls._first_line(cls._get_cfg_value(redis_cfg, "url"))
|
||||
db = cls._to_int(cls._get_cfg_value(redis_cfg, "db", cls._get_cfg_value(redis_cfg, "database", 0)), default=0)
|
||||
if not url:
|
||||
host = cls._first_line(cls._get_cfg_value(redis_cfg, "host"))
|
||||
port = cls._first_line(cls._get_cfg_value(redis_cfg, "port", "6379")) or "6379"
|
||||
password = cls._first_line(cls._get_cfg_value(redis_cfg, "password", ""))
|
||||
username = cls._first_line(cls._get_cfg_value(redis_cfg, "username", ""))
|
||||
database = cls._first_line(cls._get_cfg_value(redis_cfg, "database", str(db))) or str(db)
|
||||
if host:
|
||||
from urllib.parse import quote_plus
|
||||
if username and password:
|
||||
auth = f"{quote_plus(username)}:{quote_plus(password)}@"
|
||||
elif password:
|
||||
auth = f":{quote_plus(password)}@"
|
||||
else:
|
||||
auth = ""
|
||||
url = f"redis://{auth}{host}:{port}/{database}"
|
||||
|
||||
if not url:
|
||||
return NoopCache()
|
||||
try:
|
||||
return RedisCache(url=url, db=db)
|
||||
except Exception:
|
||||
return NoopCache()
|
||||
|
||||
@staticmethod
|
||||
def _safe_filename(name: str) -> str:
|
||||
return name.replace("..", "").replace("/", "_").replace("\\", "_")
|
||||
|
||||
def _redis_key(self, table_name: str) -> str:
|
||||
return f"{self.KEY_PREFIX}:{table_name}"
|
||||
|
||||
def get_prompt(self, table_name: str) -> Optional[Dict[str, Any]]:
|
||||
"""读取指定表的提示词,优先 Redis,回退本地文件"""
|
||||
if not table_name:
|
||||
return None
|
||||
|
||||
if self._use_redis_primary:
|
||||
prompt = self._get_from_redis(table_name)
|
||||
if prompt:
|
||||
return prompt
|
||||
|
||||
prompt = self._get_from_file(table_name)
|
||||
return prompt
|
||||
|
||||
def get_prompt_with_source(self, table_name: str) -> tuple[Optional[Dict[str, Any]], str]:
|
||||
"""读取指定表的提示词,返回 (prompt, source) 元组"""
|
||||
if not table_name:
|
||||
return None, self.SOURCE_FILE
|
||||
|
||||
if self._use_redis_primary:
|
||||
prompt = self._get_from_redis(table_name)
|
||||
if prompt:
|
||||
return prompt, self.SOURCE_REDIS
|
||||
|
||||
prompt = self._get_from_file(table_name)
|
||||
source = self.SOURCE_FILE if prompt else self.SOURCE_FILE
|
||||
return prompt, source
|
||||
|
||||
def _get_from_redis(self, table_name: str) -> Optional[Dict[str, Any]]:
|
||||
key = self._redis_key(table_name)
|
||||
try:
|
||||
data = self._cache.get(key)
|
||||
if data:
|
||||
try:
|
||||
return json.loads(data)
|
||||
except Exception:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
def _get_from_file(self, table_name: str) -> Optional[Dict[str, Any]]:
|
||||
safe_name = self._safe_filename(table_name)
|
||||
filename = safe_name + ".json"
|
||||
path = os.path.join(self._fallback_dir, filename)
|
||||
if not os.path.exists(path):
|
||||
return None
|
||||
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
prompt = json.load(f)
|
||||
return prompt
|
||||
|
||||
def save_prompt(self, table_name: str, prompt: Dict[str, Any]) -> bool:
|
||||
"""保存提示词到 Redis"""
|
||||
if not table_name or not prompt:
|
||||
return False
|
||||
|
||||
key = self._redis_key(table_name)
|
||||
try:
|
||||
self._cache.set(key, json.dumps(prompt, ensure_ascii=False), self._cache_ttl)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def delete_prompt(self, table_name: str) -> bool:
|
||||
"""从 Redis 删除提示词"""
|
||||
if not table_name:
|
||||
return False
|
||||
|
||||
key = self._redis_key(table_name)
|
||||
try:
|
||||
self._cache.delete(key)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def list_tables(self) -> List[str]:
|
||||
"""列出 Redis 中所有表名"""
|
||||
pattern = f"{self.KEY_PREFIX}:*"
|
||||
keys = self._cache.keys(pattern)
|
||||
tables = []
|
||||
for key in keys:
|
||||
if key == self.TABLE_LIST_KEY:
|
||||
continue
|
||||
parts = key.split(":", 1)
|
||||
if len(parts) == 2:
|
||||
tables.append(parts[1])
|
||||
return tables
|
||||
|
||||
def sync_from_files(self, tables: Optional[List[str]] = None) -> Dict[str, bool]:
|
||||
"""从本地文件同步到 Redis"""
|
||||
results: Dict[str, bool] = {}
|
||||
|
||||
if tables:
|
||||
files_to_sync = [f"{self._safe_filename(t)}.json" for t in tables]
|
||||
else:
|
||||
try:
|
||||
files_to_sync = [f for f in os.listdir(self._fallback_dir) if f.endswith(".json")]
|
||||
except Exception:
|
||||
return results
|
||||
|
||||
for filename in files_to_sync:
|
||||
table_name = filename[:-5]
|
||||
prompt = self._get_from_file(table_name)
|
||||
if prompt:
|
||||
results[table_name] = self.save_prompt(table_name, prompt)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
_GLOBAL_SQL_PROMPT_MANAGER: Optional[SqlPromptManager] = None
|
||||
|
||||
|
||||
def get_sql_prompt_manager(base_dir: Optional[str] = None) -> SqlPromptManager:
|
||||
"""获取全局 SqlPromptManager(单例)"""
|
||||
global _GLOBAL_SQL_PROMPT_MANAGER
|
||||
if _GLOBAL_SQL_PROMPT_MANAGER is None:
|
||||
_GLOBAL_SQL_PROMPT_MANAGER = SqlPromptManager(base_dir=base_dir)
|
||||
return _GLOBAL_SQL_PROMPT_MANAGER
|
||||
@@ -0,0 +1,201 @@
|
||||
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
|
||||
@@ -0,0 +1,24 @@
|
||||
"""外部集成模块"""
|
||||
|
||||
from .ragflow_client import RagflowClient, extract_table_name
|
||||
from .ragflow_sync import RagflowSync
|
||||
|
||||
# 延迟导入 nacos(可选依赖)
|
||||
try:
|
||||
from .nacos_service import NacosManager, NacosConfig, ServiceConfig, load_nacos_config, load_service_config
|
||||
__all__ = [
|
||||
"RagflowClient",
|
||||
"extract_table_name",
|
||||
"RagflowSync",
|
||||
"NacosManager",
|
||||
"NacosConfig",
|
||||
"ServiceConfig",
|
||||
"load_nacos_config",
|
||||
"load_service_config",
|
||||
]
|
||||
except ImportError:
|
||||
__all__ = [
|
||||
"RagflowClient",
|
||||
"extract_table_name",
|
||||
"RagflowSync",
|
||||
]
|
||||
@@ -1,9 +1,15 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import socket
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
import nacos
|
||||
|
||||
try:
|
||||
import nacos # type: ignore
|
||||
except ImportError:
|
||||
nacos = None
|
||||
|
||||
from config import Config
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -22,6 +28,7 @@ class NacosConfig:
|
||||
heartbeat_interval: int
|
||||
weight: float
|
||||
ephemeral: bool
|
||||
register_port: Optional[int]
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -36,21 +43,34 @@ class ServiceConfig:
|
||||
|
||||
def _get_local_ip() -> str:
|
||||
"""获取本地 IP 地址"""
|
||||
for env_name in ("POD_IP", "HOST_IP"):
|
||||
env_ip = os.getenv(env_name)
|
||||
if env_ip:
|
||||
return env_ip
|
||||
|
||||
try:
|
||||
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||
s.connect(("8.8.8.8", 80))
|
||||
ip = s.getsockname()[0]
|
||||
s.close()
|
||||
return ip
|
||||
except Exception as e:
|
||||
logger.warning(f"获取本地 IP 失败,使用 127.0.0.1: {e}")
|
||||
return "127.0.0.1"
|
||||
except Exception:
|
||||
try:
|
||||
host_ip = socket.gethostbyname(socket.gethostname())
|
||||
if host_ip and host_ip != "127.0.0.1":
|
||||
return host_ip
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logger.warning("获取本地 IP 失败,使用 127.0.0.1")
|
||||
return "127.0.0.1"
|
||||
|
||||
|
||||
def load_nacos_config() -> NacosConfig:
|
||||
"""从 config.ini 读取 Nacos 配置"""
|
||||
section = "nacos"
|
||||
enabled = Config._config.getboolean(section, "enabled", fallback=False)
|
||||
register_port = Config._config.getint(section, "register_port", fallback=0)
|
||||
return NacosConfig(
|
||||
enabled=enabled,
|
||||
server_addresses=Config._config.get(section, "server", fallback="localhost:8848"),
|
||||
@@ -62,22 +82,23 @@ def load_nacos_config() -> NacosConfig:
|
||||
heartbeat_interval=Config._config.getint(section, "heartbeat_interval", fallback=5),
|
||||
weight=Config._config.getfloat(section, "weight", fallback=1.0),
|
||||
ephemeral=Config._config.getboolean(section, "ephemeral", fallback=True),
|
||||
register_port=register_port if register_port > 0 else None,
|
||||
)
|
||||
|
||||
|
||||
def load_service_config() -> ServiceConfig:
|
||||
"""从 config.ini 读取服务配置"""
|
||||
section = "app"
|
||||
service_name = Config._config.get(section, "service_name", fallback="more-dots-api")
|
||||
service_name = Config._config.get(section, "service_name", fallback="apbo-boat-agent")
|
||||
host = Config._config.get(section, "host", fallback="0.0.0.0")
|
||||
port = Config._config.getint(section, "port", fallback=8000)
|
||||
ip = host if host != "0.0.0.0" else _get_local_ip()
|
||||
ip = host if host not in ("0.0.0.0", "::") else _get_local_ip()
|
||||
|
||||
metadata = {
|
||||
"version": Config._config.get(section, "version", fallback="1.0.0"),
|
||||
"service_type": "fastapi",
|
||||
"api_paths": "/health,/api/workflows,/api/workflows/stream,/nacos/status",
|
||||
"streaming": "false",
|
||||
"streaming": "true",
|
||||
"model_section": Config._config.get(section, "model_section", fallback=Config.DEFAULT_MODEL_SECTION),
|
||||
}
|
||||
|
||||
@@ -105,6 +126,10 @@ class NacosManager:
|
||||
self._stop_event = asyncio.Event()
|
||||
self.is_registered = False
|
||||
|
||||
def _registration_port(self) -> int:
|
||||
"""Nacos 注册端口:优先使用 nacos.register_port,未配置时回退 app.port。"""
|
||||
return int(self.nacos_config.register_port or self.service_config.port)
|
||||
|
||||
def _init_client(self) -> bool:
|
||||
"""初始化 Nacos 客户端"""
|
||||
if nacos is None:
|
||||
@@ -132,7 +157,7 @@ class NacosManager:
|
||||
self.client.add_naming_instance(
|
||||
service_name=self.service_config.service_name,
|
||||
ip=self.service_config.ip,
|
||||
port=self.service_config.port,
|
||||
port=self._registration_port(),
|
||||
cluster_name=self.nacos_config.cluster_name,
|
||||
group_name=self.nacos_config.group_name,
|
||||
weight=self.nacos_config.weight,
|
||||
@@ -144,7 +169,7 @@ class NacosManager:
|
||||
"✅ 服务注册成功: %s (%s:%s)",
|
||||
self.service_config.service_name,
|
||||
self.service_config.ip,
|
||||
self.service_config.port,
|
||||
self._registration_port(),
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
@@ -161,7 +186,7 @@ class NacosManager:
|
||||
self.client.remove_naming_instance(
|
||||
service_name=self.service_config.service_name,
|
||||
ip=self.service_config.ip,
|
||||
port=self.service_config.port,
|
||||
port=self._registration_port(),
|
||||
cluster_name=self.nacos_config.cluster_name,
|
||||
group_name=self.nacos_config.group_name,
|
||||
)
|
||||
@@ -180,7 +205,7 @@ class NacosManager:
|
||||
self.client.send_heartbeat(
|
||||
service_name=self.service_config.service_name,
|
||||
ip=self.service_config.ip,
|
||||
port=self.service_config.port,
|
||||
port=self._registration_port(),
|
||||
cluster_name=self.nacos_config.cluster_name,
|
||||
group_name=self.nacos_config.group_name,
|
||||
)
|
||||
@@ -192,8 +217,14 @@ class NacosManager:
|
||||
|
||||
while not self._stop_event.is_set():
|
||||
try:
|
||||
self._send_heartbeat()
|
||||
logger.debug("心跳发送成功: %s", self.service_config.service_name)
|
||||
if not self.is_registered:
|
||||
if self.register_service():
|
||||
logger.info("✅ Nacos 重试注册成功: %s", self.service_config.service_name)
|
||||
else:
|
||||
logger.warning("⚠️ Nacos 注册重试失败: %s", self.service_config.service_name)
|
||||
else:
|
||||
self._send_heartbeat()
|
||||
logger.debug("心跳发送成功: %s", self.service_config.service_name)
|
||||
except Exception as e:
|
||||
logger.warning(f"心跳发送失败: {e}")
|
||||
# 尝试重新注册
|
||||
@@ -213,12 +244,13 @@ class NacosManager:
|
||||
logger.info("Nacos 未启用,跳过注册")
|
||||
return
|
||||
|
||||
if self.register_service():
|
||||
self._heartbeat_task = asyncio.create_task(self._heartbeat_loop())
|
||||
logger.info("✅ Nacos 心跳任务已启动")
|
||||
else:
|
||||
if not self.register_service():
|
||||
logger.warning("⚠️ Nacos 注册失败,服务继续运行")
|
||||
|
||||
# 无论首次注册是否成功,都启动循环以便持续重试注册
|
||||
self._heartbeat_task = asyncio.create_task(self._heartbeat_loop())
|
||||
logger.info("✅ Nacos 心跳任务已启动")
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""停止心跳并注销"""
|
||||
self._stop_event.set()
|
||||
@@ -238,6 +270,7 @@ class NacosManager:
|
||||
"service_name": self.service_config.service_name,
|
||||
"ip": self.service_config.ip,
|
||||
"port": self.service_config.port,
|
||||
"register_port": self._registration_port(),
|
||||
"namespace": self.nacos_config.namespace,
|
||||
"group": self.nacos_config.group_name,
|
||||
"cluster": self.nacos_config.cluster_name,
|
||||
@@ -6,6 +6,11 @@ import httpx
|
||||
from config import Config
|
||||
|
||||
|
||||
TABLE_NAME_ALIASES = {
|
||||
"apbo_tp_multiple_impact": "apbo_eta_multiple_impact",
|
||||
}
|
||||
|
||||
|
||||
class RagflowClient:
|
||||
"""RAGFlow 客户端(仅检索)"""
|
||||
|
||||
@@ -55,17 +60,23 @@ class RagflowClient:
|
||||
|
||||
def extract_table_name(record: Dict[str, Any]) -> Optional[str]:
|
||||
"""从检索结果中提取表名"""
|
||||
def _canonicalize(table_name: Any) -> Optional[str]:
|
||||
normalized = str(table_name or "").strip()
|
||||
if not normalized:
|
||||
return None
|
||||
return TABLE_NAME_ALIASES.get(normalized, normalized)
|
||||
|
||||
if not record:
|
||||
return None
|
||||
|
||||
metadata = record.get("metadata") or {}
|
||||
for key in ("table", "table_name"):
|
||||
if key in metadata:
|
||||
return metadata.get(key)
|
||||
return _canonicalize(metadata.get(key))
|
||||
|
||||
for key in ("table", "table_name"):
|
||||
if key in record:
|
||||
return record.get(key)
|
||||
return _canonicalize(record.get(key))
|
||||
|
||||
content = record.get("content") or record.get("text") or ""
|
||||
|
||||
@@ -75,12 +86,12 @@ def extract_table_name(record: Dict[str, Any]) -> Optional[str]:
|
||||
if isinstance(parsed, dict):
|
||||
for key in ("table", "table_name"):
|
||||
if parsed.get(key):
|
||||
return str(parsed.get(key))
|
||||
return _canonicalize(parsed.get(key))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for line in str(content).splitlines():
|
||||
if line.lower().startswith("table:"):
|
||||
return line.split(":", 1)[1].strip()
|
||||
return _canonicalize(line.split(":", 1)[1].strip())
|
||||
|
||||
return None
|
||||
@@ -1,6 +1,7 @@
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, List
|
||||
from pathlib import Path
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -31,6 +32,38 @@ class RagflowSync:
|
||||
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()
|
||||
|
||||
@staticmethod
|
||||
def _project_root() -> Path:
|
||||
"""返回项目根目录。"""
|
||||
return Path(__file__).resolve().parents[2]
|
||||
|
||||
def _collect_sql_gen_documents(self) -> tuple[List[Dict[str, Any]], List[Dict[str, str]]]:
|
||||
"""收集 SQL 生成文档,并跳过空文件/非法 JSON。"""
|
||||
prompts_dir = self._project_root() / "config" / "sql_gen_prompts"
|
||||
documents: List[Dict[str, Any]] = []
|
||||
warnings: List[Dict[str, str]] = []
|
||||
|
||||
for name in os.listdir(prompts_dir):
|
||||
if not name.endswith(".json"):
|
||||
continue
|
||||
|
||||
path = prompts_dir / name
|
||||
raw_text = path.read_text(encoding="utf-8")
|
||||
if not raw_text.strip():
|
||||
warnings.append({"file": name, "reason": "empty_file"})
|
||||
continue
|
||||
|
||||
try:
|
||||
prompt = json.loads(raw_text)
|
||||
except json.JSONDecodeError as exc:
|
||||
warnings.append({"file": name, "reason": f"invalid_json:{exc}"})
|
||||
continue
|
||||
|
||||
table = prompt.get("table") or path.stem
|
||||
documents.append({"filename": f"{table}.txt", "content": _dump_json_content(prompt)})
|
||||
|
||||
return documents, warnings
|
||||
|
||||
def upload_documents(self, dataset_id: str, documents: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
"""上传文档到指定知识库(每个文档单独上传)"""
|
||||
if not self._base_url:
|
||||
@@ -262,8 +295,7 @@ class RagflowSync:
|
||||
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")
|
||||
tables_file = self._project_root() / "config" / "table_retrieval_prompts" / "tables.json"
|
||||
with open(tables_file, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
|
||||
@@ -277,34 +309,35 @@ class RagflowSync:
|
||||
]
|
||||
return self.replace_documents(self._table_retrieval_dataset_id, documents)
|
||||
|
||||
def sync_table_retrieval(self) -> Dict[str, Any]:
|
||||
"""兼容旧脚本:同步表名检索文档,采用覆盖更新避免旧表残留。"""
|
||||
return self.update_table_retrieval_documents()
|
||||
|
||||
def update_sql_gen_documents(self) -> Dict[str, Any]:
|
||||
"""更新 SQL 生成文档(仅文档内容)"""
|
||||
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: List[Dict[str, Any]] = []
|
||||
documents, warnings = self._collect_sql_gen_documents()
|
||||
if not documents:
|
||||
raise RuntimeError(f"SQL 生成提示词目录中没有可同步的有效 JSON 文档,warnings={warnings}")
|
||||
|
||||
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]
|
||||
documents.append({"filename": f"{table}.txt", "content": _dump_json_content(prompt)})
|
||||
result = self.replace_documents(self._sql_gen_dataset_id, documents)
|
||||
result["warnings"] = warnings
|
||||
result["valid_document_count"] = len(documents)
|
||||
return result
|
||||
|
||||
return self.replace_documents(self._sql_gen_dataset_id, documents)
|
||||
def sync_sql_gen_prompts(self) -> Dict[str, Any]:
|
||||
"""兼容旧脚本:同步 SQL 生成提示词文档,采用覆盖更新避免旧 prompt 残留。"""
|
||||
return self.update_sql_gen_documents()
|
||||
|
||||
def upload_table_retrieval(self) -> Dict[str, Any]:
|
||||
"""上传表名检索模板文档 - 直接上传整个 JSON 文件"""
|
||||
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")
|
||||
|
||||
tables_file = self._project_root() / "config" / "table_retrieval_prompts" / "tables.json"
|
||||
|
||||
if not os.path.exists(tables_file):
|
||||
raise RuntimeError(f"表名检索模板文件不存在: {tables_file}")
|
||||
|
||||
@@ -328,22 +361,16 @@ class RagflowSync:
|
||||
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")
|
||||
|
||||
prompts_dir = self._project_root() / "config" / "sql_gen_prompts"
|
||||
|
||||
if not os.path.exists(prompts_dir):
|
||||
raise RuntimeError(f"SQL 生成提示词目录不存在: {prompts_dir}")
|
||||
|
||||
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]
|
||||
json_content = _dump_json_content(prompt)
|
||||
documents.append({"filename": f"{table}.txt", "content": json_content})
|
||||
documents, warnings = self._collect_sql_gen_documents()
|
||||
if not documents:
|
||||
raise RuntimeError(f"SQL 生成提示词目录中没有可上传的有效 JSON 文档,warnings={warnings}")
|
||||
|
||||
return self.upload_documents(self._sql_gen_dataset_id, documents)
|
||||
result = self.upload_documents(self._sql_gen_dataset_id, documents)
|
||||
result["warnings"] = warnings
|
||||
result["valid_document_count"] = len(documents)
|
||||
return result
|
||||
@@ -1,41 +0,0 @@
|
||||
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
|
||||
safe_name = self._safe_filename(table_name)
|
||||
filename = safe_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:
|
||||
prompt = json.load(f)
|
||||
|
||||
return prompt
|
||||
|
||||
|
||||
_GLOBAL_SQL_PROMPT_MANAGER: Optional[SqlPromptManager] = None
|
||||
|
||||
|
||||
def get_sql_prompt_manager(base_dir: Optional[str] = None) -> SqlPromptManager:
|
||||
"""获取全局 SqlPromptManager(单例)"""
|
||||
global _GLOBAL_SQL_PROMPT_MANAGER
|
||||
if _GLOBAL_SQL_PROMPT_MANAGER is None:
|
||||
_GLOBAL_SQL_PROMPT_MANAGER = SqlPromptManager(base_dir=base_dir)
|
||||
return _GLOBAL_SQL_PROMPT_MANAGER
|
||||
@@ -0,0 +1,14 @@
|
||||
"""数据存储模块"""
|
||||
|
||||
from .message_storage import get_message_storage, MessageStorage
|
||||
from .structured_logger import get_structured_logger
|
||||
from .cache import CacheBase, NoopCache, RedisCache
|
||||
|
||||
__all__ = [
|
||||
"get_message_storage",
|
||||
"MessageStorage",
|
||||
"get_structured_logger",
|
||||
"CacheBase",
|
||||
"NoopCache",
|
||||
"RedisCache",
|
||||
]
|
||||
@@ -0,0 +1,64 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
try:
|
||||
import redis
|
||||
except Exception:
|
||||
redis = None
|
||||
|
||||
|
||||
class CacheBase:
|
||||
"""缓存接口"""
|
||||
|
||||
def get(self, key: str) -> Optional[str]:
|
||||
raise NotImplementedError
|
||||
|
||||
def set(self, key: str, value: str, ttl: int = None) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def delete(self, key: str) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def keys(self, pattern: str) -> List[str]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class NoopCache(CacheBase):
|
||||
"""空实现缓存"""
|
||||
|
||||
def get(self, key: str) -> Optional[str]:
|
||||
return None
|
||||
|
||||
def set(self, key: str, value: str, ttl: int = None) -> None:
|
||||
return None
|
||||
|
||||
def delete(self, key: str) -> None:
|
||||
return None
|
||||
|
||||
def keys(self, pattern: str) -> List[str]:
|
||||
return []
|
||||
|
||||
|
||||
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) -> None:
|
||||
if ttl:
|
||||
self._client.set(key, value, ex=ttl)
|
||||
else:
|
||||
self._client.set(key, value)
|
||||
|
||||
def delete(self, key: str) -> None:
|
||||
self._client.delete(key)
|
||||
|
||||
def keys(self, pattern: str) -> List[str]:
|
||||
return self._client.keys(pattern)
|
||||
@@ -0,0 +1,885 @@
|
||||
"""
|
||||
消息存储服务 - 将每次查询的消息记录存储到 MySQL
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import pymysql
|
||||
|
||||
from config import Config
|
||||
from services.common.datetime_utils import DateTimeGenerator
|
||||
from schemas.messages import MessagesDTO
|
||||
|
||||
|
||||
class MessageStorage:
|
||||
"""消息存储服务"""
|
||||
|
||||
def __init__(self):
|
||||
cfg = Config.get_section("logging_mysql")
|
||||
self.enabled = str(cfg.get("enabled", "false")).lower() in ("1", "true", "yes")
|
||||
self.entity_debug_enabled = str(cfg.get("entity_debug_enabled", "false")).lower() in ("1", "true", "yes")
|
||||
self.host = cfg.get("host", "127.0.0.1")
|
||||
self.port = int(cfg.get("port", 3306))
|
||||
self.user = cfg.get("user", "root")
|
||||
self.password = cfg.get("password", "")
|
||||
self.database = cfg.get("database", "more_dots")
|
||||
# 消息落库与结构化日志分表,避免误用 logging_mysql.table=structured_logs
|
||||
self.table = cfg.get("messages_table", "ipc_apbo.messages")
|
||||
self.conversation_table = cfg.get("conversation_table", "ipc_apbo.conversations")
|
||||
self.connect_timeout = int(cfg.get("connect_timeout", 5))
|
||||
self._inited = False
|
||||
self._conversation_schema_checked = False
|
||||
|
||||
def _get_conn(self):
|
||||
"""获取数据库连接"""
|
||||
return pymysql.connect(
|
||||
host=self.host,
|
||||
port=self.port,
|
||||
user=self.user,
|
||||
password=self.password,
|
||||
database=self.database,
|
||||
charset="utf8mb4",
|
||||
autocommit=True,
|
||||
connect_timeout=self.connect_timeout,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _log_local(event: str, payload: Optional[Dict[str, Any]] = None) -> None:
|
||||
print(json.dumps({
|
||||
"level": "ERROR",
|
||||
"event": event,
|
||||
"payload": payload or {},
|
||||
"created_at": DateTimeGenerator.now().epoch_millis,
|
||||
}, ensure_ascii=False))
|
||||
|
||||
def _log_entity_debug(self, event: str, payload: Optional[Dict[str, Any]] = None) -> None:
|
||||
if not self.entity_debug_enabled:
|
||||
return
|
||||
print(json.dumps({
|
||||
"level": "DEBUG",
|
||||
"event": event,
|
||||
"payload": payload or {},
|
||||
"created_at": DateTimeGenerator.now().epoch_millis,
|
||||
}, ensure_ascii=False))
|
||||
|
||||
@staticmethod
|
||||
def _now_ms() -> int:
|
||||
return DateTimeGenerator.now().epoch_millis
|
||||
|
||||
@staticmethod
|
||||
def _build_audit_fields(
|
||||
*,
|
||||
random_code: str,
|
||||
user: Optional[str],
|
||||
created_value: Any = None,
|
||||
updated_value: Any = None,
|
||||
) -> Dict[str, Any]:
|
||||
created_ms = MessageStorage._resolve_epoch_millis(created_value)
|
||||
updated_ms = MessageStorage._resolve_epoch_millis(updated_value, fallback=created_ms)
|
||||
created_bundle = DateTimeGenerator.bundle(created_ms, default_to_now=True)
|
||||
updated_bundle = DateTimeGenerator.bundle(updated_ms, default_to_now=True)
|
||||
operator = (user or "system").strip() if isinstance(user, str) else "system"
|
||||
return {
|
||||
"random_code": random_code,
|
||||
"create_user": operator,
|
||||
"create_date": created_bundle.db_datetime,
|
||||
"update_user": operator,
|
||||
"update_date": updated_bundle.db_datetime,
|
||||
"create_user_name": operator,
|
||||
"update_user_name": operator,
|
||||
"created_at": created_ms,
|
||||
"updated_at": updated_ms,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _resolve_epoch_millis(value: Any, fallback: Optional[int] = None) -> int:
|
||||
if value is None:
|
||||
return fallback if fallback is not None else DateTimeGenerator.now().epoch_millis
|
||||
|
||||
if isinstance(value, (int, float)):
|
||||
raw = int(value)
|
||||
digits = len(str(abs(raw)))
|
||||
if digits == 10:
|
||||
return raw * 1000
|
||||
return raw
|
||||
|
||||
if isinstance(value, str):
|
||||
text = value.strip()
|
||||
if text.isdigit():
|
||||
raw = int(text)
|
||||
digits = len(text)
|
||||
if digits == 10:
|
||||
return raw * 1000
|
||||
if digits == 13:
|
||||
return raw
|
||||
|
||||
parsed = DateTimeGenerator.bundle(value, default_to_now=True)
|
||||
return parsed.epoch_millis
|
||||
|
||||
parsed = DateTimeGenerator.bundle(value, default_to_now=True)
|
||||
return parsed.epoch_millis
|
||||
|
||||
def _log_entity_stage(self, entity: str, action: str, stage: str, started_at: int, payload: Optional[Dict[str, Any]] = None) -> None:
|
||||
debug_payload = dict(payload or {})
|
||||
debug_payload.update({
|
||||
"entity": entity,
|
||||
"action": action,
|
||||
"stage": stage,
|
||||
"elapsed_ms": max(0, self._now_ms() - started_at),
|
||||
})
|
||||
self._log_entity_debug(f"message_storage.{entity}.{action}.{stage}", debug_payload)
|
||||
|
||||
@staticmethod
|
||||
def _split_table_reference(table_name: str, default_schema: str) -> tuple[str, str]:
|
||||
cleaned = str(table_name or "").strip()
|
||||
if "." in cleaned:
|
||||
schema_name, physical_table_name = cleaned.split(".", 1)
|
||||
else:
|
||||
schema_name, physical_table_name = default_schema, cleaned
|
||||
return schema_name.strip().strip("`"), physical_table_name.strip().strip("`")
|
||||
|
||||
def _ensure_conversation_schema(self, conn) -> None:
|
||||
if self._conversation_schema_checked or not self.enabled:
|
||||
return
|
||||
|
||||
started_at = self._now_ms()
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"ensure_schema",
|
||||
"start",
|
||||
started_at,
|
||||
{"conversation_table": self.conversation_table},
|
||||
)
|
||||
|
||||
schema_name, table_name = self._split_table_reference(self.conversation_table, self.database)
|
||||
probe_sql = """
|
||||
SELECT 1
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = %s AND table_name = %s AND column_name = %s
|
||||
LIMIT 1
|
||||
"""
|
||||
|
||||
with conn.cursor() as cur:
|
||||
cur.execute(probe_sql, (schema_name, table_name, "name"))
|
||||
if cur.fetchone() is None:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"ensure_schema",
|
||||
"alter_needed",
|
||||
started_at,
|
||||
{"conversation_table": self.conversation_table, "missing_column": "name"},
|
||||
)
|
||||
alter_sql = f"""
|
||||
ALTER TABLE {self.conversation_table}
|
||||
ADD COLUMN name VARCHAR(255) NULL COMMENT '会话名称' AFTER user
|
||||
"""
|
||||
cur.execute(alter_sql)
|
||||
|
||||
self._conversation_schema_checked = True
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"ensure_schema",
|
||||
"success",
|
||||
started_at,
|
||||
{"conversation_table": self.conversation_table, "schema_checked": True},
|
||||
)
|
||||
|
||||
def _ensure_table(self) -> None:
|
||||
"""确保消息表存在"""
|
||||
if self._inited or not self.enabled:
|
||||
return
|
||||
|
||||
started_at = self._now_ms()
|
||||
self._log_entity_stage(
|
||||
"storage",
|
||||
"ensure_table",
|
||||
"start",
|
||||
started_at,
|
||||
{"messages_table": self.table, "conversation_table": self.conversation_table},
|
||||
)
|
||||
|
||||
message_sql = f"""
|
||||
CREATE TABLE IF NOT EXISTS {self.table} (
|
||||
id BIGINT PRIMARY KEY AUTO_INCREMENT,
|
||||
random_code VARCHAR(100) NULL COMMENT '业务主键',
|
||||
create_user VARCHAR(100) NULL COMMENT '创建人',
|
||||
create_date DATETIME NULL COMMENT '创建时间',
|
||||
update_user VARCHAR(100) NULL COMMENT '修改人',
|
||||
update_date DATETIME NULL COMMENT '修改时间',
|
||||
create_user_name VARCHAR(255) NULL COMMENT '创建人姓名',
|
||||
update_user_name VARCHAR(255) NULL COMMENT '修改人姓名',
|
||||
message_id VARCHAR(255) NULL COMMENT '消息 ID',
|
||||
conversation_id VARCHAR(255) NULL COMMENT '会话 ID',
|
||||
user VARCHAR(255) NULL COMMENT '用户标识',
|
||||
query LONGTEXT NULL COMMENT '用户查询',
|
||||
answer LONGTEXT NULL COMMENT '回答消息内容',
|
||||
feedback VARCHAR(255) NULL COMMENT '点赞 like / 点踩 dislike',
|
||||
feedback_content TEXT NULL COMMENT '点踩内容',
|
||||
created_at BIGINT NULL COMMENT '创建时间(毫秒时间戳)',
|
||||
updated_at BIGINT NULL COMMENT '更新时间(毫秒时间戳)',
|
||||
`log` JSON NULL COMMENT '当前对话日志(JSON字符串)',
|
||||
UNIQUE KEY uk_random_code (random_code),
|
||||
UNIQUE KEY uk_message_id (message_id),
|
||||
INDEX idx_conversation_id (conversation_id),
|
||||
INDEX idx_created_at (created_at)
|
||||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='消息记录表';
|
||||
"""
|
||||
|
||||
conversation_sql = f"""
|
||||
CREATE TABLE IF NOT EXISTS {self.conversation_table} (
|
||||
id BIGINT PRIMARY KEY AUTO_INCREMENT,
|
||||
random_code VARCHAR(100) NULL COMMENT '业务主键',
|
||||
create_user VARCHAR(100) NULL COMMENT '创建人',
|
||||
create_date DATETIME NULL COMMENT '创建时间',
|
||||
update_user VARCHAR(100) NULL COMMENT '修改人',
|
||||
update_date DATETIME NULL COMMENT '修改时间',
|
||||
create_user_name VARCHAR(255) NULL COMMENT '创建人姓名',
|
||||
update_user_name VARCHAR(255) NULL COMMENT '修改人姓名',
|
||||
conversation_id VARCHAR(255) NULL COMMENT '会话 ID',
|
||||
user VARCHAR(255) NULL COMMENT '用户',
|
||||
name VARCHAR(512) NULL COMMENT '会话名称',
|
||||
status VARCHAR(255) NULL COMMENT '状态',
|
||||
introduction VARCHAR(255) NULL COMMENT '开场白',
|
||||
created_at BIGINT NULL COMMENT '创建时间(毫秒时间戳)',
|
||||
updated_at BIGINT NULL COMMENT '更新时间(毫秒时间戳)',
|
||||
UNIQUE KEY uk_conversation_random_code (random_code),
|
||||
UNIQUE KEY uk_conversation_id (conversation_id),
|
||||
INDEX idx_conversation_updated_at (updated_at)
|
||||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='会话记录表';
|
||||
"""
|
||||
|
||||
try:
|
||||
with self._get_conn() as conn:
|
||||
with conn.cursor() as cur:
|
||||
cur.execute(message_sql)
|
||||
cur.execute(conversation_sql)
|
||||
self._ensure_conversation_schema(conn)
|
||||
self._inited = True
|
||||
self._log_entity_stage(
|
||||
"storage",
|
||||
"ensure_table",
|
||||
"success",
|
||||
started_at,
|
||||
{"messages_table": self.table, "conversation_table": self.conversation_table},
|
||||
)
|
||||
except Exception as exc:
|
||||
self._log_entity_stage(
|
||||
"storage",
|
||||
"ensure_table",
|
||||
"failed",
|
||||
started_at,
|
||||
{"messages_table": self.table, "conversation_table": self.conversation_table, "error": str(exc)},
|
||||
)
|
||||
self._log_local("message_storage.ensure_table_failed", {
|
||||
"error": str(exc),
|
||||
"messages_table": self.table,
|
||||
"conversation_table": self.conversation_table,
|
||||
})
|
||||
# 开发阶段容错,避免初始化失败影响主流程
|
||||
self.enabled = False
|
||||
|
||||
def create_conversation(
|
||||
self,
|
||||
conversation_id: str,
|
||||
user: Optional[str],
|
||||
name: Optional[str],
|
||||
status: str,
|
||||
introduction: Optional[str],
|
||||
created_at: int,
|
||||
updated_at: int,
|
||||
) -> bool:
|
||||
started_at = self._now_ms()
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"create",
|
||||
"start",
|
||||
started_at,
|
||||
{
|
||||
"conversation_id": conversation_id,
|
||||
"user": user,
|
||||
"name_len": len(name or ""),
|
||||
"status": status,
|
||||
},
|
||||
)
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"create",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "reason": "storage_disabled"},
|
||||
)
|
||||
return False
|
||||
|
||||
self._ensure_table()
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"create",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "reason": "storage_disabled_after_init"},
|
||||
)
|
||||
return False
|
||||
|
||||
insert_sql = f"""
|
||||
INSERT INTO {self.conversation_table}(
|
||||
random_code, create_user, create_date, update_user, update_date,
|
||||
create_user_name, update_user_name,
|
||||
conversation_id, user, name, status, introduction, created_at, updated_at
|
||||
) VALUES(%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
||||
"""
|
||||
try:
|
||||
audit = self._build_audit_fields(
|
||||
random_code=conversation_id,
|
||||
user=user,
|
||||
created_value=created_at,
|
||||
updated_value=updated_at,
|
||||
)
|
||||
with self._get_conn() as conn:
|
||||
self._ensure_conversation_schema(conn)
|
||||
with conn.cursor() as cur:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"create",
|
||||
"sql_execute",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "table": self.conversation_table},
|
||||
)
|
||||
cur.execute(
|
||||
insert_sql,
|
||||
(
|
||||
audit["random_code"],
|
||||
audit["create_user"],
|
||||
audit["create_date"],
|
||||
audit["update_user"],
|
||||
audit["update_date"],
|
||||
audit["create_user_name"],
|
||||
audit["update_user_name"],
|
||||
conversation_id,
|
||||
user,
|
||||
name,
|
||||
status,
|
||||
introduction,
|
||||
audit["created_at"],
|
||||
audit["updated_at"],
|
||||
),
|
||||
)
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"create",
|
||||
"success",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "table": self.conversation_table},
|
||||
)
|
||||
return True
|
||||
except Exception as exc:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"create",
|
||||
"failed",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "table": self.conversation_table, "error": str(exc)},
|
||||
)
|
||||
self._log_local("message_storage.create_conversation_failed", {
|
||||
"error": str(exc),
|
||||
"conversation_table": self.conversation_table,
|
||||
"conversation_id": conversation_id,
|
||||
})
|
||||
return False
|
||||
|
||||
def get_conversation_by_id(self, conversation_id: str) -> Optional[Dict[str, Any]]:
|
||||
started_at = self._now_ms()
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"get_by_id",
|
||||
"start",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id},
|
||||
)
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"get_by_id",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "reason": "storage_disabled"},
|
||||
)
|
||||
return None
|
||||
|
||||
self._ensure_table()
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"get_by_id",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "reason": "storage_disabled_after_init"},
|
||||
)
|
||||
return None
|
||||
|
||||
select_sql = f"""
|
||||
SELECT conversation_id, user, name, status, introduction, created_at, updated_at
|
||||
FROM {self.conversation_table}
|
||||
WHERE conversation_id = %s
|
||||
LIMIT 1
|
||||
"""
|
||||
try:
|
||||
with self._get_conn() as conn:
|
||||
self._ensure_conversation_schema(conn)
|
||||
with conn.cursor(pymysql.cursors.DictCursor) as cur:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"get_by_id",
|
||||
"sql_execute",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "table": self.conversation_table},
|
||||
)
|
||||
cur.execute(select_sql, (conversation_id,))
|
||||
result = cur.fetchone()
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"get_by_id",
|
||||
"success",
|
||||
started_at,
|
||||
{
|
||||
"conversation_id": conversation_id,
|
||||
"table": self.conversation_table,
|
||||
"found": bool(result),
|
||||
},
|
||||
)
|
||||
return dict(result) if result else None
|
||||
except Exception as exc:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"get_by_id",
|
||||
"failed",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "table": self.conversation_table, "error": str(exc)},
|
||||
)
|
||||
self._log_local("message_storage.get_conversation_failed", {
|
||||
"error": str(exc),
|
||||
"conversation_table": self.conversation_table,
|
||||
"conversation_id": conversation_id,
|
||||
})
|
||||
return None
|
||||
|
||||
def update_conversation_updated_at(self, conversation_id: str, updated_at: int) -> bool:
|
||||
started_at = self._now_ms()
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"update_updated_at",
|
||||
"start",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "updated_at": updated_at},
|
||||
)
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"update_updated_at",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "reason": "storage_disabled"},
|
||||
)
|
||||
return False
|
||||
|
||||
self._ensure_table()
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"update_updated_at",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "reason": "storage_disabled_after_init"},
|
||||
)
|
||||
return False
|
||||
|
||||
update_sql = f"""
|
||||
UPDATE {self.conversation_table}
|
||||
SET updated_at = %s,
|
||||
update_date = %s
|
||||
WHERE conversation_id = %s
|
||||
"""
|
||||
try:
|
||||
update_ms = self._resolve_epoch_millis(updated_at)
|
||||
update_bundle = DateTimeGenerator.bundle(update_ms, default_to_now=True)
|
||||
with self._get_conn() as conn:
|
||||
self._ensure_conversation_schema(conn)
|
||||
with conn.cursor() as cur:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"update_updated_at",
|
||||
"sql_execute",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "table": self.conversation_table},
|
||||
)
|
||||
affected_rows = cur.execute(
|
||||
update_sql,
|
||||
(update_ms, update_bundle.db_datetime, conversation_id),
|
||||
)
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"update_updated_at",
|
||||
"success",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "affected_rows": int(affected_rows or 0)},
|
||||
)
|
||||
return bool(affected_rows)
|
||||
except Exception as exc:
|
||||
self._log_entity_stage(
|
||||
"conversations",
|
||||
"update_updated_at",
|
||||
"failed",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "table": self.conversation_table, "error": str(exc)},
|
||||
)
|
||||
self._log_local("message_storage.update_conversation_failed", {
|
||||
"error": str(exc),
|
||||
"conversation_table": self.conversation_table,
|
||||
"conversation_id": conversation_id,
|
||||
})
|
||||
return False
|
||||
|
||||
def save_message(
|
||||
self,
|
||||
conversation_id: str,
|
||||
message_id: str,
|
||||
query: str,
|
||||
answer: Optional[str] = None,
|
||||
workflow_type: Optional[str] = None,
|
||||
user: Optional[str] = None,
|
||||
sql_query: Optional[str] = None,
|
||||
execution_result: Optional[Dict[str, Any]] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
created_at: Optional[int] = None,
|
||||
updated_at: Optional[int] = None,
|
||||
logs: Optional[List[str]] = None,
|
||||
) -> bool:
|
||||
started_at = self._now_ms()
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"create",
|
||||
"start",
|
||||
started_at,
|
||||
{
|
||||
"conversation_id": conversation_id,
|
||||
"message_id": message_id,
|
||||
"query_len": len(query or ""),
|
||||
"answer_len": len(answer or ""),
|
||||
"workflow_type": workflow_type,
|
||||
},
|
||||
)
|
||||
"""
|
||||
保存消息记录
|
||||
|
||||
Args:
|
||||
conversation_id: 会话 ID
|
||||
message_id: 消息 ID
|
||||
query: 用户查询
|
||||
answer: AI 回复
|
||||
workflow_type: 工作流类型 (conversation/tool_using)
|
||||
user: 用户标识
|
||||
sql_query: 生成的 SQL
|
||||
execution_result: SQL 执行结果
|
||||
metadata: 其他元数据
|
||||
logs: 过程日志(兼容 Java saveMessageToDB)
|
||||
|
||||
Returns:
|
||||
bool: 是否保存成功
|
||||
"""
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"create",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "message_id": message_id, "reason": "storage_disabled"},
|
||||
)
|
||||
return False
|
||||
|
||||
self._ensure_table()
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"create",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "message_id": message_id, "reason": "storage_disabled_after_init"},
|
||||
)
|
||||
return False
|
||||
|
||||
audit = self._build_audit_fields(
|
||||
random_code=message_id,
|
||||
user=user,
|
||||
created_value=created_at,
|
||||
updated_value=updated_at,
|
||||
)
|
||||
log_payload = {
|
||||
"workflow_type": workflow_type,
|
||||
"sql_query": sql_query,
|
||||
"execution_result": execution_result or {},
|
||||
"metadata": metadata or {},
|
||||
}
|
||||
normalized_logs = [str(item) for item in (logs or []) if str(item).strip()]
|
||||
if normalized_logs:
|
||||
log_payload["data"] = "\n".join(normalized_logs)
|
||||
message_record = MessagesDTO(
|
||||
message_id=message_id,
|
||||
conversation_id=conversation_id,
|
||||
user=user,
|
||||
query=query,
|
||||
answer=answer,
|
||||
feedback=None,
|
||||
feedback_content=None,
|
||||
created_at=audit["created_at"],
|
||||
updated_at=audit["updated_at"],
|
||||
log=log_payload,
|
||||
)
|
||||
|
||||
insert_sql = f"""
|
||||
INSERT INTO {self.table}(
|
||||
random_code, create_user, create_date, update_user, update_date,
|
||||
create_user_name, update_user_name,
|
||||
message_id, conversation_id, user, query, answer,
|
||||
feedback, feedback_content, created_at, updated_at, `log`
|
||||
) VALUES(%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
|
||||
"""
|
||||
|
||||
try:
|
||||
with self._get_conn() as conn:
|
||||
with conn.cursor() as cur:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"create",
|
||||
"sql_execute",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "message_id": message_id, "table": self.table},
|
||||
)
|
||||
cur.execute(
|
||||
insert_sql,
|
||||
(
|
||||
audit["random_code"],
|
||||
audit["create_user"],
|
||||
audit["create_date"],
|
||||
audit["update_user"],
|
||||
audit["update_date"],
|
||||
audit["create_user_name"],
|
||||
audit["update_user_name"],
|
||||
message_record.message_id,
|
||||
message_record.conversation_id,
|
||||
message_record.user,
|
||||
message_record.query,
|
||||
message_record.answer,
|
||||
message_record.feedback,
|
||||
message_record.feedback_content,
|
||||
message_record.created_at,
|
||||
message_record.updated_at,
|
||||
json.dumps(message_record.log, ensure_ascii=False),
|
||||
),
|
||||
)
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"create",
|
||||
"success",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "message_id": message_id, "table": self.table},
|
||||
)
|
||||
return True
|
||||
except Exception as exc:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"create",
|
||||
"failed",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "message_id": message_id, "table": self.table, "error": str(exc)},
|
||||
)
|
||||
self._log_local("message_storage.save_failed", {
|
||||
"error": str(exc),
|
||||
"messages_table": self.table,
|
||||
"message_id": message_id,
|
||||
"conversation_id": conversation_id,
|
||||
})
|
||||
# 开发阶段容错,避免日志失败影响主流程
|
||||
return False
|
||||
|
||||
def get_conversation_history(
|
||||
self,
|
||||
conversation_id: str,
|
||||
limit: int = 20
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
获取会话历史消息
|
||||
|
||||
Args:
|
||||
conversation_id: 会话 ID
|
||||
limit: 返回消息数量
|
||||
|
||||
Returns:
|
||||
List[Dict[str, Any]]: 消息列表
|
||||
"""
|
||||
started_at = self._now_ms()
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"get_history",
|
||||
"start",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "limit": limit},
|
||||
)
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"get_history",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "reason": "storage_disabled"},
|
||||
)
|
||||
return []
|
||||
|
||||
self._ensure_table()
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"get_history",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "reason": "storage_disabled_after_init"},
|
||||
)
|
||||
return []
|
||||
|
||||
select_sql = f"""
|
||||
SELECT * FROM {self.table}
|
||||
WHERE conversation_id = %s
|
||||
ORDER BY created_at DESC
|
||||
LIMIT %s
|
||||
"""
|
||||
|
||||
try:
|
||||
with self._get_conn() as conn:
|
||||
with conn.cursor(pymysql.cursors.DictCursor) as cur:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"get_history",
|
||||
"sql_execute",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "limit": limit, "table": self.table},
|
||||
)
|
||||
cur.execute(select_sql, (conversation_id, limit))
|
||||
results = cur.fetchall()
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"get_history",
|
||||
"success",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "count": len(results or [])},
|
||||
)
|
||||
return list(results)
|
||||
except Exception as exc:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"get_history",
|
||||
"failed",
|
||||
started_at,
|
||||
{"conversation_id": conversation_id, "error": str(exc)},
|
||||
)
|
||||
return []
|
||||
|
||||
def update_feedback_by_message_id(
|
||||
self,
|
||||
message_id: str,
|
||||
feedback: str,
|
||||
feedback_content: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""按 message_id 回写点赞/点踩反馈。"""
|
||||
started_at = self._now_ms()
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"update_feedback",
|
||||
"start",
|
||||
started_at,
|
||||
{"message_id": message_id, "feedback": feedback},
|
||||
)
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"update_feedback",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"message_id": message_id, "reason": "storage_disabled"},
|
||||
)
|
||||
return False
|
||||
|
||||
self._ensure_table()
|
||||
if not self.enabled:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"update_feedback",
|
||||
"skipped",
|
||||
started_at,
|
||||
{"message_id": message_id, "reason": "storage_disabled_after_init"},
|
||||
)
|
||||
return False
|
||||
|
||||
update_sql = f"""
|
||||
UPDATE {self.table}
|
||||
SET feedback = %s,
|
||||
feedback_content = %s,
|
||||
updated_at = %s,
|
||||
update_date = %s
|
||||
WHERE message_id = %s
|
||||
"""
|
||||
|
||||
normalized_feedback_content = (feedback_content or "").strip() or None
|
||||
now_bundle = DateTimeGenerator.now()
|
||||
try:
|
||||
with self._get_conn() as conn:
|
||||
with conn.cursor() as cur:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"update_feedback",
|
||||
"sql_execute",
|
||||
started_at,
|
||||
{"message_id": message_id, "table": self.table},
|
||||
)
|
||||
affected_rows = cur.execute(
|
||||
update_sql,
|
||||
(
|
||||
feedback,
|
||||
normalized_feedback_content,
|
||||
now_bundle.epoch_millis,
|
||||
now_bundle.db_datetime,
|
||||
message_id,
|
||||
),
|
||||
)
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"update_feedback",
|
||||
"success",
|
||||
started_at,
|
||||
{"message_id": message_id, "affected_rows": int(affected_rows or 0)},
|
||||
)
|
||||
return bool(affected_rows)
|
||||
except Exception as exc:
|
||||
self._log_entity_stage(
|
||||
"messages",
|
||||
"update_feedback",
|
||||
"failed",
|
||||
started_at,
|
||||
{"message_id": message_id, "error": str(exc)},
|
||||
)
|
||||
# 开发阶段容错,避免日志失败影响主流程
|
||||
return False
|
||||
|
||||
|
||||
# 全局单例
|
||||
_GLOBAL_MESSAGE_STORAGE: Optional[MessageStorage] = None
|
||||
|
||||
|
||||
def get_message_storage() -> MessageStorage:
|
||||
"""获取消息存储服务实例"""
|
||||
global _GLOBAL_MESSAGE_STORAGE
|
||||
if _GLOBAL_MESSAGE_STORAGE is None:
|
||||
_GLOBAL_MESSAGE_STORAGE = MessageStorage()
|
||||
return _GLOBAL_MESSAGE_STORAGE
|
||||
@@ -0,0 +1,104 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import pymysql
|
||||
|
||||
from config import Config
|
||||
from services.common.datetime_utils import DateTimeGenerator
|
||||
|
||||
|
||||
class StructuredLogger:
|
||||
def __init__(self):
|
||||
cfg = Config.get_section("logging_mysql")
|
||||
self.enabled = str(cfg.get("enabled", "false")).lower() in ("1", "true", "yes")
|
||||
self.host = cfg.get("host", "127.0.0.1")
|
||||
self.port = int(cfg.get("port", 3306))
|
||||
self.user = cfg.get("user", "root")
|
||||
self.password = cfg.get("password", "")
|
||||
self.database = cfg.get("database", "more_dots")
|
||||
self.table = cfg.get("table", "structured_logs")
|
||||
self.connect_timeout = int(cfg.get("connect_timeout", 5))
|
||||
self._inited = False
|
||||
|
||||
def _get_conn(self):
|
||||
return pymysql.connect(
|
||||
host=self.host,
|
||||
port=self.port,
|
||||
user=self.user,
|
||||
password=self.password,
|
||||
database=self.database,
|
||||
charset="utf8mb4",
|
||||
autocommit=True,
|
||||
connect_timeout=self.connect_timeout,
|
||||
)
|
||||
|
||||
def _ensure_table(self) -> None:
|
||||
if self._inited or not self.enabled:
|
||||
return
|
||||
sql = f"""
|
||||
CREATE TABLE IF NOT EXISTS {self.table} (
|
||||
id BIGINT PRIMARY KEY AUTO_INCREMENT,
|
||||
trace_id VARCHAR(64) NOT NULL,
|
||||
level VARCHAR(16) NOT NULL,
|
||||
event VARCHAR(128) NOT NULL,
|
||||
error_code VARCHAR(64) NULL,
|
||||
payload JSON NULL,
|
||||
created_at DATETIME NOT NULL
|
||||
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4;
|
||||
"""
|
||||
try:
|
||||
with self._get_conn() as conn:
|
||||
with conn.cursor() as cur:
|
||||
cur.execute(sql)
|
||||
self._inited = True
|
||||
except Exception:
|
||||
# 开发阶段容错,避免日志失败影响主流程
|
||||
self.enabled = False
|
||||
|
||||
def log(self, level: str, event: str, trace_id: str, payload: Optional[Dict[str, Any]] = None, error_code: Optional[str] = None) -> None:
|
||||
print(json.dumps({
|
||||
"trace_id": trace_id,
|
||||
"level": level,
|
||||
"event": event,
|
||||
"error_code": error_code,
|
||||
"payload": payload or {},
|
||||
"created_at": DateTimeGenerator.now().iso_str,
|
||||
}, ensure_ascii=False))
|
||||
|
||||
if not self.enabled:
|
||||
return
|
||||
|
||||
self._ensure_table()
|
||||
if not self.enabled:
|
||||
return
|
||||
|
||||
insert_sql = f"INSERT INTO {self.table}(trace_id, level, event, error_code, payload, created_at) VALUES(%s,%s,%s,%s,%s,%s)"
|
||||
try:
|
||||
with self._get_conn() as conn:
|
||||
with conn.cursor() as cur:
|
||||
cur.execute(
|
||||
insert_sql,
|
||||
(
|
||||
trace_id,
|
||||
level,
|
||||
event,
|
||||
error_code,
|
||||
json.dumps(payload or {}, ensure_ascii=False),
|
||||
DateTimeGenerator.now().db_datetime,
|
||||
),
|
||||
)
|
||||
except Exception:
|
||||
# 开发阶段容错,避免日志失败影响主流程
|
||||
return
|
||||
|
||||
|
||||
_GLOBAL_STRUCTURED_LOGGER: Optional[StructuredLogger] = None
|
||||
|
||||
|
||||
def get_structured_logger() -> StructuredLogger:
|
||||
global _GLOBAL_STRUCTURED_LOGGER
|
||||
if _GLOBAL_STRUCTURED_LOGGER is None:
|
||||
_GLOBAL_STRUCTURED_LOGGER = StructuredLogger()
|
||||
return _GLOBAL_STRUCTURED_LOGGER
|
||||
@@ -1,30 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
|
||||
class StructuredLogger:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def log(self, level: str, event: str, trace_id: str, payload: Optional[Dict[str, Any]] = None, error_code: Optional[str] = None) -> None:
|
||||
print(json.dumps({
|
||||
"trace_id": trace_id,
|
||||
"level": level,
|
||||
"event": event,
|
||||
"error_code": error_code,
|
||||
"payload": payload or {},
|
||||
"created_at": datetime.now().isoformat(),
|
||||
}, ensure_ascii=False))
|
||||
|
||||
|
||||
_GLOBAL_STRUCTURED_LOGGER: Optional[StructuredLogger] = None
|
||||
|
||||
|
||||
def get_structured_logger() -> StructuredLogger:
|
||||
global _GLOBAL_STRUCTURED_LOGGER
|
||||
if _GLOBAL_STRUCTURED_LOGGER is None:
|
||||
_GLOBAL_STRUCTURED_LOGGER = StructuredLogger()
|
||||
return _GLOBAL_STRUCTURED_LOGGER
|
||||
@@ -1,55 +0,0 @@
|
||||
from typing import Any, Dict
|
||||
|
||||
from config import Config
|
||||
from services.ragflow_client import RagflowClient, extract_table_name
|
||||
|
||||
|
||||
class TemplateMatcher:
|
||||
"""模板匹配器:RAGFlow 检索"""
|
||||
|
||||
def __init__(self):
|
||||
self._ragflow = RagflowClient()
|
||||
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))
|
||||
|
||||
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]:
|
||||
"""返回匹配的表名与原始响应"""
|
||||
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, "raw": {"error": str(e)}}
|
||||
|
||||
candidates = []
|
||||
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 table_name:
|
||||
candidates.append(table_name)
|
||||
|
||||
matched = candidates[0] if candidates else None
|
||||
return {"table_name": matched, "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
|
||||
@@ -1,272 +0,0 @@
|
||||
"""
|
||||
工具路由器模块
|
||||
|
||||
支持动态注册和管理工具
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Any, Callable, Dict, List, Optional, Type
|
||||
|
||||
from langchain_core.tools import BaseTool
|
||||
|
||||
from tools.calculator import CalculatorTool
|
||||
from tools.web_search import WebSearchTool
|
||||
from tools.rest_api_tool import RestApiTool
|
||||
from tools.sr_api_tool import SrApiQueryTool
|
||||
from core.registry import ToolRegistry, ToolMetadata
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ToolRouter:
|
||||
"""
|
||||
工具路由器:统一调用入口
|
||||
|
||||
支持特性:
|
||||
- 动态注册工具
|
||||
- 工具元数据管理
|
||||
- 执行监控
|
||||
"""
|
||||
|
||||
def __init__(self, tools: Optional[List[BaseTool]] = None):
|
||||
self._tools: Dict[str, BaseTool] = {}
|
||||
self._tool_metadata: Dict[str, ToolMetadata] = {}
|
||||
self._execution_stats: Dict[str, Dict[str, Any]] = {}
|
||||
|
||||
if tools is not None:
|
||||
for tool in tools:
|
||||
self.register_tool(tool)
|
||||
else:
|
||||
self._register_default_tools()
|
||||
|
||||
def _register_default_tools(self) -> None:
|
||||
"""注册默认工具"""
|
||||
default_tools = [
|
||||
CalculatorTool(),
|
||||
WebSearchTool(),
|
||||
RestApiTool(),
|
||||
SrApiQueryTool(),
|
||||
]
|
||||
for tool in default_tools:
|
||||
self.register_tool(tool)
|
||||
|
||||
def register_tool(
|
||||
self,
|
||||
tool: BaseTool,
|
||||
description: str = "",
|
||||
version: str = "1.0.0",
|
||||
timeout: int = 30,
|
||||
retry: int = 0,
|
||||
tags: Optional[List[str]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
注册工具
|
||||
|
||||
Args:
|
||||
tool: 工具实例
|
||||
description: 描述(默认使用 tool.description)
|
||||
version: 版本
|
||||
timeout: 超时时间
|
||||
retry: 重试次数
|
||||
tags: 标签
|
||||
"""
|
||||
name = tool.name
|
||||
metadata = ToolMetadata(
|
||||
name=name,
|
||||
description=description or tool.description,
|
||||
version=version,
|
||||
timeout=timeout,
|
||||
retry=retry,
|
||||
tags=tags or [],
|
||||
)
|
||||
|
||||
self._tools[name] = tool
|
||||
self._tool_metadata[name] = metadata
|
||||
self._execution_stats[name] = {
|
||||
"total_calls": 0,
|
||||
"success_calls": 0,
|
||||
"failed_calls": 0,
|
||||
"total_time_ms": 0,
|
||||
}
|
||||
|
||||
ToolRegistry._entries[name] = type(
|
||||
"RegistryEntry",
|
||||
(),
|
||||
{"instance": tool, "metadata": {"tool_metadata": metadata}}
|
||||
)()
|
||||
|
||||
logger.info(f"Registered tool: {name} (v{version})")
|
||||
|
||||
def unregister_tool(self, name: str) -> bool:
|
||||
"""
|
||||
注销工具
|
||||
|
||||
Args:
|
||||
name: 工具名称
|
||||
|
||||
Returns:
|
||||
是否成功注销
|
||||
"""
|
||||
if name in self._tools:
|
||||
del self._tools[name]
|
||||
del self._tool_metadata[name]
|
||||
del self._execution_stats[name]
|
||||
ToolRegistry.unregister(name)
|
||||
logger.info(f"Unregistered tool: {name}")
|
||||
return True
|
||||
return False
|
||||
|
||||
def get_tool(self, name: str) -> Optional[BaseTool]:
|
||||
"""获取工具实例"""
|
||||
return self._tools.get(name)
|
||||
|
||||
def get_tool_metadata(self, name: str) -> Optional[ToolMetadata]:
|
||||
"""获取工具元数据"""
|
||||
return self._tool_metadata.get(name)
|
||||
|
||||
def list_tools(self) -> List[str]:
|
||||
"""列出可用工具名称"""
|
||||
return list(self._tools.keys())
|
||||
|
||||
def get_tool_info(self, name: str) -> Optional[Dict[str, Any]]:
|
||||
"""获取工具详细信息"""
|
||||
if name not in self._tools:
|
||||
return None
|
||||
|
||||
tool = self._tools[name]
|
||||
metadata = self._tool_metadata.get(name)
|
||||
stats = self._execution_stats.get(name, {})
|
||||
|
||||
return {
|
||||
"name": name,
|
||||
"description": metadata.description if metadata else tool.description,
|
||||
"version": metadata.version if metadata else "unknown",
|
||||
"timeout": metadata.timeout if metadata else 30,
|
||||
"tags": metadata.tags if metadata else [],
|
||||
"stats": {
|
||||
"total_calls": stats.get("total_calls", 0),
|
||||
"success_rate": self._calculate_success_rate(name),
|
||||
},
|
||||
}
|
||||
|
||||
def call(self, tool_name: str, payload: Any) -> Dict[str, Any]:
|
||||
"""
|
||||
调用工具并返回标准化结果
|
||||
|
||||
Args:
|
||||
tool_name: 工具名称
|
||||
payload: 输入参数
|
||||
|
||||
Returns:
|
||||
标准化结果 {ok, data, error}
|
||||
"""
|
||||
tool = self._tools.get(tool_name)
|
||||
if not tool:
|
||||
return {"ok": False, "data": None, "error": f"工具不存在: {tool_name}"}
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
if isinstance(payload, (dict, list)):
|
||||
input_value = json.dumps(payload, ensure_ascii=False)
|
||||
elif payload is None:
|
||||
input_value = ""
|
||||
else:
|
||||
input_value = str(payload)
|
||||
|
||||
result = tool.run(input_value)
|
||||
|
||||
self._record_success(tool_name, time.time() - start_time)
|
||||
|
||||
return {"ok": True, "data": result, "error": None}
|
||||
|
||||
except Exception as e:
|
||||
self._record_failure(tool_name, time.time() - start_time)
|
||||
return {"ok": False, "data": None, "error": str(e)}
|
||||
|
||||
def call_with_metadata(
|
||||
self,
|
||||
tool_name: str,
|
||||
payload: Any,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
调用工具并返回包含元数据的结果
|
||||
|
||||
Args:
|
||||
tool_name: 工具名称
|
||||
payload: 输入参数
|
||||
|
||||
Returns:
|
||||
包含元数据的结果
|
||||
"""
|
||||
result = self.call(tool_name, payload)
|
||||
metadata = self.get_tool_metadata(tool_name)
|
||||
|
||||
return {
|
||||
**result,
|
||||
"tool_name": tool_name,
|
||||
"tool_version": metadata.version if metadata else "unknown",
|
||||
"execution_time_ms": self._execution_stats.get(tool_name, {}).get("last_time_ms", 0),
|
||||
}
|
||||
|
||||
def _record_success(self, tool_name: str, elapsed: float) -> None:
|
||||
"""记录成功执行"""
|
||||
if tool_name in self._execution_stats:
|
||||
stats = self._execution_stats[tool_name]
|
||||
stats["total_calls"] += 1
|
||||
stats["success_calls"] += 1
|
||||
stats["total_time_ms"] += elapsed * 1000
|
||||
stats["last_time_ms"] = elapsed * 1000
|
||||
|
||||
def _record_failure(self, tool_name: str, elapsed: float) -> None:
|
||||
"""记录失败执行"""
|
||||
if tool_name in self._execution_stats:
|
||||
stats = self._execution_stats[tool_name]
|
||||
stats["total_calls"] += 1
|
||||
stats["failed_calls"] += 1
|
||||
stats["total_time_ms"] += elapsed * 1000
|
||||
stats["last_time_ms"] = elapsed * 1000
|
||||
|
||||
def _calculate_success_rate(self, tool_name: str) -> float:
|
||||
"""计算成功率"""
|
||||
stats = self._execution_stats.get(tool_name)
|
||||
if not stats or stats["total_calls"] == 0:
|
||||
return 0.0
|
||||
return stats["success_calls"] / stats["total_calls"]
|
||||
|
||||
def get_all_stats(self) -> Dict[str, Dict[str, Any]]:
|
||||
"""获取所有工具的执行统计"""
|
||||
result = {}
|
||||
for name in self._tools:
|
||||
result[name] = {
|
||||
**self._execution_stats.get(name, {}),
|
||||
"success_rate": self._calculate_success_rate(name),
|
||||
}
|
||||
return result
|
||||
|
||||
def register_function(
|
||||
self,
|
||||
name: str,
|
||||
func: Callable,
|
||||
description: str = "",
|
||||
timeout: int = 30,
|
||||
) -> None:
|
||||
"""
|
||||
将普通函数注册为工具
|
||||
|
||||
Args:
|
||||
name: 工具名称
|
||||
func: 函数
|
||||
description: 描述
|
||||
timeout: 超时时间
|
||||
"""
|
||||
from langchain_core.tools import Tool
|
||||
|
||||
tool = Tool(
|
||||
name=name,
|
||||
description=description,
|
||||
func=func,
|
||||
)
|
||||
self.register_tool(tool, description=description, timeout=timeout)
|
||||
@@ -0,0 +1,5 @@
|
||||
"""工具服务模块"""
|
||||
|
||||
from .tool_router import ToolRouter
|
||||
|
||||
__all__ = ["ToolRouter"]
|
||||
@@ -0,0 +1,41 @@
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from langchain_core.tools import BaseTool
|
||||
|
||||
from tools.calculator import CalculatorTool
|
||||
from tools.web_search import WebSearchTool
|
||||
from tools.rest_api_tool import RestApiTool
|
||||
from tools.sr_api_tool import SrApiQueryTool
|
||||
|
||||
|
||||
class ToolRouter:
|
||||
"""工具路由器:统一调用入口"""
|
||||
|
||||
def __init__(self, tools: Optional[list[BaseTool]] = None):
|
||||
if tools is None:
|
||||
tools = [CalculatorTool(), WebSearchTool(), RestApiTool(), SrApiQueryTool()]
|
||||
self._tools: Dict[str, BaseTool] = {tool.name: tool for tool in tools}
|
||||
|
||||
def list_tools(self) -> list[str]:
|
||||
"""列出可用工具名称"""
|
||||
return list(self._tools.keys())
|
||||
|
||||
def call(self, tool_name: str, payload: Any) -> Dict[str, Any]:
|
||||
"""调用工具并返回标准化结果"""
|
||||
tool = self._tools.get(tool_name)
|
||||
if not tool:
|
||||
return {"ok": False, "data": None, "error": f"工具不存在: {tool_name}"}
|
||||
|
||||
try:
|
||||
if isinstance(payload, (dict, list)):
|
||||
input_value = json.dumps(payload, ensure_ascii=False)
|
||||
elif payload is None:
|
||||
input_value = ""
|
||||
else:
|
||||
input_value = str(payload)
|
||||
|
||||
result = tool.run(input_value)
|
||||
return {"ok": True, "data": result, "error": None}
|
||||
except Exception as e:
|
||||
return {"ok": False, "data": None, "error": str(e)}
|
||||
Reference in New Issue
Block a user