x
This commit is contained in:
@@ -2,11 +2,6 @@ from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
try:
|
||||
import redis
|
||||
except Exception:
|
||||
redis = None
|
||||
|
||||
|
||||
class CacheBase:
|
||||
"""缓存接口"""
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user