Merge branch 'mangger-api-voice-print' into py_test

This commit is contained in:
hrz
2025-07-09 16:51:10 +08:00
committed by GitHub
57 changed files with 2980 additions and 284 deletions
+58
View File
@@ -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
View File
@@ -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()
+43
View File
@@ -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
+11 -13
View File
@@ -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
+37
View File
@@ -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