This commit is contained in:
2026-03-24 18:07:22 +08:00
parent e062368ef2
commit 9a16f738d8
121 changed files with 8904 additions and 3940 deletions
+29
View File
@@ -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/`:错误码与应用异常
+74
View File
@@ -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",
]
-23
View File
@@ -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
+6
View File
@@ -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"
+152
View File
@@ -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)
+13
View File
@@ -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()
+238
View File
@@ -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
+201
View File
@@ -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
+24
View File
@@ -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
-41
View File
@@ -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
+14
View File
@@ -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",
]
+64
View File
@@ -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)
+885
View File
@@ -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
+104
View File
@@ -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
-30
View File
@@ -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
-55
View File
@@ -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
-272
View File
@@ -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)
+5
View File
@@ -0,0 +1,5 @@
"""工具服务模块"""
from .tool_router import ToolRouter
__all__ = ["ToolRouter"]
+41
View File
@@ -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)}