mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-25 00:23:53 +08:00
Merge branch 'mangger-api-voice-print' into py_test
This commit is contained in:
+58
@@ -0,0 +1,58 @@
|
||||
"""
|
||||
缓存配置管理
|
||||
"""
|
||||
|
||||
from enum import Enum
|
||||
from typing import Dict, Any, Optional
|
||||
from dataclasses import dataclass
|
||||
from .strategies import CacheStrategy
|
||||
|
||||
|
||||
class CacheType(Enum):
|
||||
"""缓存类型枚举"""
|
||||
|
||||
LOCATION = "location"
|
||||
WEATHER = "weather"
|
||||
LUNAR = "lunar"
|
||||
INTENT = "intent"
|
||||
IP_INFO = "ip_info"
|
||||
CONFIG = "config"
|
||||
DEVICE_PROMPT = "device_prompt"
|
||||
|
||||
|
||||
@dataclass
|
||||
class CacheConfig:
|
||||
"""缓存配置类"""
|
||||
|
||||
strategy: CacheStrategy = CacheStrategy.TTL
|
||||
ttl: Optional[float] = 300 # 默认5分钟
|
||||
max_size: Optional[int] = 1000 # 默认最大1000条
|
||||
cleanup_interval: float = 60 # 清理间隔(秒)
|
||||
|
||||
@classmethod
|
||||
def for_type(cls, cache_type: CacheType) -> "CacheConfig":
|
||||
"""根据缓存类型返回预设配置"""
|
||||
configs = {
|
||||
CacheType.LOCATION: cls(
|
||||
strategy=CacheStrategy.TTL, ttl=None, max_size=1000 # 手动失效
|
||||
),
|
||||
CacheType.IP_INFO: cls(
|
||||
strategy=CacheStrategy.TTL, ttl=86400, max_size=1000 # 24小时
|
||||
),
|
||||
CacheType.WEATHER: cls(
|
||||
strategy=CacheStrategy.TTL, ttl=28800, max_size=1000 # 8小时
|
||||
),
|
||||
CacheType.LUNAR: cls(
|
||||
strategy=CacheStrategy.TTL, ttl=2592000, max_size=365 # 30天过期
|
||||
),
|
||||
CacheType.INTENT: cls(
|
||||
strategy=CacheStrategy.TTL_LRU, ttl=600, max_size=1000 # 10分钟
|
||||
),
|
||||
CacheType.CONFIG: cls(
|
||||
strategy=CacheStrategy.FIXED_SIZE, ttl=None, max_size=20 # 手动失效
|
||||
),
|
||||
CacheType.DEVICE_PROMPT: cls(
|
||||
strategy=CacheStrategy.TTL, ttl=None, max_size=1000 # 手动失效
|
||||
),
|
||||
}
|
||||
return configs.get(cache_type, cls())
|
||||
+216
@@ -0,0 +1,216 @@
|
||||
"""
|
||||
全局缓存管理器
|
||||
"""
|
||||
|
||||
import time
|
||||
import threading
|
||||
from typing import Any, Optional, Dict
|
||||
from collections import OrderedDict
|
||||
from .strategies import CacheStrategy, CacheEntry
|
||||
from .config import CacheConfig, CacheType
|
||||
|
||||
|
||||
class GlobalCacheManager:
|
||||
"""全局缓存管理器"""
|
||||
|
||||
def __init__(self):
|
||||
self._logger = None
|
||||
self._caches: Dict[str, Dict[str, CacheEntry]] = {}
|
||||
self._configs: Dict[str, CacheConfig] = {}
|
||||
self._locks: Dict[str, threading.RLock] = {}
|
||||
self._global_lock = threading.RLock()
|
||||
self._last_cleanup = time.time()
|
||||
self._stats = {"hits": 0, "misses": 0, "evictions": 0, "cleanups": 0}
|
||||
|
||||
@property
|
||||
def logger(self):
|
||||
"""延迟初始化 logger 以避免循环导入"""
|
||||
if self._logger is None:
|
||||
from config.logger import setup_logging
|
||||
|
||||
self._logger = setup_logging()
|
||||
return self._logger
|
||||
|
||||
def _get_cache_name(self, cache_type: CacheType, namespace: str = "") -> str:
|
||||
"""生成缓存名称"""
|
||||
if namespace:
|
||||
return f"{cache_type.value}:{namespace}"
|
||||
return cache_type.value
|
||||
|
||||
def _get_or_create_cache(
|
||||
self, cache_name: str, config: CacheConfig
|
||||
) -> Dict[str, CacheEntry]:
|
||||
"""获取或创建缓存空间"""
|
||||
with self._global_lock:
|
||||
if cache_name not in self._caches:
|
||||
self._caches[cache_name] = (
|
||||
OrderedDict()
|
||||
if config.strategy in [CacheStrategy.LRU, CacheStrategy.TTL_LRU]
|
||||
else {}
|
||||
)
|
||||
self._configs[cache_name] = config
|
||||
self._locks[cache_name] = threading.RLock()
|
||||
return self._caches[cache_name]
|
||||
|
||||
def set(
|
||||
self,
|
||||
cache_type: CacheType,
|
||||
key: str,
|
||||
value: Any,
|
||||
ttl: Optional[float] = None,
|
||||
namespace: str = "",
|
||||
) -> None:
|
||||
"""设置缓存值"""
|
||||
cache_name = self._get_cache_name(cache_type, namespace)
|
||||
config = self._configs.get(cache_name) or CacheConfig.for_type(cache_type)
|
||||
cache = self._get_or_create_cache(cache_name, config)
|
||||
|
||||
# 使用配置的TTL或传入的TTL
|
||||
effective_ttl = ttl if ttl is not None else config.ttl
|
||||
|
||||
with self._locks[cache_name]:
|
||||
# 创建缓存条目
|
||||
entry = CacheEntry(value=value, timestamp=time.time(), ttl=effective_ttl)
|
||||
|
||||
# 处理不同策略
|
||||
if config.strategy in [CacheStrategy.LRU, CacheStrategy.TTL_LRU]:
|
||||
# LRU策略:如果已存在则移动到末尾
|
||||
if key in cache:
|
||||
del cache[key]
|
||||
cache[key] = entry
|
||||
|
||||
# 检查大小限制
|
||||
if config.max_size and len(cache) > config.max_size:
|
||||
# 移除最旧的条目
|
||||
oldest_key = next(iter(cache))
|
||||
del cache[oldest_key]
|
||||
self._stats["evictions"] += 1
|
||||
|
||||
else:
|
||||
cache[key] = entry
|
||||
|
||||
# 检查大小限制
|
||||
if config.max_size and len(cache) > config.max_size:
|
||||
# 简单策略:随机移除一个条目
|
||||
victim_key = next(iter(cache))
|
||||
del cache[victim_key]
|
||||
self._stats["evictions"] += 1
|
||||
|
||||
# 定期清理过期条目
|
||||
self._maybe_cleanup(cache_name)
|
||||
|
||||
def get(
|
||||
self, cache_type: CacheType, key: str, namespace: str = ""
|
||||
) -> Optional[Any]:
|
||||
"""获取缓存值"""
|
||||
cache_name = self._get_cache_name(cache_type, namespace)
|
||||
|
||||
if cache_name not in self._caches:
|
||||
self._stats["misses"] += 1
|
||||
return None
|
||||
|
||||
cache = self._caches[cache_name]
|
||||
config = self._configs[cache_name]
|
||||
|
||||
with self._locks[cache_name]:
|
||||
if key not in cache:
|
||||
self._stats["misses"] += 1
|
||||
return None
|
||||
|
||||
entry = cache[key]
|
||||
|
||||
# 检查过期
|
||||
if entry.is_expired():
|
||||
del cache[key]
|
||||
self._stats["misses"] += 1
|
||||
return None
|
||||
|
||||
# 更新访问信息
|
||||
entry.touch()
|
||||
|
||||
# LRU策略:移动到末尾
|
||||
if config.strategy in [CacheStrategy.LRU, CacheStrategy.TTL_LRU]:
|
||||
del cache[key]
|
||||
cache[key] = entry
|
||||
|
||||
self._stats["hits"] += 1
|
||||
return entry.value
|
||||
|
||||
def delete(self, cache_type: CacheType, key: str, namespace: str = "") -> bool:
|
||||
"""删除缓存条目"""
|
||||
cache_name = self._get_cache_name(cache_type, namespace)
|
||||
|
||||
if cache_name not in self._caches:
|
||||
return False
|
||||
|
||||
cache = self._caches[cache_name]
|
||||
|
||||
with self._locks[cache_name]:
|
||||
if key in cache:
|
||||
del cache[key]
|
||||
return True
|
||||
return False
|
||||
|
||||
def clear(self, cache_type: CacheType, namespace: str = "") -> None:
|
||||
"""清空指定缓存"""
|
||||
cache_name = self._get_cache_name(cache_type, namespace)
|
||||
|
||||
if cache_name not in self._caches:
|
||||
return
|
||||
|
||||
with self._locks[cache_name]:
|
||||
self._caches[cache_name].clear()
|
||||
|
||||
def invalidate_pattern(
|
||||
self, cache_type: CacheType, pattern: str, namespace: str = ""
|
||||
) -> int:
|
||||
"""按模式失效缓存条目"""
|
||||
cache_name = self._get_cache_name(cache_type, namespace)
|
||||
|
||||
if cache_name not in self._caches:
|
||||
return 0
|
||||
|
||||
cache = self._caches[cache_name]
|
||||
deleted_count = 0
|
||||
|
||||
with self._locks[cache_name]:
|
||||
keys_to_delete = [key for key in cache.keys() if pattern in key]
|
||||
for key in keys_to_delete:
|
||||
del cache[key]
|
||||
deleted_count += 1
|
||||
|
||||
return deleted_count
|
||||
|
||||
def _cleanup_expired(self, cache_name: str) -> int:
|
||||
"""清理过期条目"""
|
||||
if cache_name not in self._caches:
|
||||
return 0
|
||||
|
||||
cache = self._caches[cache_name]
|
||||
deleted_count = 0
|
||||
|
||||
with self._locks[cache_name]:
|
||||
expired_keys = [key for key, entry in cache.items() if entry.is_expired()]
|
||||
for key in expired_keys:
|
||||
del cache[key]
|
||||
deleted_count += 1
|
||||
|
||||
return deleted_count
|
||||
|
||||
def _maybe_cleanup(self, cache_name: str):
|
||||
"""定期清理检查"""
|
||||
config = self._configs.get(cache_name)
|
||||
if not config:
|
||||
return
|
||||
|
||||
now = time.time()
|
||||
if now - self._last_cleanup > config.cleanup_interval:
|
||||
self._last_cleanup = now
|
||||
deleted = self._cleanup_expired(cache_name)
|
||||
if deleted > 0:
|
||||
self._stats["cleanups"] += 1
|
||||
self.logger.debug(f"清理缓存 {cache_name}: 删除 {deleted} 个过期条目")
|
||||
|
||||
|
||||
# 创建全局缓存管理器实例
|
||||
cache_manager = GlobalCacheManager()
|
||||
@@ -0,0 +1,43 @@
|
||||
"""
|
||||
缓存策略和数据结构定义
|
||||
"""
|
||||
|
||||
import time
|
||||
from enum import Enum
|
||||
from typing import Any, Optional
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
class CacheStrategy(Enum):
|
||||
"""缓存策略枚举"""
|
||||
|
||||
TTL = "ttl" # 基于时间过期
|
||||
LRU = "lru" # 最近最少使用
|
||||
FIXED_SIZE = "fixed_size" # 固定大小
|
||||
TTL_LRU = "ttl_lru" # TTL + LRU混合策略
|
||||
|
||||
|
||||
@dataclass
|
||||
class CacheEntry:
|
||||
"""缓存条目数据结构"""
|
||||
|
||||
value: Any
|
||||
timestamp: float
|
||||
ttl: Optional[float] = None # 生存时间(秒)
|
||||
access_count: int = 0
|
||||
last_access: float = None
|
||||
|
||||
def __post_init__(self):
|
||||
if self.last_access is None:
|
||||
self.last_access = self.timestamp
|
||||
|
||||
def is_expired(self) -> bool:
|
||||
"""检查是否过期"""
|
||||
if self.ttl is None:
|
||||
return False
|
||||
return time.time() - self.timestamp > self.ttl
|
||||
|
||||
def touch(self):
|
||||
"""更新访问时间和计数"""
|
||||
self.last_access = time.time()
|
||||
self.access_count += 1
|
||||
@@ -1,4 +1,5 @@
|
||||
import uuid
|
||||
import re
|
||||
from typing import List, Dict
|
||||
from datetime import datetime
|
||||
from config.settings import load_config
|
||||
@@ -74,13 +75,6 @@ class Dialogue:
|
||||
# 基础系统提示
|
||||
enhanced_system_prompt = system_message.content
|
||||
|
||||
# 添加说话人识别功能说明
|
||||
speaker_guidance = "\n\n[说话人识别功能说明]\n" \
|
||||
"当用户消息为JSON格式包含speaker字段时(如:{\"speaker\": \"张三\", \"content\": \"消息内容\"}),表示系统已识别出说话人身份。\n" \
|
||||
"请根据说话人的身份特征来调整回应风格和内容。\n" \
|
||||
"你可以称呼说话人的名字,并参考他们的特点进行个性化回应。"
|
||||
enhanced_system_prompt += speaker_guidance
|
||||
|
||||
# 添加说话人个性化描述
|
||||
try:
|
||||
config = load_config()
|
||||
@@ -88,7 +82,7 @@ class Dialogue:
|
||||
speakers = voiceprint_config.get("speakers", [])
|
||||
|
||||
if speakers:
|
||||
enhanced_system_prompt += "\n\n[已知说话人信息]"
|
||||
enhanced_system_prompt += "\n\n<speaker>"
|
||||
for speaker_str in speakers:
|
||||
try:
|
||||
parts = speaker_str.split(",", 2)
|
||||
@@ -99,15 +93,19 @@ class Dialogue:
|
||||
description = parts[2].strip() if len(parts) >= 3 else ""
|
||||
enhanced_system_prompt += f"\n- {name}:{description}"
|
||||
except:
|
||||
continue
|
||||
pass
|
||||
enhanced_system_prompt += "\n\n</speaker>"
|
||||
except:
|
||||
# 配置读取失败时忽略错误,不影响其他功能
|
||||
pass
|
||||
|
||||
# 只有当有记忆时才添加记忆部分
|
||||
if memory_str and len(memory_str) > 0:
|
||||
enhanced_system_prompt += f"\n\n以下是用户的历史记忆:\n```\n{memory_str}\n```"
|
||||
|
||||
# 使用正则表达式匹配 <memory> 标签,不管中间有什么内容
|
||||
enhanced_system_prompt = re.sub(
|
||||
r"<memory>.*?</memory>",
|
||||
f"<memory>\n{memory_str}\n</memory>",
|
||||
system_message.content,
|
||||
flags=re.DOTALL,
|
||||
)
|
||||
dialogue.append({"role": "system", "content": enhanced_system_prompt})
|
||||
|
||||
# 添加用户和助手的对话
|
||||
|
||||
@@ -0,0 +1,219 @@
|
||||
"""
|
||||
系统提示词管理器模块
|
||||
负责管理和更新系统提示词,包括快速初始化和异步增强功能
|
||||
"""
|
||||
|
||||
import os
|
||||
import cnlunar
|
||||
from typing import Dict, Any
|
||||
from config.logger import setup_logging
|
||||
|
||||
TAG = __name__
|
||||
|
||||
WEEKDAY_MAP = {
|
||||
"Monday": "星期一",
|
||||
"Tuesday": "星期二",
|
||||
"Wednesday": "星期三",
|
||||
"Thursday": "星期四",
|
||||
"Friday": "星期五",
|
||||
"Saturday": "星期六",
|
||||
"Sunday": "星期日",
|
||||
}
|
||||
|
||||
|
||||
class PromptManager:
|
||||
"""系统提示词管理器,负责管理和更新系统提示词"""
|
||||
|
||||
def __init__(self, config: Dict[str, Any], logger=None):
|
||||
self.config = config
|
||||
self.logger = logger or setup_logging()
|
||||
self.base_prompt_template = None
|
||||
self.last_update_time = 0
|
||||
|
||||
# 导入全局缓存管理器
|
||||
from core.utils.cache.manager import cache_manager, CacheType
|
||||
|
||||
self.cache_manager = cache_manager
|
||||
self.CacheType = CacheType
|
||||
|
||||
self._load_base_template()
|
||||
|
||||
def _load_base_template(self):
|
||||
"""加载基础提示词模板"""
|
||||
try:
|
||||
template_path = "agent-base-prompt.txt"
|
||||
cache_key = f"prompt_template:{template_path}"
|
||||
|
||||
# 先从缓存获取
|
||||
cached_template = self.cache_manager.get(self.CacheType.CONFIG, cache_key)
|
||||
if cached_template is not None:
|
||||
self.base_prompt_template = cached_template
|
||||
self.logger.bind(tag=TAG).debug("从缓存加载基础提示词模板")
|
||||
return
|
||||
|
||||
# 缓存未命中,从文件读取
|
||||
if os.path.exists(template_path):
|
||||
with open(template_path, "r", encoding="utf-8") as f:
|
||||
template_content = f.read()
|
||||
|
||||
# 存入缓存(CONFIG类型默认不自动过期,需要手动失效)
|
||||
self.cache_manager.set(
|
||||
self.CacheType.CONFIG, cache_key, template_content
|
||||
)
|
||||
self.base_prompt_template = template_content
|
||||
self.logger.bind(tag=TAG).debug("成功加载基础提示词模板并缓存")
|
||||
else:
|
||||
self.logger.bind(tag=TAG).warning("未找到agent-base-prompt.txt文件")
|
||||
except Exception as e:
|
||||
self.logger.bind(tag=TAG).error(f"加载提示词模板失败: {e}")
|
||||
|
||||
def get_quick_prompt(self, user_prompt: str, device_id: str = None) -> str:
|
||||
"""快速获取系统提示词(使用用户配置)"""
|
||||
device_cache_key = f"device_prompt:{device_id}"
|
||||
cached_device_prompt = self.cache_manager.get(
|
||||
self.CacheType.DEVICE_PROMPT, device_cache_key
|
||||
)
|
||||
if cached_device_prompt is not None:
|
||||
self.logger.bind(tag=TAG).debug(f"使用设备 {device_id} 的缓存提示词")
|
||||
return cached_device_prompt
|
||||
else:
|
||||
self.logger.bind(tag=TAG).debug(
|
||||
f"设备 {device_id} 无缓存提示词,使用传入的提示词"
|
||||
)
|
||||
|
||||
# 使用传入的提示词并缓存(如果有设备ID)
|
||||
if device_id:
|
||||
device_cache_key = f"device_prompt:{device_id}"
|
||||
self.cache_manager.set(self.CacheType.CONFIG, device_cache_key, user_prompt)
|
||||
self.logger.bind(tag=TAG).debug(f"设备 {device_id} 的提示词已缓存")
|
||||
|
||||
self.logger.bind(tag=TAG).info(f"使用快速提示词: {user_prompt[:50]}...")
|
||||
return user_prompt
|
||||
|
||||
def _get_current_time_info(self) -> tuple:
|
||||
"""获取当前时间信息"""
|
||||
from datetime import datetime
|
||||
|
||||
now = datetime.now()
|
||||
current_time = now.strftime("%H:%M")
|
||||
today_date = now.strftime("%Y-%m-%d")
|
||||
today_weekday = WEEKDAY_MAP[now.strftime("%A")]
|
||||
today_lunar = cnlunar.Lunar(now, godType="8char")
|
||||
lunar_date = "%s年%s%s\n" % (
|
||||
today_lunar.lunarYearCn,
|
||||
today_lunar.lunarMonthCn[:-1],
|
||||
today_lunar.lunarDayCn,
|
||||
)
|
||||
|
||||
return current_time, today_date, today_weekday, lunar_date
|
||||
|
||||
def _get_location_info(self, client_ip: str) -> str:
|
||||
"""获取位置信息"""
|
||||
try:
|
||||
# 先从缓存获取
|
||||
cached_location = self.cache_manager.get(self.CacheType.LOCATION, client_ip)
|
||||
if cached_location is not None:
|
||||
return cached_location
|
||||
|
||||
# 缓存未命中,调用API获取
|
||||
from core.utils.util import get_ip_info
|
||||
|
||||
ip_info = get_ip_info(client_ip, self.logger)
|
||||
city = ip_info.get("city", "未知位置")
|
||||
location = f"{city}"
|
||||
|
||||
# 存入缓存
|
||||
self.cache_manager.set(self.CacheType.LOCATION, client_ip, location)
|
||||
return location
|
||||
except Exception as e:
|
||||
self.logger.bind(tag=TAG).error(f"获取位置信息失败: {e}")
|
||||
return "未知位置"
|
||||
|
||||
def _get_weather_info(self, conn, location: str) -> str:
|
||||
"""获取天气信息"""
|
||||
try:
|
||||
# 先从缓存获取
|
||||
cached_weather = self.cache_manager.get(self.CacheType.WEATHER, location)
|
||||
if cached_weather is not None:
|
||||
return cached_weather
|
||||
|
||||
# 缓存未命中,调用get_weather函数获取
|
||||
from plugins_func.functions.get_weather import get_weather
|
||||
from plugins_func.register import ActionResponse
|
||||
|
||||
# 调用get_weather函数
|
||||
result = get_weather(conn, location=location, lang="zh_CN")
|
||||
if isinstance(result, ActionResponse):
|
||||
weather_report = result.result
|
||||
self.cache_manager.set(self.CacheType.WEATHER, location, weather_report)
|
||||
return weather_report
|
||||
return "天气信息获取失败"
|
||||
|
||||
except Exception as e:
|
||||
self.logger.bind(tag=TAG).error(f"获取天气信息失败: {e}")
|
||||
return "天气信息获取失败"
|
||||
|
||||
def update_context_info(self, conn, client_ip: str):
|
||||
"""同步更新上下文信息"""
|
||||
try:
|
||||
# 获取位置信息(使用全局缓存)
|
||||
local_address = self._get_location_info(client_ip)
|
||||
# 获取天气信息(使用全局缓存)
|
||||
self._get_weather_info(conn, local_address)
|
||||
self.logger.bind(tag=TAG).info(f"上下文信息更新完成")
|
||||
|
||||
except Exception as e:
|
||||
self.logger.bind(tag=TAG).error(f"更新上下文信息失败: {e}")
|
||||
|
||||
def build_enhanced_prompt(
|
||||
self, user_prompt: str, device_id: str, client_ip: str = None
|
||||
) -> str:
|
||||
"""构建增强的系统提示词"""
|
||||
if not self.base_prompt_template:
|
||||
return user_prompt
|
||||
|
||||
try:
|
||||
# 获取最新的时间信息(不缓存)
|
||||
current_time, today_date, today_weekday, lunar_date = (
|
||||
self._get_current_time_info()
|
||||
)
|
||||
|
||||
# 获取缓存的上下文信息
|
||||
local_address = ""
|
||||
weather_info = ""
|
||||
|
||||
if client_ip:
|
||||
# 获取位置信息(从全局缓存)
|
||||
local_address = (
|
||||
self.cache_manager.get(self.CacheType.LOCATION, client_ip) or ""
|
||||
)
|
||||
|
||||
# 获取天气信息(从全局缓存)
|
||||
if local_address:
|
||||
weather_info = (
|
||||
self.cache_manager.get(self.CacheType.WEATHER, local_address)
|
||||
or ""
|
||||
)
|
||||
|
||||
# 替换模板变量
|
||||
enhanced_prompt = self.base_prompt_template.format(
|
||||
base_prompt=user_prompt,
|
||||
current_time=current_time,
|
||||
today_date=today_date,
|
||||
today_weekday=today_weekday,
|
||||
lunar_date=lunar_date,
|
||||
local_address=local_address,
|
||||
weather_info=weather_info,
|
||||
)
|
||||
device_cache_key = f"device_prompt:{device_id}"
|
||||
self.cache_manager.set(
|
||||
self.CacheType.DEVICE_PROMPT, device_cache_key, enhanced_prompt
|
||||
)
|
||||
self.logger.bind(tag=TAG).info(
|
||||
f"构建增强提示词成功,长度: {len(enhanced_prompt)}"
|
||||
)
|
||||
return enhanced_prompt
|
||||
|
||||
except Exception as e:
|
||||
self.logger.bind(tag=TAG).error(f"构建增强提示词失败: {e}")
|
||||
return user_prompt
|
||||
@@ -96,11 +96,23 @@ def is_private_ip(ip_addr):
|
||||
|
||||
def get_ip_info(ip_addr, logger):
|
||||
try:
|
||||
# 导入全局缓存管理器
|
||||
from core.utils.cache.manager import cache_manager, CacheType
|
||||
|
||||
# 先从缓存获取
|
||||
cached_ip_info = cache_manager.get(CacheType.IP_INFO, ip_addr)
|
||||
if cached_ip_info is not None:
|
||||
return cached_ip_info
|
||||
|
||||
# 缓存未命中,调用API
|
||||
if is_private_ip(ip_addr):
|
||||
ip_addr = ""
|
||||
url = f"https://whois.pconline.com.cn/ipJson.jsp?json=true&ip={ip_addr}"
|
||||
resp = requests.get(url).json()
|
||||
ip_info = {"city": resp.get("city")}
|
||||
|
||||
# 存入缓存
|
||||
cache_manager.set(CacheType.IP_INFO, ip_addr, ip_info)
|
||||
return ip_info
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"Error getting client ip info: {e}")
|
||||
@@ -982,3 +994,28 @@ def sanitize_tool_name(name: str) -> str:
|
||||
"""Sanitize tool names for OpenAI compatibility."""
|
||||
# 支持中文、英文字母、数字、下划线和连字符
|
||||
return re.sub(r"[^a-zA-Z0-9_\-\u4e00-\u9fff]", "_", name)
|
||||
|
||||
|
||||
def validate_mcp_endpoint(mcp_endpoint: str) -> bool:
|
||||
"""
|
||||
校验MCP接入点格式
|
||||
|
||||
Args:
|
||||
mcp_endpoint: MCP接入点字符串
|
||||
|
||||
Returns:
|
||||
bool: 是否有效
|
||||
"""
|
||||
# 1. 检查是否以ws开头
|
||||
if not mcp_endpoint.startswith("ws"):
|
||||
return False
|
||||
|
||||
# 2. 检查是否包含key、call字样
|
||||
if "key" in mcp_endpoint.lower() or "call" in mcp_endpoint.lower():
|
||||
return False
|
||||
|
||||
# 3. 检查是否包含/mcp/字样
|
||||
if "/mcp/" not in mcp_endpoint:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
Reference in New Issue
Block a user