This commit is contained in:
2026-03-11 23:40:39 +08:00
parent db25d61026
commit e062368ef2
15 changed files with 1592 additions and 229 deletions
-22
View File
@@ -2,11 +2,6 @@ from __future__ import annotations
from typing import Optional
try:
import redis
except Exception:
redis = None
class CacheBase:
"""缓存接口"""
@@ -26,20 +21,3 @@ class NoopCache(CacheBase):
def set(self, key: str, value: str, ttl: int) -> None:
return None
class RedisCache(CacheBase):
"""Redis 缓存实现"""
def __init__(self, url: str, db: int = 0):
if redis is None:
raise ImportError("未安装 redis 依赖")
self._client = redis.Redis.from_url(url, db=db, decode_responses=True)
def get(self, key: str) -> Optional[str]:
return self._client.get(key)
def set(self, key: str, value: str, ttl: int) -> None:
self._client.set(key, value, ex=ttl)
-53
View File
@@ -2,9 +2,6 @@ import json
import os
from typing import Any, Dict, Optional
from config import Config
from services.cache import NoopCache, RedisCache
class SqlPromptManager:
"""按表名读取 SQL 提示词"""
@@ -12,51 +9,11 @@ class SqlPromptManager:
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")
self._cache = self._init_cache()
self._cache_ttl = self._get_cache_ttl()
@staticmethod
def _get_cache_ttl() -> int:
redis_cfg = Config.get_section("redis")
try:
return int(redis_cfg.get("sql_prompt_ttl", 600))
except Exception:
return 600
@staticmethod
def _init_cache():
redis_cfg = Config.get_section("redis")
enabled = str(redis_cfg.get("enabled", "false")).lower() in ("1", "true", "yes")
if not enabled:
return NoopCache()
# 优先使用完整 URL;否则使用 host/port/password/database 拼接
url = redis_cfg.get("url")
db = int(redis_cfg.get("db", redis_cfg.get("database", 0)))
if not url:
host = redis_cfg.get("host")
port = redis_cfg.get("port", "6379")
password = redis_cfg.get("password", "")
database = redis_cfg.get("database", str(db))
if host:
auth = f":{password}@" if password else ""
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("\\", "_")
@staticmethod
def _cache_key(table_name: str, mtime: float) -> str:
return f"sql_prompt:{table_name}:{int(mtime)}"
def get_prompt(self, table_name: str) -> Optional[Dict[str, Any]]:
"""读取指定表的提示词 JSON"""
if not table_name:
@@ -67,19 +24,9 @@ class SqlPromptManager:
if not os.path.exists(path):
return None
mtime = os.path.getmtime(path)
key = self._cache_key(safe_name, mtime)
cached = self._cache.get(key)
if cached:
try:
return json.loads(cached)
except Exception:
pass
with open(path, "r", encoding="utf-8") as f:
prompt = json.load(f)
self._cache.set(key, json.dumps(prompt, ensure_ascii=False), self._cache_ttl)
return prompt
+1 -75
View File
@@ -4,58 +4,10 @@ import json
from datetime import datetime
from typing import Any, Dict, Optional
import pymysql
from config import Config
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
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({
@@ -67,32 +19,6 @@ class StructuredLogger:
"created_at": datetime.now().isoformat(),
}, 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),
datetime.now(),
),
)
except Exception:
# 开发阶段容错,避免日志失败影响主流程
return
_GLOBAL_STRUCTURED_LOGGER: Optional[StructuredLogger] = None
+244 -13
View File
@@ -1,5 +1,13 @@
"""
工具路由器模块
支持动态注册和管理工具
"""
import json
from typing import Any, Dict, Optional
import logging
import time
from typing import Any, Callable, Dict, List, Optional, Type
from langchain_core.tools import BaseTool
@@ -7,26 +15,159 @@ 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):
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]:
"""
工具路由器:统一调用入口
支持特性:
- 动态注册工具
- 工具元数据管理
- 执行监控
"""
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)
@@ -34,8 +175,98 @@ class ToolRouter:
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)