forked from xiaohei/taiji-AI-PAD
✅ 主要功能: - 实现函数注册表 (function_registry.py) - 16个内置安全函数(数学、字符串、日期、JSON、哈希等) - 函数元数据管理 - 函数查询和列表 - 实现沙箱执行器 (sandbox_executor.py) - 超时控制 (默认5秒) - 参数验证和资源限制 - 安全的执行环境 - 实现 _execute_function_tool 方法 - 函数白名单验证 - 参数类型验证 - 沙箱执行 - 完善的错误处理 🔒 安全特性: - 只允许执行注册表中的函数 - 禁止网络和文件系统访问 - 超时和资源限制 - 参数验证和类型检查 📚 文档: - 添加实现方案文档 - 添加实现完成文档 完成度: 95% (核心功能完成,单元测试待完成) 版本: v1.2.0
259 lines
8.8 KiB
Python
259 lines
8.8 KiB
Python
"""
|
|
函数注册表
|
|
管理安全的 Python 函数,用于 MCP 工具调用
|
|
"""
|
|
|
|
import logging
|
|
from typing import Any, Callable, Dict, Optional, List
|
|
from datetime import datetime
|
|
import math
|
|
import json
|
|
import re
|
|
import hashlib
|
|
import base64
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class FunctionRegistry:
|
|
"""函数注册表,管理安全的 Python 函数"""
|
|
|
|
def __init__(self):
|
|
self._functions: Dict[str, Dict[str, Any]] = {}
|
|
self._register_builtin_functions()
|
|
|
|
def _register_builtin_functions(self):
|
|
"""注册内置安全函数"""
|
|
|
|
# 数学函数
|
|
self.register(
|
|
name="math_add",
|
|
func=lambda a, b: float(a) + float(b),
|
|
description="两个数字相加",
|
|
parameters=[
|
|
{"name": "a", "type": "number", "description": "第一个数字"},
|
|
{"name": "b", "type": "number", "description": "第二个数字"}
|
|
],
|
|
returns={"type": "number", "description": "两数之和"}
|
|
)
|
|
|
|
self.register(
|
|
name="math_subtract",
|
|
func=lambda a, b: float(a) - float(b),
|
|
description="两个数字相减",
|
|
parameters=[
|
|
{"name": "a", "type": "number", "description": "被减数"},
|
|
{"name": "b", "type": "number", "description": "减数"}
|
|
],
|
|
returns={"type": "number", "description": "两数之差"}
|
|
)
|
|
|
|
self.register(
|
|
name="math_multiply",
|
|
func=lambda a, b: float(a) * float(b),
|
|
description="两个数字相乘",
|
|
parameters=[
|
|
{"name": "a", "type": "number", "description": "第一个数字"},
|
|
{"name": "b", "type": "number", "description": "第二个数字"}
|
|
],
|
|
returns={"type": "number", "description": "两数之积"}
|
|
)
|
|
|
|
self.register(
|
|
name="math_divide",
|
|
func=lambda a, b: float(a) / float(b) if float(b) != 0 else None,
|
|
description="两个数字相除",
|
|
parameters=[
|
|
{"name": "a", "type": "number", "description": "被除数"},
|
|
{"name": "b", "type": "number", "description": "除数"}
|
|
],
|
|
returns={"type": "number", "description": "两数之商"}
|
|
)
|
|
|
|
self.register(
|
|
name="math_power",
|
|
func=lambda a, b: math.pow(float(a), float(b)),
|
|
description="计算 a 的 b 次方",
|
|
parameters=[
|
|
{"name": "a", "type": "number", "description": "底数"},
|
|
{"name": "b", "type": "number", "description": "指数"}
|
|
],
|
|
returns={"type": "number", "description": "a 的 b 次方"}
|
|
)
|
|
|
|
# 字符串函数
|
|
self.register(
|
|
name="string_upper",
|
|
func=lambda s: str(s).upper(),
|
|
description="将字符串转换为大写",
|
|
parameters=[
|
|
{"name": "s", "type": "string", "description": "输入字符串"}
|
|
],
|
|
returns={"type": "string", "description": "大写字符串"}
|
|
)
|
|
|
|
self.register(
|
|
name="string_lower",
|
|
func=lambda s: str(s).lower(),
|
|
description="将字符串转换为小写",
|
|
parameters=[
|
|
{"name": "s", "type": "string", "description": "输入字符串"}
|
|
],
|
|
returns={"type": "string", "description": "小写字符串"}
|
|
)
|
|
|
|
self.register(
|
|
name="string_length",
|
|
func=lambda s: len(str(s)),
|
|
description="获取字符串长度",
|
|
parameters=[
|
|
{"name": "s", "type": "string", "description": "输入字符串"}
|
|
],
|
|
returns={"type": "number", "description": "字符串长度"}
|
|
)
|
|
|
|
self.register(
|
|
name="string_replace",
|
|
func=lambda s, old, new: str(s).replace(str(old), str(new)),
|
|
description="替换字符串中的子串",
|
|
parameters=[
|
|
{"name": "s", "type": "string", "description": "原字符串"},
|
|
{"name": "old", "type": "string", "description": "要替换的子串"},
|
|
{"name": "new", "type": "string", "description": "新子串"}
|
|
],
|
|
returns={"type": "string", "description": "替换后的字符串"}
|
|
)
|
|
|
|
# 日期时间函数
|
|
self.register(
|
|
name="datetime_now",
|
|
func=lambda: datetime.utcnow().isoformat(),
|
|
description="获取当前 UTC 时间",
|
|
parameters=[],
|
|
returns={"type": "string", "description": "ISO 格式的时间字符串"}
|
|
)
|
|
|
|
# JSON 函数
|
|
self.register(
|
|
name="json_parse",
|
|
func=lambda s: json.loads(str(s)),
|
|
description="解析 JSON 字符串",
|
|
parameters=[
|
|
{"name": "s", "type": "string", "description": "JSON 字符串"}
|
|
],
|
|
returns={"type": "object", "description": "解析后的对象"}
|
|
)
|
|
|
|
self.register(
|
|
name="json_stringify",
|
|
func=lambda obj: json.dumps(obj),
|
|
description="将对象转换为 JSON 字符串",
|
|
parameters=[
|
|
{"name": "obj", "type": "object", "description": "要转换的对象"}
|
|
],
|
|
returns={"type": "string", "description": "JSON 字符串"}
|
|
)
|
|
|
|
# 哈希函数
|
|
self.register(
|
|
name="hash_md5",
|
|
func=lambda s: hashlib.md5(str(s).encode()).hexdigest(),
|
|
description="计算字符串的 MD5 哈希值",
|
|
parameters=[
|
|
{"name": "s", "type": "string", "description": "输入字符串"}
|
|
],
|
|
returns={"type": "string", "description": "MD5 哈希值"}
|
|
)
|
|
|
|
self.register(
|
|
name="hash_sha256",
|
|
func=lambda s: hashlib.sha256(str(s).encode()).hexdigest(),
|
|
description="计算字符串的 SHA256 哈希值",
|
|
parameters=[
|
|
{"name": "s", "type": "string", "description": "输入字符串"}
|
|
],
|
|
returns={"type": "string", "description": "SHA256 哈希值"}
|
|
)
|
|
|
|
# Base64 编码/解码
|
|
self.register(
|
|
name="base64_encode",
|
|
func=lambda s: base64.b64encode(str(s).encode()).decode(),
|
|
description="Base64 编码",
|
|
parameters=[
|
|
{"name": "s", "type": "string", "description": "输入字符串"}
|
|
],
|
|
returns={"type": "string", "description": "Base64 编码后的字符串"}
|
|
)
|
|
|
|
self.register(
|
|
name="base64_decode",
|
|
func=lambda s: base64.b64decode(str(s)).decode(),
|
|
description="Base64 解码",
|
|
parameters=[
|
|
{"name": "s", "type": "string", "description": "输入字符串"}
|
|
],
|
|
returns={"type": "string", "description": "Base64 解码后的字符串"}
|
|
)
|
|
|
|
logger.info(f"已注册 {len(self._functions)} 个内置函数")
|
|
|
|
def register(
|
|
self,
|
|
name: str,
|
|
func: Callable,
|
|
description: str,
|
|
parameters: List[Dict[str, Any]],
|
|
returns: Dict[str, Any],
|
|
category: str = "builtin"
|
|
):
|
|
"""注册一个函数"""
|
|
if name in self._functions:
|
|
logger.warning(f"函数 {name} 已存在,将被覆盖")
|
|
|
|
self._functions[name] = {
|
|
"name": name,
|
|
"func": func,
|
|
"description": description,
|
|
"parameters": parameters,
|
|
"returns": returns,
|
|
"category": category,
|
|
"registered_at": datetime.utcnow().isoformat()
|
|
}
|
|
|
|
logger.debug(f"注册函数: {name}")
|
|
|
|
def get(self, name: str) -> Optional[Dict[str, Any]]:
|
|
"""获取函数信息"""
|
|
return self._functions.get(name)
|
|
|
|
def list_all(self) -> List[Dict[str, Any]]:
|
|
"""列出所有注册的函数"""
|
|
return [
|
|
{
|
|
"name": info["name"],
|
|
"description": info["description"],
|
|
"parameters": info["parameters"],
|
|
"returns": info["returns"],
|
|
"category": info["category"]
|
|
}
|
|
for info in self._functions.values()
|
|
]
|
|
|
|
def exists(self, name: str) -> bool:
|
|
"""检查函数是否存在"""
|
|
return name in self._functions
|
|
|
|
|
|
# 全局函数注册表实例
|
|
_function_registry = None
|
|
|
|
|
|
def get_function_registry() -> FunctionRegistry:
|
|
"""获取全局函数注册表实例"""
|
|
global _function_registry
|
|
if _function_registry is None:
|
|
_function_registry = FunctionRegistry()
|
|
return _function_registry
|
|
|