mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-30 09:10:01 +08:00
Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
eb6d90ae88 | ||
|
|
4b8013d2d6 | ||
|
|
d528711641 | ||
|
|
1cc18bb195 | ||
|
|
74a9f3a843 | ||
|
|
e7f3c210df | ||
|
|
f94121080f | ||
|
|
761c8daac4 | ||
|
|
c667fc215e |
@@ -26,7 +26,7 @@ __plugin_meta__ = PluginMetadata(
|
||||
""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="0.1",
|
||||
version="0.2",
|
||||
plugin_type=PluginType.SUPERUSER,
|
||||
configs=[
|
||||
RegisterConfig(
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import contextlib
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
from pathlib import Path
|
||||
@@ -18,7 +19,47 @@ BAIDU_URL = "https://www.baidu.com/"
|
||||
GOOGLE_URL = "https://www.google.com/"
|
||||
|
||||
VERSION_FILE = Path() / "__version__"
|
||||
ARM_KEY = "aarch64"
|
||||
|
||||
|
||||
def get_arm_cpu_freq_safe():
|
||||
"""获取ARM设备CPU频率"""
|
||||
# 方法1: 优先从系统频率文件读取
|
||||
freq_files = [
|
||||
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_max_freq",
|
||||
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_max_freq",
|
||||
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_cur_freq",
|
||||
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_cur_freq",
|
||||
]
|
||||
|
||||
for freq_file in freq_files:
|
||||
try:
|
||||
with open(freq_file) as f:
|
||||
frequency = int(f.read().strip())
|
||||
return round(frequency / 1000000, 2) # 转换为GHz
|
||||
except (OSError, ValueError):
|
||||
continue
|
||||
|
||||
# 方法2: 解析/proc/cpuinfo
|
||||
with contextlib.suppress(OSError, FileNotFoundError, ValueError, PermissionError):
|
||||
with open("/proc/cpuinfo") as f:
|
||||
for line in f:
|
||||
if "CPU MHz" in line:
|
||||
freq = float(line.split(":")[1].strip())
|
||||
return round(freq / 1000, 2) # 转换为GHz
|
||||
# 方法3: 使用lscpu命令
|
||||
with contextlib.suppress(OSError, subprocess.SubprocessError, ValueError):
|
||||
env = os.environ.copy()
|
||||
env["LC_ALL"] = "C"
|
||||
result = subprocess.run(
|
||||
["lscpu"], capture_output=True, text=True, env=env, timeout=10
|
||||
)
|
||||
|
||||
if result.returncode == 0:
|
||||
for line in result.stdout.split("\n"):
|
||||
if "CPU max MHz" in line or "CPU MHz" in line:
|
||||
freq = float(line.split(":")[1].strip())
|
||||
return round(freq / 1000, 2) # 转换为GHz
|
||||
return 0 # 如果所有方法都失败,返回0
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -37,7 +78,7 @@ class CPUInfo:
|
||||
if _cpu_freq := psutil.cpu_freq():
|
||||
cpu_freq = round(_cpu_freq.current / 1000, 2)
|
||||
else:
|
||||
cpu_freq = 0
|
||||
cpu_freq = get_arm_cpu_freq_safe()
|
||||
return CPUInfo(core=cpu_core, usage=cpu_usage, freq=cpu_freq)
|
||||
|
||||
|
||||
@@ -160,44 +201,13 @@ def __get_version() -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def __get_arm_cpu():
|
||||
env = os.environ.copy()
|
||||
env["LC_ALL"] = "en_US.UTF-8"
|
||||
cpu_info = subprocess.check_output(["lscpu"], env=env).decode()
|
||||
model_name = ""
|
||||
cpu_freq = 0
|
||||
for line in cpu_info.splitlines():
|
||||
if "Model name" in line:
|
||||
model_name = line.split(":")[1].strip()
|
||||
if "CPU MHz" in line:
|
||||
cpu_freq = float(line.split(":")[1].strip())
|
||||
return model_name, cpu_freq
|
||||
|
||||
|
||||
def __get_arm_oracle_cpu_freq():
|
||||
cpu_freq = subprocess.check_output(
|
||||
["dmidecode", "-s", "processor-frequency"]
|
||||
).decode()
|
||||
return round(float(cpu_freq.split()[0]) / 1000, 2)
|
||||
|
||||
|
||||
async def get_status_info() -> dict:
|
||||
"""获取信息"""
|
||||
data = await __build_status()
|
||||
|
||||
system = platform.uname()
|
||||
if system.machine == ARM_KEY and not (
|
||||
cpuinfo.get_cpu_info().get("brand_raw") and data.cpu.freq
|
||||
):
|
||||
model_name, cpu_freq = __get_arm_cpu()
|
||||
if not data.cpu.freq:
|
||||
data.cpu.freq = cpu_freq or __get_arm_oracle_cpu_freq()
|
||||
data = data.get_system_info()
|
||||
data["brand_raw"] = model_name
|
||||
else:
|
||||
data = data.get_system_info()
|
||||
data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
|
||||
|
||||
data = data.get_system_info()
|
||||
data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
|
||||
baidu, google = await __get_network_info()
|
||||
data["baidu"] = "#8CC265" if baidu else "red"
|
||||
data["google"] = "#8CC265" if google else "red"
|
||||
|
||||
@@ -74,8 +74,8 @@ async def _(matcher: Matcher, message: UniMsg, session: EventSession):
|
||||
message_list.append(image)
|
||||
message_list.append(
|
||||
"桀桀桀,预判到会有 '笨蛋' 把功能名称当命令用,特地前来嘲笑!"
|
||||
f"但还是好心来帮帮你啦!\n请at我发送 '帮助{plugin.name}' 或者"
|
||||
f" '帮助{plugin.id}' 来获取该功能帮助!"
|
||||
f"但还是好心来帮帮你啦!\n请at我发送 '帮助 {plugin.name}' 或者"
|
||||
f" '帮助 {plugin.id}' 来获取该功能帮助!"
|
||||
)
|
||||
logger.info("检测到功能名称当命令使用,已发送帮助信息", "功能帮助", session=session)
|
||||
await MessageUtils.build_message(message_list).send(reply_to=True)
|
||||
|
||||
@@ -58,5 +58,14 @@ Config.add_plugin_config(
|
||||
type=bool,
|
||||
)
|
||||
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
"AUTH_HOOKS_CONCURRENCY_LIMIT",
|
||||
5,
|
||||
help="同步进入权限钩子最大并发数",
|
||||
default_value=5,
|
||||
type=int,
|
||||
)
|
||||
|
||||
|
||||
nonebot.load_plugins(str(Path(__file__).parent.resolve()))
|
||||
|
||||
@@ -96,7 +96,6 @@ async def is_ban(user_id: str | None, group_id: str | None) -> int:
|
||||
f"查询ban记录超时: user_id={user_id}, group_id={group_id}",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
# 超时时返回0,避免阻塞
|
||||
return 0
|
||||
|
||||
# 检查记录并计算ban时间
|
||||
@@ -199,7 +198,7 @@ async def group_handle(group_id: str) -> None:
|
||||
)
|
||||
|
||||
|
||||
async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
|
||||
async def user_handle(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) -> None:
|
||||
"""用户ban检查
|
||||
|
||||
参数:
|
||||
@@ -217,22 +216,12 @@ async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
|
||||
if not time_val:
|
||||
return
|
||||
time_str = format_time(time_val)
|
||||
plugin_dao = DataAccess(PluginInfo)
|
||||
try:
|
||||
db_plugin = await asyncio.wait_for(
|
||||
plugin_dao.safe_get_or_none(module=module), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"查询插件信息超时: {module}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
raise SkipPluginException("用户处于黑名单中...")
|
||||
|
||||
if (
|
||||
db_plugin
|
||||
and not db_plugin.ignore_prompt
|
||||
plugin
|
||||
and time_val != -1
|
||||
and ban_result
|
||||
and freq.is_send_limit_message(db_plugin, entity.user_id, False)
|
||||
and freq.is_send_limit_message(plugin, entity.user_id, False)
|
||||
):
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
@@ -260,7 +249,9 @@ async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
|
||||
)
|
||||
|
||||
|
||||
async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
|
||||
async def auth_ban(
|
||||
matcher: Matcher, bot: Bot, session: Uninfo, plugin: PluginInfo
|
||||
) -> None:
|
||||
"""权限检查 - ban 检查
|
||||
|
||||
参数:
|
||||
@@ -289,7 +280,7 @@ async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
|
||||
if entity.user_id:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_handle(matcher.plugin_name, entity, session),
|
||||
user_handle(plugin, entity, session),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
|
||||
@@ -1,50 +1,36 @@
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from nonebot_plugin_alconna import UniMsg
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.utils import EntityIDs
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
|
||||
async def auth_group(
|
||||
plugin: PluginInfo,
|
||||
group: GroupConsole | None,
|
||||
message: UniMsg,
|
||||
group_id: str | None,
|
||||
):
|
||||
"""群黑名单检测 群总开关检测
|
||||
|
||||
参数:
|
||||
plugin: PluginInfo
|
||||
entity: EntityIDs
|
||||
group: GroupConsole
|
||||
message: UniMsg
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
if not entity.group_id:
|
||||
if not group_id:
|
||||
return
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
text = message.extract_plain_text()
|
||||
|
||||
# 从数据库或缓存中获取群组信息
|
||||
group_dao = DataAccess(GroupConsole)
|
||||
|
||||
try:
|
||||
group: GroupConsole | None = await asyncio.wait_for(
|
||||
group_dao.safe_get_or_none(
|
||||
group_id=entity.group_id, channel_id__isnull=True
|
||||
),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("查询群组信息超时", LOGGER_COMMAND, session=entity.user_id)
|
||||
# 超时时不阻塞,继续执行
|
||||
return
|
||||
|
||||
if not group:
|
||||
raise SkipPluginException("群组信息不存在...")
|
||||
if group.level < 0:
|
||||
@@ -63,6 +49,5 @@ async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
|
||||
logger.warning(
|
||||
f"auth_group 耗时: {elapsed:.3f}s, plugin={plugin.module}",
|
||||
LOGGER_COMMAND,
|
||||
session=entity.user_id,
|
||||
group_id=entity.group_id,
|
||||
group_id=group_id,
|
||||
)
|
||||
|
||||
@@ -6,12 +6,10 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.enum import BlockType
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .exception import IsSuperuserException, SkipPluginException
|
||||
@@ -20,30 +18,17 @@ from .utils import freq, is_poke, send_message
|
||||
|
||||
class GroupCheck:
|
||||
def __init__(
|
||||
self, plugin: PluginInfo, group_id: str, session: Uninfo, is_poke: bool
|
||||
self, plugin: PluginInfo, group: GroupConsole, session: Uninfo, is_poke: bool
|
||||
) -> None:
|
||||
self.group_id = group_id
|
||||
self.session = session
|
||||
self.is_poke = is_poke
|
||||
self.plugin = plugin
|
||||
self.group_dao = DataAccess(GroupConsole)
|
||||
self.group_data = None
|
||||
self.group_data = group
|
||||
self.group_id = group.group_id
|
||||
|
||||
async def check(self):
|
||||
start_time = time.time()
|
||||
try:
|
||||
# 只查询一次数据库,使用 DataAccess 的缓存机制
|
||||
try:
|
||||
self.group_data = await asyncio.wait_for(
|
||||
self.group_dao.safe_get_or_none(
|
||||
group_id=self.group_id, channel_id__isnull=True
|
||||
),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"查询群组数据超时: {self.group_id}", LOGGER_COMMAND)
|
||||
return # 超时时不阻塞,继续执行
|
||||
|
||||
# 检查超级用户禁用
|
||||
if (
|
||||
self.group_data
|
||||
@@ -113,12 +98,13 @@ class GroupCheck:
|
||||
|
||||
|
||||
class PluginCheck:
|
||||
def __init__(self, group_id: str | None, session: Uninfo, is_poke: bool):
|
||||
def __init__(self, group: GroupConsole | None, session: Uninfo, is_poke: bool):
|
||||
self.session = session
|
||||
self.is_poke = is_poke
|
||||
self.group_id = group_id
|
||||
self.group_dao = DataAccess(GroupConsole)
|
||||
self.group_data = None
|
||||
self.group_data = group
|
||||
self.group_id = None
|
||||
if group:
|
||||
self.group_id = group.group_id
|
||||
|
||||
async def check_user(self, plugin: PluginInfo):
|
||||
"""全局私聊禁用检测
|
||||
@@ -156,21 +142,8 @@ class PluginCheck:
|
||||
if plugin.status or plugin.block_type != BlockType.ALL:
|
||||
return
|
||||
"""全局状态"""
|
||||
if self.group_id:
|
||||
# 使用 DataAccess 的缓存机制
|
||||
try:
|
||||
self.group_data = await asyncio.wait_for(
|
||||
self.group_dao.safe_get_or_none(
|
||||
group_id=self.group_id, channel_id__isnull=True
|
||||
),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"查询群组数据超时: {self.group_id}", LOGGER_COMMAND)
|
||||
return # 超时时不阻塞,继续执行
|
||||
|
||||
if self.group_data and self.group_data.is_super:
|
||||
raise IsSuperuserException()
|
||||
if self.group_data and self.group_data.is_super:
|
||||
raise IsSuperuserException()
|
||||
|
||||
sid = self.group_id or self.session.user.id
|
||||
if freq.is_send_limit_message(plugin, sid, self.is_poke):
|
||||
@@ -193,7 +166,9 @@ class PluginCheck:
|
||||
)
|
||||
|
||||
|
||||
async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
|
||||
async def auth_plugin(
|
||||
plugin: PluginInfo, group: GroupConsole | None, session: Uninfo, event: Event
|
||||
):
|
||||
"""插件状态
|
||||
|
||||
参数:
|
||||
@@ -203,35 +178,23 @@ async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
|
||||
"""
|
||||
start_time = time.time()
|
||||
try:
|
||||
entity = get_entity_ids(session)
|
||||
is_poke_event = is_poke(event)
|
||||
user_check = PluginCheck(entity.group_id, session, is_poke_event)
|
||||
user_check = PluginCheck(group, session, is_poke_event)
|
||||
|
||||
if entity.group_id:
|
||||
group_check = GroupCheck(plugin, entity.group_id, session, is_poke_event)
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
group_check.check(), timeout=DB_TIMEOUT_SECONDS * 2
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"群组检查超时: {entity.group_id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
tasks = []
|
||||
if group:
|
||||
tasks.append(GroupCheck(plugin, group, session, is_poke_event).check())
|
||||
else:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_check.check_user(plugin), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("用户检查超时", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
tasks.append(user_check.check_user(plugin))
|
||||
tasks.append(user_check.check_global(plugin))
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_check.check_global(plugin), timeout=DB_TIMEOUT_SECONDS
|
||||
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS * 2
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("全局检查超时", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
logger.error("插件用户/群组/全局检查超时...", LOGGER_COMMAND)
|
||||
|
||||
finally:
|
||||
# 记录总执行时间
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
@@ -85,7 +85,7 @@ class FreqUtils:
|
||||
return False
|
||||
if plugin.plugin_type == PluginType.DEPENDANT:
|
||||
return False
|
||||
return plugin.module != "ai" if self._flmt_s.check(sid) else False
|
||||
return False if plugin.ignore_prompt else self._flmt_s.check(sid)
|
||||
|
||||
|
||||
freq = FreqUtils()
|
||||
|
||||
@@ -8,6 +8,7 @@ from nonebot_plugin_alconna import UniMsg
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
from tortoise.exceptions import IntegrityError
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
@@ -31,6 +32,7 @@ from .auth.exception import (
|
||||
PermissionExemption,
|
||||
SkipPluginException,
|
||||
)
|
||||
from .auth.utils import base_config
|
||||
|
||||
# 超时设置(秒)
|
||||
TIMEOUT_SECONDS = 5.0
|
||||
@@ -46,6 +48,16 @@ CIRCUIT_BREAKERS = {
|
||||
# 熔断重置时间(秒)
|
||||
CIRCUIT_RESET_TIME = 300 # 5分钟
|
||||
|
||||
# 并发控制:限制同时进入 hooks 并行检查的协程数
|
||||
|
||||
# 默认为 6,可通过环境变量 AUTH_HOOKS_CONCURRENCY_LIMIT 调整
|
||||
HOOKS_CONCURRENCY_LIMIT = base_config.get("AUTH_HOOKS_CONCURRENCY_LIMIT")
|
||||
|
||||
# 全局信号量与计数器
|
||||
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT)
|
||||
HOOKS_ACTIVE_COUNT = 0
|
||||
HOOKS_ACTIVE_LOCK = asyncio.Lock()
|
||||
|
||||
|
||||
# 超时装饰器
|
||||
async def with_timeout(coro, timeout=TIMEOUT_SECONDS, name=None):
|
||||
@@ -259,6 +271,30 @@ async def time_hook(coro, name, time_dict):
|
||||
time_dict[name] = f"{time.time() - start:.3f}s"
|
||||
|
||||
|
||||
async def _enter_hooks_section():
|
||||
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。"""
|
||||
global HOOKS_ACTIVE_COUNT
|
||||
# 队列模式:如果达到上限,协程将排队等待直到获取到信号量
|
||||
await HOOKS_SEMAPHORE.acquire()
|
||||
async with HOOKS_ACTIVE_LOCK:
|
||||
HOOKS_ACTIVE_COUNT += 1
|
||||
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
|
||||
|
||||
|
||||
async def _leave_hooks_section():
|
||||
"""释放信号量并更新计数器。"""
|
||||
global HOOKS_ACTIVE_COUNT
|
||||
from contextlib import suppress
|
||||
|
||||
with suppress(Exception):
|
||||
HOOKS_SEMAPHORE.release()
|
||||
async with HOOKS_ACTIVE_LOCK:
|
||||
HOOKS_ACTIVE_COUNT -= 1
|
||||
# 保证计数不为负
|
||||
HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT, 0)
|
||||
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
|
||||
|
||||
|
||||
async def auth(
|
||||
matcher: Matcher,
|
||||
event: Event,
|
||||
@@ -285,6 +321,9 @@ async def auth(
|
||||
hook_times = {}
|
||||
hooks_time = 0 # 初始化 hooks_time 变量
|
||||
|
||||
# 记录是否已进入 hooks 区域(用于 finally 中释放)
|
||||
entered_hooks = False
|
||||
|
||||
try:
|
||||
if not module:
|
||||
raise PermissionExemption("Matcher插件名称不存在...")
|
||||
@@ -304,6 +343,10 @@ async def auth(
|
||||
)
|
||||
raise PermissionExemption("获取插件和用户数据超时,请稍后再试...")
|
||||
|
||||
# 进入 hooks 并行检查区域(会在高并发时排队)
|
||||
await _enter_hooks_section()
|
||||
entered_hooks = True
|
||||
|
||||
# 获取插件费用
|
||||
cost_start = time.time()
|
||||
try:
|
||||
@@ -320,16 +363,32 @@ async def auth(
|
||||
# 执行 bot_filter
|
||||
bot_filter(session)
|
||||
|
||||
group = None
|
||||
if entity.group_id:
|
||||
group_dao = DataAccess(GroupConsole)
|
||||
group = await with_timeout(
|
||||
group_dao.safe_get_or_none(
|
||||
group_id=entity.group_id, channel_id__isnull=True
|
||||
),
|
||||
name="get_group",
|
||||
)
|
||||
|
||||
# 并行执行所有 hook 检查,并记录执行时间
|
||||
hooks_start = time.time()
|
||||
|
||||
# 创建所有 hook 任务
|
||||
hook_tasks = [
|
||||
time_hook(auth_ban(matcher, bot, session), "auth_ban", hook_times),
|
||||
time_hook(auth_ban(matcher, bot, session, plugin), "auth_ban", hook_times),
|
||||
time_hook(auth_bot(plugin, bot.self_id), "auth_bot", hook_times),
|
||||
time_hook(auth_group(plugin, entity, message), "auth_group", hook_times),
|
||||
time_hook(
|
||||
auth_group(plugin, group, message, entity.group_id),
|
||||
"auth_group",
|
||||
hook_times,
|
||||
),
|
||||
time_hook(auth_admin(plugin, session), "auth_admin", hook_times),
|
||||
time_hook(auth_plugin(plugin, session, event), "auth_plugin", hook_times),
|
||||
time_hook(
|
||||
auth_plugin(plugin, group, session, event), "auth_plugin", hook_times
|
||||
),
|
||||
time_hook(auth_limit(plugin, session), "auth_limit", hook_times),
|
||||
]
|
||||
|
||||
@@ -358,7 +417,17 @@ async def auth(
|
||||
logger.debug("超级用户跳过权限检测...", LOGGER_COMMAND, session=session)
|
||||
except PermissionExemption as e:
|
||||
logger.info(str(e), LOGGER_COMMAND, session=session)
|
||||
|
||||
finally:
|
||||
# 如果进入过 hooks 区域,确保释放信号量(即使上层处理抛出了异常)
|
||||
if entered_hooks:
|
||||
try:
|
||||
await _leave_hooks_section()
|
||||
except Exception:
|
||||
logger.error(
|
||||
"释放 hooks 信号量时出错",
|
||||
LOGGER_COMMAND,
|
||||
session=session,
|
||||
)
|
||||
# 扣除金币
|
||||
if not ignore_flag and cost_gold > 0:
|
||||
gold_start = time.time()
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
from typing import Any
|
||||
|
||||
from nonebot.adapters import Bot, Message
|
||||
from nonebot.adapters.onebot.v11 import MessageSegment
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.bot_message_store import BotMessageStore
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import BotSentType
|
||||
from zhenxun.utils.log_sanitizer import sanitize_for_logging
|
||||
from zhenxun.utils.manager.message_manager import MessageManager
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
@@ -41,35 +41,6 @@ def replace_message(message: Message) -> str:
|
||||
return result
|
||||
|
||||
|
||||
def format_message_for_log(message: Message) -> str:
|
||||
"""
|
||||
将消息对象转换为适合日志记录的字符串,对base64等长内容进行摘要处理。
|
||||
"""
|
||||
if not isinstance(message, Message):
|
||||
return str(message)
|
||||
|
||||
log_parts = []
|
||||
for seg in message:
|
||||
seg: MessageSegment
|
||||
if seg.type == "text":
|
||||
log_parts.append(seg.data.get("text", ""))
|
||||
elif seg.type in ("image", "record", "video"):
|
||||
file_info = seg.data.get("file", "")
|
||||
if isinstance(file_info, str) and file_info.startswith("base64://"):
|
||||
b64_data = file_info[9:]
|
||||
data_size_bytes = (len(b64_data) * 3) / 4 - b64_data.count("=", -2)
|
||||
log_parts.append(
|
||||
f"[{seg.type}: base64, size={data_size_bytes / 1024:.2f}KB]"
|
||||
)
|
||||
else:
|
||||
log_parts.append(f"[{seg.type}]")
|
||||
elif seg.type == "at":
|
||||
log_parts.append(f"[@{seg.data.get('qq', 'unknown')}]")
|
||||
else:
|
||||
log_parts.append(f"[{seg.type}]")
|
||||
return "".join(log_parts)
|
||||
|
||||
|
||||
@Bot.on_called_api
|
||||
async def handle_api_result(
|
||||
bot: Bot, exception: Exception | None, api: str, data: dict[str, Any], result: Any
|
||||
@@ -82,7 +53,6 @@ async def handle_api_result(
|
||||
message: Message = data.get("message", "")
|
||||
message_type = data.get("message_type")
|
||||
try:
|
||||
# 记录消息id
|
||||
if user_id and message_id:
|
||||
MessageManager.add(str(user_id), str(message_id))
|
||||
logger.debug(
|
||||
@@ -108,7 +78,8 @@ async def handle_api_result(
|
||||
else replace_message(message),
|
||||
platform=PlatformUtils.get_platform(bot),
|
||||
)
|
||||
logger.debug(f"消息发送记录,message: {format_message_for_log(message)}")
|
||||
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
|
||||
logger.debug(f"消息发送记录,message: {sanitized_message}")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"消息发送记录发生错误...data: {data}, result: {result}",
|
||||
|
||||
@@ -43,18 +43,20 @@ class BanCheckLimiter:
|
||||
|
||||
def check(self, key: str | float) -> bool:
|
||||
if time.time() - self.mtime[key] > self.default_check_time:
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return False
|
||||
return self._extracted_from_check_3(key, False)
|
||||
if (
|
||||
self.mint[key] >= self.default_count
|
||||
and time.time() - self.mtime[key] < self.default_check_time
|
||||
):
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return True
|
||||
return self._extracted_from_check_3(key, True)
|
||||
return False
|
||||
|
||||
# TODO Rename this here and in `check`
|
||||
def _extracted_from_check_3(self, key, arg1):
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return arg1
|
||||
|
||||
|
||||
_blmt = BanCheckLimiter(
|
||||
malicious_check_time,
|
||||
@@ -70,16 +72,15 @@ async def _(
|
||||
module = None
|
||||
if plugin := matcher.plugin:
|
||||
module = plugin.module_name
|
||||
if metadata := plugin.metadata:
|
||||
extra = metadata.extra
|
||||
if extra.get("plugin_type") in [
|
||||
PluginType.HIDDEN,
|
||||
PluginType.DEPENDANT,
|
||||
PluginType.ADMIN,
|
||||
PluginType.SUPERUSER,
|
||||
]:
|
||||
return
|
||||
else:
|
||||
if not (metadata := plugin.metadata):
|
||||
return
|
||||
extra = metadata.extra
|
||||
if extra.get("plugin_type") in [
|
||||
PluginType.HIDDEN,
|
||||
PluginType.DEPENDANT,
|
||||
PluginType.ADMIN,
|
||||
PluginType.SUPERUSER,
|
||||
]:
|
||||
return
|
||||
if matcher.type == "notice":
|
||||
return
|
||||
@@ -88,32 +89,31 @@ async def _(
|
||||
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
|
||||
if not malicious_ban_time:
|
||||
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0")
|
||||
if user_id:
|
||||
if module:
|
||||
if _blmt.check(f"{user_id}__{module}"):
|
||||
await BanConsole.ban(
|
||||
user_id,
|
||||
group_id,
|
||||
9,
|
||||
"恶意触发命令检测",
|
||||
malicious_ban_time * 60,
|
||||
bot.self_id,
|
||||
)
|
||||
logger.info(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
await MessageUtils.build_message(
|
||||
[
|
||||
At(flag="user", target=user_id),
|
||||
"检测到恶意触发命令,您将被封禁 30 分钟",
|
||||
]
|
||||
).send()
|
||||
logger.debug(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
raise IgnoredException("检测到恶意触发命令")
|
||||
_blmt.add(f"{user_id}__{module}")
|
||||
if user_id and module:
|
||||
if _blmt.check(f"{user_id}__{module}"):
|
||||
await BanConsole.ban(
|
||||
user_id,
|
||||
group_id,
|
||||
9,
|
||||
"恶意触发命令检测",
|
||||
malicious_ban_time * 60,
|
||||
bot.self_id,
|
||||
)
|
||||
logger.info(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
await MessageUtils.build_message(
|
||||
[
|
||||
At(flag="user", target=user_id),
|
||||
"检测到恶意触发命令,您将被封禁 30 分钟",
|
||||
]
|
||||
).send()
|
||||
logger.debug(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
raise IgnoredException("检测到恶意触发命令")
|
||||
_blmt.add(f"{user_id}__{module}")
|
||||
|
||||
@@ -263,10 +263,9 @@ class StoreManager:
|
||||
"""安装插件
|
||||
|
||||
参数:
|
||||
github_url: 仓库地址
|
||||
module_path: 模块路径
|
||||
is_dir: 是否是文件夹
|
||||
plugin_info: 插件信息
|
||||
is_external: 是否是外部仓库
|
||||
source: 源
|
||||
"""
|
||||
repo_type = RepoType.GITHUB if is_external else None
|
||||
if source == "ali":
|
||||
|
||||
@@ -367,7 +367,7 @@ class ShopManage:
|
||||
else:
|
||||
goods_info = await GoodsInfo.get_or_none(goods_name=goods_name)
|
||||
if not goods_info:
|
||||
return f"{goods_name} 不存在..."
|
||||
return "对应的道具不存在..."
|
||||
if goods_info.is_passive:
|
||||
return f"{goods_info.goods_name} 是被动道具, 无法使用..."
|
||||
goods = cls.uuid2goods.get(goods_info.uuid)
|
||||
|
||||
@@ -344,7 +344,9 @@ class ConfigsManager:
|
||||
返回:
|
||||
ConfigGroup: ConfigGroup
|
||||
"""
|
||||
return self._data.get(key) or ConfigGroup(module="")
|
||||
if key not in self._data:
|
||||
self._data[key] = ConfigGroup(module=key)
|
||||
return self._data[key]
|
||||
|
||||
def save(self, path: str | Path | None = None, save_simple_data: bool = False):
|
||||
"""保存数据
|
||||
|
||||
Vendored
+2
-2
@@ -98,6 +98,7 @@ from .cache_containers import CacheDict, CacheList
|
||||
from .config import (
|
||||
CACHE_KEY_PREFIX,
|
||||
CACHE_KEY_SEPARATOR,
|
||||
CACHE_TIMEOUT,
|
||||
DEFAULT_EXPIRE,
|
||||
LOG_COMMAND,
|
||||
SPECIAL_KEY_FORMATS,
|
||||
@@ -551,7 +552,6 @@ class CacheManager:
|
||||
返回:
|
||||
Any: 缓存数据,如果不存在返回默认值
|
||||
"""
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
|
||||
# 如果缓存被禁用或缓存模式为NONE,直接返回默认值
|
||||
if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
|
||||
@@ -561,7 +561,7 @@ class CacheManager:
|
||||
cache_key = self._build_key(cache_type, key)
|
||||
data = await asyncio.wait_for(
|
||||
self.cache_backend.get(cache_key), # type: ignore
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
timeout=CACHE_TIMEOUT,
|
||||
)
|
||||
|
||||
if data is None:
|
||||
|
||||
Vendored
+3
@@ -5,6 +5,9 @@
|
||||
# 日志标识
|
||||
LOG_COMMAND = "CacheRoot"
|
||||
|
||||
# 缓存获取超时时间(秒)
|
||||
CACHE_TIMEOUT = 10
|
||||
|
||||
# 默认缓存过期时间(秒)
|
||||
DEFAULT_EXPIRE = 600
|
||||
|
||||
|
||||
@@ -27,5 +27,8 @@ async def with_db_timeout(
|
||||
return result
|
||||
except asyncio.TimeoutError:
|
||||
if operation:
|
||||
logger.error(f"数据库操作超时: {operation} (>{timeout}s)", LOG_COMMAND)
|
||||
logger.error(
|
||||
f"数据库操作超时: {operation} (>{timeout}s) 来源: {source}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -7,6 +7,7 @@ LLM 服务模块 - 公共 API 入口
|
||||
from .api import (
|
||||
chat,
|
||||
code,
|
||||
create_image,
|
||||
embed,
|
||||
generate,
|
||||
generate_structured,
|
||||
@@ -74,6 +75,7 @@ __all__ = [
|
||||
"chat",
|
||||
"clear_model_cache",
|
||||
"code",
|
||||
"create_image",
|
||||
"create_multimodal_message",
|
||||
"embed",
|
||||
"function_tool",
|
||||
|
||||
@@ -3,6 +3,9 @@ LLM 适配器基类和通用数据结构
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
import base64
|
||||
import binascii
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from pydantic import BaseModel
|
||||
@@ -32,6 +35,7 @@ class ResponseData(BaseModel):
|
||||
"""响应数据封装 - 支持所有高级功能"""
|
||||
|
||||
text: str
|
||||
images: list[bytes] | None = None
|
||||
usage_info: dict[str, Any] | None = None
|
||||
raw_response: dict[str, Any] | None = None
|
||||
tool_calls: list[LLMToolCall] | None = None
|
||||
@@ -242,6 +246,38 @@ class BaseAdapter(ABC):
|
||||
if content:
|
||||
content = content.strip()
|
||||
|
||||
images_bytes: list[bytes] = []
|
||||
if content and content.startswith("{") and content.endswith("}"):
|
||||
try:
|
||||
content_json = json.loads(content)
|
||||
if "b64_json" in content_json:
|
||||
images_bytes.append(base64.b64decode(content_json["b64_json"]))
|
||||
content = "[图片已生成]"
|
||||
elif "data" in content_json and isinstance(
|
||||
content_json["data"], str
|
||||
):
|
||||
images_bytes.append(base64.b64decode(content_json["data"]))
|
||||
content = "[图片已生成]"
|
||||
|
||||
except (json.JSONDecodeError, KeyError, binascii.Error):
|
||||
pass
|
||||
elif (
|
||||
"images" in message
|
||||
and isinstance(message["images"], list)
|
||||
and message["images"]
|
||||
):
|
||||
image_info = message["images"][0]
|
||||
if image_info.get("type") == "image_url":
|
||||
image_url_obj = image_info.get("image_url", {})
|
||||
url_str = image_url_obj.get("url", "")
|
||||
if url_str.startswith("data:image/png;base64,"):
|
||||
try:
|
||||
b64_data = url_str.split(",", 1)[1]
|
||||
images_bytes.append(base64.b64decode(b64_data))
|
||||
content = content if content else "[图片已生成]"
|
||||
except (IndexError, binascii.Error) as e:
|
||||
logger.warning(f"解析OpenRouter Base64图片数据失败: {e}")
|
||||
|
||||
parsed_tool_calls: list[LLMToolCall] | None = None
|
||||
if message_tool_calls := message.get("tool_calls"):
|
||||
from ..types.models import LLMToolFunction
|
||||
@@ -280,6 +316,7 @@ class BaseAdapter(ABC):
|
||||
text=final_text,
|
||||
tool_calls=parsed_tool_calls,
|
||||
usage_info=usage_info,
|
||||
images=images_bytes if images_bytes else None,
|
||||
raw_response=response_json,
|
||||
)
|
||||
|
||||
@@ -450,6 +487,13 @@ class OpenAICompatAdapter(BaseAdapter):
|
||||
"""准备高级请求 - OpenAI兼容格式"""
|
||||
url = self.get_api_url(model, self.get_chat_endpoint(model))
|
||||
headers = self.get_base_headers(api_key)
|
||||
if model.api_type == "openrouter":
|
||||
headers.update(
|
||||
{
|
||||
"HTTP-Referer": "https://github.com/zhenxun-org/zhenxun_bot",
|
||||
"X-Title": "Zhenxun Bot",
|
||||
}
|
||||
)
|
||||
openai_messages = self.convert_messages_to_openai_format(messages)
|
||||
|
||||
body = {
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
Gemini API 适配器
|
||||
"""
|
||||
|
||||
import base64
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
@@ -373,7 +374,16 @@ class GeminiAdapter(BaseAdapter):
|
||||
self.validate_response(response_json)
|
||||
|
||||
try:
|
||||
candidates = response_json.get("candidates", [])
|
||||
if "image_generation" in response_json and isinstance(
|
||||
response_json["image_generation"], dict
|
||||
):
|
||||
candidates_source = response_json["image_generation"]
|
||||
else:
|
||||
candidates_source = response_json
|
||||
|
||||
candidates = candidates_source.get("candidates", [])
|
||||
usage_info = response_json.get("usageMetadata")
|
||||
|
||||
if not candidates:
|
||||
logger.debug("Gemini响应中没有candidates。")
|
||||
return ResponseData(text="", raw_response=response_json)
|
||||
@@ -398,6 +408,7 @@ class GeminiAdapter(BaseAdapter):
|
||||
parts = content_data.get("parts", [])
|
||||
|
||||
text_content = ""
|
||||
images_bytes: list[bytes] = []
|
||||
parsed_tool_calls: list["LLMToolCall"] | None = None
|
||||
thought_summary_parts = []
|
||||
answer_parts = []
|
||||
@@ -409,6 +420,11 @@ class GeminiAdapter(BaseAdapter):
|
||||
thought_summary_parts.append(part["thought"])
|
||||
elif "thoughtSummary" in part:
|
||||
thought_summary_parts.append(part["thoughtSummary"])
|
||||
elif "inlineData" in part:
|
||||
inline_data = part["inlineData"]
|
||||
if "data" in inline_data:
|
||||
images_bytes.append(base64.b64decode(inline_data["data"]))
|
||||
|
||||
elif "functionCall" in part:
|
||||
if parsed_tool_calls is None:
|
||||
parsed_tool_calls = []
|
||||
@@ -475,6 +491,7 @@ class GeminiAdapter(BaseAdapter):
|
||||
return ResponseData(
|
||||
text=text_content,
|
||||
tool_calls=parsed_tool_calls,
|
||||
images=images_bytes if images_bytes else None,
|
||||
usage_info=usage_info,
|
||||
raw_response=response_json,
|
||||
grounding_metadata=grounding_metadata_obj,
|
||||
|
||||
@@ -21,7 +21,14 @@ class OpenAIAdapter(OpenAICompatAdapter):
|
||||
|
||||
@property
|
||||
def supported_api_types(self) -> list[str]:
|
||||
return ["openai", "deepseek", "zhipu", "general_openai_compat", "ark"]
|
||||
return [
|
||||
"openai",
|
||||
"deepseek",
|
||||
"zhipu",
|
||||
"general_openai_compat",
|
||||
"ark",
|
||||
"openrouter",
|
||||
]
|
||||
|
||||
def get_chat_endpoint(self, model: "LLMModel") -> str:
|
||||
"""返回聊天完成端点"""
|
||||
|
||||
+100
-2
@@ -2,7 +2,8 @@
|
||||
LLM 服务的高级 API 接口 - 便捷函数入口 (无状态)
|
||||
"""
|
||||
|
||||
from typing import Any, TypeVar
|
||||
from pathlib import Path
|
||||
from typing import Any, TypeVar, overload
|
||||
|
||||
from nonebot_plugin_alconna.uniseg import UniMessage
|
||||
from pydantic import BaseModel
|
||||
@@ -10,7 +11,7 @@ from pydantic import BaseModel
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import CommonOverrides
|
||||
from .config.generation import create_generation_config_from_kwargs
|
||||
from .config.generation import LLMGenerationConfig, create_generation_config_from_kwargs
|
||||
from .manager import get_model_instance
|
||||
from .session import AI
|
||||
from .tools.manager import tool_provider_manager
|
||||
@@ -23,6 +24,7 @@ from .types import (
|
||||
LLMResponse,
|
||||
ModelName,
|
||||
)
|
||||
from .utils import create_multimodal_message
|
||||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
|
||||
@@ -303,3 +305,99 @@ async def run_with_tools(
|
||||
raise LLMException(
|
||||
"带工具的执行循环未能产生有效的助手回复。", code=LLMErrorCode.GENERATION_FAILED
|
||||
)
|
||||
|
||||
|
||||
async def _generate_image_from_message(
|
||||
message: UniMessage,
|
||||
model: ModelName = None,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
[内部] 从 UniMessage 生成图片的核心辅助函数。
|
||||
"""
|
||||
from .utils import normalize_to_llm_messages
|
||||
|
||||
config = (
|
||||
create_generation_config_from_kwargs(**kwargs)
|
||||
if kwargs
|
||||
else LLMGenerationConfig()
|
||||
)
|
||||
|
||||
config.validation_policy = {"require_image": True}
|
||||
config.response_modalities = ["IMAGE", "TEXT"]
|
||||
|
||||
try:
|
||||
messages = await normalize_to_llm_messages(message)
|
||||
|
||||
async with await get_model_instance(model) as model_instance:
|
||||
if not model_instance.can_generate_images():
|
||||
raise LLMException(
|
||||
f"模型 '{model_instance.provider_name}/{model_instance.model_name}'"
|
||||
f"不支持图片生成",
|
||||
code=LLMErrorCode.CONFIGURATION_ERROR,
|
||||
)
|
||||
|
||||
response = await model_instance.generate_response(messages, config=config)
|
||||
|
||||
if not response.images:
|
||||
error_text = response.text or "模型未返回图片数据。"
|
||||
logger.warning(f"图片生成调用未返回图片,返回文本内容: {error_text}")
|
||||
|
||||
return response
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"执行图片生成时发生未知错误: {e}", e=e)
|
||||
raise LLMException(f"图片生成失败: {e}", cause=e)
|
||||
|
||||
|
||||
@overload
|
||||
async def create_image(
|
||||
prompt: str | UniMessage,
|
||||
*,
|
||||
images: None = None,
|
||||
model: ModelName = None,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""根据文本提示生成一张新图片。"""
|
||||
...
|
||||
|
||||
|
||||
@overload
|
||||
async def create_image(
|
||||
prompt: str | UniMessage,
|
||||
*,
|
||||
images: list[Path | bytes | str] | Path | bytes | str,
|
||||
model: ModelName = None,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""在给定图片的基础上,根据文本提示进行编辑或重新生成。"""
|
||||
...
|
||||
|
||||
|
||||
async def create_image(
|
||||
prompt: str | UniMessage,
|
||||
*,
|
||||
images: list[Path | bytes | str] | Path | bytes | str | None = None,
|
||||
model: ModelName = None,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
智能图片生成/编辑函数。
|
||||
- 如果 `images` 为 None,执行文生图。
|
||||
- 如果提供了 `images`,执行图+文生图,支持多张图片输入。
|
||||
"""
|
||||
text_prompt = (
|
||||
prompt.extract_plain_text() if isinstance(prompt, UniMessage) else str(prompt)
|
||||
)
|
||||
|
||||
image_list = []
|
||||
if images:
|
||||
if isinstance(images, list):
|
||||
image_list.extend(images)
|
||||
else:
|
||||
image_list.append(images)
|
||||
|
||||
message = create_multimodal_message(text=text_prompt, images=image_list)
|
||||
|
||||
return await _generate_image_from_message(message, model=model, **kwargs)
|
||||
|
||||
@@ -2,13 +2,15 @@
|
||||
LLM 生成配置相关类和函数
|
||||
"""
|
||||
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
from ..types import LLMResponse
|
||||
from ..types.enums import ResponseFormat
|
||||
from ..types.exceptions import LLMErrorCode, LLMException
|
||||
|
||||
@@ -64,6 +66,15 @@ class ModelConfigOverride(BaseModel):
|
||||
|
||||
custom_params: dict[str, Any] | None = Field(default=None, description="自定义参数")
|
||||
|
||||
validation_policy: dict[str, Any] | None = Field(
|
||||
default=None, description="声明式的响应验证策略 (例如: {'require_image': True})"
|
||||
)
|
||||
response_validator: Callable[[LLMResponse], None] | None = Field(
|
||||
default=None, description="一个高级回调函数,用于验证响应,验证失败时应抛出异常"
|
||||
)
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
"""转换为字典,排除None值"""
|
||||
|
||||
|
||||
@@ -50,8 +50,8 @@ class LLMHttpClient:
|
||||
async with self._lock:
|
||||
if self._client is None or self._client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClient: Initializing new httpx.AsyncClient "
|
||||
f"with config: {self.config}"
|
||||
f"LLMHttpClient: 正在初始化新的 httpx.AsyncClient "
|
||||
f"配置: {self.config}"
|
||||
)
|
||||
headers = get_user_agent()
|
||||
limits = httpx.Limits(
|
||||
@@ -92,7 +92,7 @@ class LLMHttpClient:
|
||||
)
|
||||
if self._client is None:
|
||||
raise LLMException(
|
||||
"HTTP client failed to initialize.", LLMErrorCode.CONFIGURATION_ERROR
|
||||
"HTTP 客户端初始化失败。", LLMErrorCode.CONFIGURATION_ERROR
|
||||
)
|
||||
return self._client
|
||||
|
||||
@@ -110,17 +110,17 @@ class LLMHttpClient:
|
||||
async with self._lock:
|
||||
if self._client and not self._client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClient: Closing with config: {self.config}. "
|
||||
f"Active requests: {self._active_requests}"
|
||||
f"LLMHttpClient: 正在关闭,配置: {self.config}. "
|
||||
f"活跃请求数: {self._active_requests}"
|
||||
)
|
||||
if self._active_requests > 0:
|
||||
logger.warning(
|
||||
f"LLMHttpClient: Closing while {self._active_requests} "
|
||||
f"requests are still active."
|
||||
f"LLMHttpClient: 关闭时仍有 {self._active_requests} "
|
||||
f"个请求处于活跃状态。"
|
||||
)
|
||||
await self._client.aclose()
|
||||
self._client = None
|
||||
logger.debug(f"LLMHttpClient for config {self.config} definitively closed.")
|
||||
logger.debug(f"配置为 {self.config} 的 LLMHttpClient 已完全关闭。")
|
||||
|
||||
@property
|
||||
def is_closed(self) -> bool:
|
||||
@@ -145,20 +145,17 @@ class LLMHttpClientManager:
|
||||
client = self._clients.get(key)
|
||||
if client and not client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClientManager: Reusing existing LLMHttpClient "
|
||||
f"for key: {key}"
|
||||
f"LLMHttpClientManager: 复用现有的 LLMHttpClient 密钥: {key}"
|
||||
)
|
||||
return client
|
||||
|
||||
if client and client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClientManager: Found a closed client for key {key}. "
|
||||
f"Creating a new one."
|
||||
f"LLMHttpClientManager: 发现密钥 {key} 对应的客户端已关闭。"
|
||||
f"正在创建新的客户端。"
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"LLMHttpClientManager: Creating new LLMHttpClient for key: {key}"
|
||||
)
|
||||
logger.debug(f"LLMHttpClientManager: 为密钥 {key} 创建新的 LLMHttpClient")
|
||||
http_client_config = HttpClientConfig(
|
||||
timeout=provider_config.timeout, proxy=provider_config.proxy
|
||||
)
|
||||
@@ -169,8 +166,7 @@ class LLMHttpClientManager:
|
||||
async def shutdown(self):
|
||||
async with self._lock:
|
||||
logger.info(
|
||||
f"LLMHttpClientManager: Shutting down. "
|
||||
f"Closing {len(self._clients)} client(s)."
|
||||
f"LLMHttpClientManager: 正在关闭。关闭 {len(self._clients)} 个客户端。"
|
||||
)
|
||||
close_tasks = [
|
||||
client.close()
|
||||
@@ -180,7 +176,7 @@ class LLMHttpClientManager:
|
||||
if close_tasks:
|
||||
await asyncio.gather(*close_tasks, return_exceptions=True)
|
||||
self._clients.clear()
|
||||
logger.info("LLMHttpClientManager: Shutdown complete.")
|
||||
logger.info("LLMHttpClientManager: 关闭完成。")
|
||||
|
||||
|
||||
http_client_manager = LLMHttpClientManager()
|
||||
|
||||
@@ -5,12 +5,12 @@ LLM 模型管理器
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import dump_json_safely
|
||||
|
||||
from .config import validate_override_params
|
||||
from .config.providers import AI_CONFIG_GROUP, PROVIDERS_CONFIG_KEY, get_ai_config
|
||||
@@ -43,7 +43,7 @@ def _make_cache_key(
|
||||
) -> str:
|
||||
"""生成缓存键"""
|
||||
config_str = (
|
||||
json.dumps(override_config, sort_keys=True) if override_config else "None"
|
||||
dump_json_safely(override_config, sort_keys=True) if override_config else "None"
|
||||
)
|
||||
key_data = f"{provider_model_name}:{config_str}"
|
||||
return hashlib.md5(key_data.encode()).hexdigest()
|
||||
@@ -118,6 +118,7 @@ def get_default_api_base_for_type(api_type: str) -> str | None:
|
||||
"deepseek": "https://api.deepseek.com",
|
||||
"zhipu": "https://open.bigmodel.cn",
|
||||
"gemini": "https://generativelanguage.googleapis.com",
|
||||
"openrouter": "https://openrouter.ai/api",
|
||||
"general_openai_compat": None,
|
||||
}
|
||||
|
||||
|
||||
@@ -12,6 +12,8 @@ from typing import Any, TypeVar
|
||||
from pydantic import BaseModel
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.log_sanitizer import sanitize_for_logging
|
||||
from zhenxun.utils.pydantic_compat import dump_json_safely
|
||||
|
||||
from .adapters.base import RequestData
|
||||
from .config import LLMGenerationConfig
|
||||
@@ -34,7 +36,6 @@ from .types import (
|
||||
ToolExecutable,
|
||||
)
|
||||
from .types.capabilities import ModelCapabilities, ModelModality
|
||||
from .utils import _sanitize_request_body_for_logging
|
||||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
|
||||
@@ -187,21 +188,32 @@ class LLMModel(LLMModelBase):
|
||||
logger.debug(f"🔑 API密钥: {masked_key}")
|
||||
logger.debug(f"📋 请求头: {dict(request_data.headers)}")
|
||||
|
||||
sanitized_body = _sanitize_request_body_for_logging(request_data.body)
|
||||
request_body_str = json.dumps(sanitized_body, ensure_ascii=False, indent=2)
|
||||
sanitizer_req_context_map = {"gemini": "gemini_request"}
|
||||
sanitizer_req_context = sanitizer_req_context_map.get(
|
||||
self.api_type, "openai_request"
|
||||
)
|
||||
sanitized_body = sanitize_for_logging(
|
||||
request_data.body, context=sanitizer_req_context
|
||||
)
|
||||
request_body_str = dump_json_safely(
|
||||
sanitized_body, ensure_ascii=False, indent=2
|
||||
)
|
||||
logger.debug(f"📦 请求体: {request_body_str}")
|
||||
|
||||
http_response = await http_client.post(
|
||||
request_data.url,
|
||||
headers=request_data.headers,
|
||||
json=request_data.body,
|
||||
content=dump_json_safely(request_data.body, ensure_ascii=False),
|
||||
)
|
||||
|
||||
logger.debug(f"📥 响应状态码: {http_response.status_code}")
|
||||
logger.debug(f"📄 响应头: {dict(http_response.headers)}")
|
||||
|
||||
response_bytes = await http_response.aread()
|
||||
logger.debug(f"📦 响应体已完整读取 ({len(response_bytes)} bytes)")
|
||||
|
||||
if http_response.status_code != 200:
|
||||
error_text = http_response.text
|
||||
error_text = response_bytes.decode("utf-8", errors="ignore")
|
||||
logger.error(
|
||||
f"❌ HTTP请求失败: {http_response.status_code} - {error_text} "
|
||||
f"[{log_context}]"
|
||||
@@ -232,13 +244,22 @@ class LLMModel(LLMModelBase):
|
||||
)
|
||||
|
||||
try:
|
||||
response_json = http_response.json()
|
||||
response_json = json.loads(response_bytes)
|
||||
|
||||
sanitizer_context_map = {"gemini": "gemini_response"}
|
||||
sanitizer_context = sanitizer_context_map.get(
|
||||
self.api_type, "openai_response"
|
||||
)
|
||||
|
||||
sanitized_for_log = sanitize_for_logging(
|
||||
response_json, context=sanitizer_context
|
||||
)
|
||||
|
||||
response_json_str = json.dumps(
|
||||
response_json, ensure_ascii=False, indent=2
|
||||
sanitized_for_log, ensure_ascii=False, indent=2
|
||||
)
|
||||
logger.debug(f"📋 响应JSON: {response_json_str}")
|
||||
parsed_data = parse_response_func(response_json)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"解析 {log_context} 响应失败: {e}", e=e)
|
||||
await self.key_store.record_failure(api_key, None, str(e))
|
||||
@@ -290,7 +311,7 @@ class LLMModel(LLMModelBase):
|
||||
adapter.validate_embedding_response(response_json)
|
||||
return adapter.parse_embedding_response(response_json)
|
||||
|
||||
parsed_data, api_key_used = await self._perform_api_call(
|
||||
parsed_data, _api_key_used = await self._perform_api_call(
|
||||
prepare_request_func=prepare_request,
|
||||
parse_response_func=parse_response,
|
||||
http_client=http_client,
|
||||
@@ -376,6 +397,7 @@ class LLMModel(LLMModelBase):
|
||||
return LLMResponse(
|
||||
text=response_data.text,
|
||||
usage_info=response_data.usage_info,
|
||||
images=response_data.images,
|
||||
raw_response=response_data.raw_response,
|
||||
tool_calls=response_tool_calls if response_tool_calls else None,
|
||||
code_executions=response_data.code_executions,
|
||||
@@ -390,6 +412,56 @@ class LLMModel(LLMModelBase):
|
||||
failed_keys=failed_keys,
|
||||
log_context="Generation",
|
||||
)
|
||||
|
||||
if config:
|
||||
if config.response_validator:
|
||||
try:
|
||||
config.response_validator(parsed_data)
|
||||
except Exception as e:
|
||||
raise LLMException(
|
||||
f"响应内容未通过自定义验证器: {e}",
|
||||
code=LLMErrorCode.API_RESPONSE_INVALID,
|
||||
details={"validator_error": str(e)},
|
||||
cause=e,
|
||||
) from e
|
||||
|
||||
policy = config.validation_policy
|
||||
if policy:
|
||||
if policy.get("require_image") and not parsed_data.images:
|
||||
if self.api_type == "gemini" and parsed_data.raw_response:
|
||||
usage_metadata = parsed_data.raw_response.get(
|
||||
"usageMetadata", {}
|
||||
)
|
||||
prompt_token_details = usage_metadata.get(
|
||||
"promptTokensDetails", []
|
||||
)
|
||||
prompt_had_image = any(
|
||||
detail.get("modality") == "IMAGE"
|
||||
for detail in prompt_token_details
|
||||
)
|
||||
|
||||
if prompt_had_image:
|
||||
raise LLMException(
|
||||
"响应验证失败:模型接收了图片输入但未生成图片。",
|
||||
code=LLMErrorCode.API_RESPONSE_INVALID,
|
||||
details={
|
||||
"policy": policy,
|
||||
"text_response": parsed_data.text,
|
||||
"raw_response": parsed_data.raw_response,
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.debug("Gemini提示词中未包含图片,跳过图片要求重试。")
|
||||
else:
|
||||
raise LLMException(
|
||||
"响应验证失败:要求返回图片但未找到图片数据。",
|
||||
code=LLMErrorCode.API_RESPONSE_INVALID,
|
||||
details={
|
||||
"policy": policy,
|
||||
"text_response": parsed_data.text,
|
||||
},
|
||||
)
|
||||
|
||||
return parsed_data, api_key_used
|
||||
|
||||
async def close(self):
|
||||
|
||||
@@ -44,6 +44,13 @@ GEMINI_CAPABILITIES = ModelCapabilities(
|
||||
supports_tool_calling=True,
|
||||
)
|
||||
|
||||
GEMINI_IMAGE_GEN_CAPABILITIES = ModelCapabilities(
|
||||
input_modalities={ModelModality.TEXT, ModelModality.IMAGE},
|
||||
output_modalities={ModelModality.TEXT, ModelModality.IMAGE},
|
||||
supports_tool_calling=True,
|
||||
)
|
||||
|
||||
|
||||
DOUBAO_ADVANCED_MULTIMODAL_CAPABILITIES = ModelCapabilities(
|
||||
input_modalities={ModelModality.TEXT, ModelModality.IMAGE, ModelModality.VIDEO},
|
||||
output_modalities={ModelModality.TEXT},
|
||||
@@ -83,6 +90,7 @@ MODEL_CAPABILITIES_REGISTRY: dict[str, ModelCapabilities] = {
|
||||
output_modalities={ModelModality.EMBEDDING},
|
||||
is_embedding_model=True,
|
||||
),
|
||||
"*gemini-*-image-preview*": GEMINI_IMAGE_GEN_CAPABILITIES,
|
||||
"gemini-2.5-pro*": GEMINI_CAPABILITIES,
|
||||
"gemini-1.5-pro*": GEMINI_CAPABILITIES,
|
||||
"gemini-2.5-flash*": GEMINI_CAPABILITIES,
|
||||
|
||||
@@ -425,6 +425,7 @@ class LLMResponse(BaseModel):
|
||||
"""LLM 响应"""
|
||||
|
||||
text: str
|
||||
images: list[bytes] | None = None
|
||||
usage_info: dict[str, Any] | None = None
|
||||
raw_response: dict[str, Any] | None = None
|
||||
tool_calls: list[Any] | None = None
|
||||
|
||||
@@ -273,54 +273,6 @@ def message_to_unimessage(message: PlatformMessage) -> UniMessage:
|
||||
return UniMessage(uni_segments)
|
||||
|
||||
|
||||
def _sanitize_request_body_for_logging(body: dict) -> dict:
|
||||
"""
|
||||
净化请求体用于日志记录,移除大数据字段并添加摘要信息
|
||||
|
||||
参数:
|
||||
body: 原始请求体字典。
|
||||
|
||||
返回:
|
||||
dict: 净化后的请求体字典。
|
||||
"""
|
||||
try:
|
||||
sanitized_body = copy.deepcopy(body)
|
||||
|
||||
if "contents" in sanitized_body and isinstance(
|
||||
sanitized_body["contents"], list
|
||||
):
|
||||
for content_item in sanitized_body["contents"]:
|
||||
if "parts" in content_item and isinstance(content_item["parts"], list):
|
||||
media_summary = []
|
||||
new_parts = []
|
||||
for part in content_item["parts"]:
|
||||
if "inlineData" in part and isinstance(
|
||||
part["inlineData"], dict
|
||||
):
|
||||
data = part["inlineData"].get("data")
|
||||
if isinstance(data, str):
|
||||
mime_type = part["inlineData"].get(
|
||||
"mimeType", "unknown"
|
||||
)
|
||||
media_summary.append(f"{mime_type} ({len(data)} chars)")
|
||||
continue
|
||||
new_parts.append(part)
|
||||
|
||||
if media_summary:
|
||||
summary_text = (
|
||||
f"[多模态内容: {len(media_summary)}个文件 - "
|
||||
f"{', '.join(media_summary)}]"
|
||||
)
|
||||
new_parts.insert(0, {"text": summary_text})
|
||||
|
||||
content_item["parts"] = new_parts
|
||||
|
||||
return sanitized_body
|
||||
except Exception as e:
|
||||
logger.warning(f"日志净化失败: {e},将记录原始请求体。")
|
||||
return body
|
||||
|
||||
|
||||
def sanitize_schema_for_llm(schema: Any, api_type: str) -> Any:
|
||||
"""
|
||||
递归地净化 JSON Schema,移除特定 LLM API 不支持的关键字。
|
||||
|
||||
@@ -22,6 +22,7 @@ from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.path_config import THEMES_PATH, UI_CACHE_PATH
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.exception import RenderingError
|
||||
from zhenxun.utils.log_sanitizer import sanitize_for_logging
|
||||
from zhenxun.utils.pydantic_compat import _dump_pydantic_obj
|
||||
|
||||
from .config import RESERVED_TEMPLATE_KEYS
|
||||
@@ -216,16 +217,17 @@ class RendererService:
|
||||
context.processed_components.add(component_id)
|
||||
|
||||
component_path_base = str(component.template_name)
|
||||
variant = getattr(component, "variant", None)
|
||||
manifest = await context.theme_manager.get_template_manifest(
|
||||
component_path_base
|
||||
component_path_base, skin=variant
|
||||
)
|
||||
|
||||
style_paths_to_load = []
|
||||
if manifest and manifest.styles:
|
||||
if manifest and "styles" in manifest:
|
||||
styles = (
|
||||
[manifest.styles]
|
||||
if isinstance(manifest.styles, str)
|
||||
else manifest.styles
|
||||
[manifest["styles"]]
|
||||
if isinstance(manifest["styles"], str)
|
||||
else manifest["styles"]
|
||||
)
|
||||
for style_path in styles:
|
||||
full_style_path = str(Path(component_path_base) / style_path).replace(
|
||||
@@ -382,6 +384,7 @@ class RendererService:
|
||||
)
|
||||
|
||||
temp_env.globals.update(context.theme_manager.jinja_env.globals)
|
||||
temp_env.filters.update(context.theme_manager.jinja_env.filters)
|
||||
temp_env.globals["asset"] = (
|
||||
context.theme_manager._create_standalone_asset_loader(template_dir)
|
||||
)
|
||||
@@ -430,10 +433,11 @@ class RendererService:
|
||||
component_render_options = {}
|
||||
|
||||
manifest_options = {}
|
||||
variant = getattr(component, "variant", None)
|
||||
if manifest := await context.theme_manager.get_template_manifest(
|
||||
component.template_name
|
||||
component.template_name, skin=variant
|
||||
):
|
||||
manifest_options = manifest.render_options or {}
|
||||
manifest_options = manifest.get("render_options", {})
|
||||
|
||||
final_render_options = component_render_options.copy()
|
||||
final_render_options.update(manifest_options)
|
||||
@@ -470,10 +474,7 @@ class RendererService:
|
||||
) from e
|
||||
|
||||
async def render(
|
||||
self,
|
||||
component: Renderable,
|
||||
use_cache: bool = False,
|
||||
**render_options,
|
||||
self, component: Renderable, use_cache: bool = False, **render_options
|
||||
) -> bytes:
|
||||
"""
|
||||
统一的、多态的渲染入口,直接返回图片字节。
|
||||
@@ -504,9 +505,12 @@ class RendererService:
|
||||
)
|
||||
result = await self._render_component(context)
|
||||
if Config.get_config("UI", "DEBUG_MODE") and result.html_content:
|
||||
sanitized_html = sanitize_for_logging(
|
||||
result.html_content, context="ui_html"
|
||||
)
|
||||
logger.info(
|
||||
f"--- [UI DEBUG] HTML for {component.__class__.__name__} ---\n"
|
||||
f"{result.html_content}\n"
|
||||
f"{sanitized_html}\n"
|
||||
f"--- [UI DEBUG] End of HTML ---"
|
||||
)
|
||||
if result.image_bytes is None:
|
||||
@@ -556,6 +560,8 @@ class RendererService:
|
||||
await self.initialize()
|
||||
assert self._theme_manager is not None, "ThemeManager 未初始化"
|
||||
|
||||
self._theme_manager._manifest_cache.clear()
|
||||
logger.debug("已清除UI清单缓存 (manifest cache)。")
|
||||
current_theme_name = Config.get_config("UI", "THEME", "default")
|
||||
await self._theme_manager.load_theme(current_theme_name)
|
||||
logger.info(f"主题 '{current_theme_name}' 已成功重载。")
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import aiofiles
|
||||
from jinja2 import (
|
||||
ChoiceLoader,
|
||||
Environment,
|
||||
@@ -21,7 +21,6 @@ import ujson as json
|
||||
|
||||
from zhenxun.configs.path_config import THEMES_PATH
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.renderer.models import TemplateManifest
|
||||
from zhenxun.services.renderer.protocols import Renderable
|
||||
from zhenxun.services.renderer.registry import asset_registry
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
@@ -32,6 +31,20 @@ if TYPE_CHECKING:
|
||||
from .config import RESERVED_TEMPLATE_KEYS
|
||||
|
||||
|
||||
def deep_merge_dict(base: dict, new: dict) -> dict:
|
||||
"""
|
||||
递归地将 new 字典合并到 base 字典中。
|
||||
new 字典中的值会覆盖 base 字典中的值。
|
||||
"""
|
||||
result = base.copy()
|
||||
for key, value in new.items():
|
||||
if isinstance(value, dict) and key in result and isinstance(result[key], dict):
|
||||
result[key] = deep_merge_dict(result[key], value)
|
||||
else:
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
|
||||
class RelativePathEnvironment(Environment):
|
||||
"""
|
||||
一个自定义的 Jinja2 环境,重写了 join_path 方法以支持模板间的相对路径引用。
|
||||
@@ -151,14 +164,42 @@ class ResourceResolver:
|
||||
|
||||
def resolve_asset_uri(self, asset_path: str, current_template_name: str) -> str:
|
||||
"""解析资源路径,实现完整的回退逻辑,并返回可用的URI。"""
|
||||
if not self.theme_manager.current_theme:
|
||||
if (
|
||||
not self.theme_manager.current_theme
|
||||
or not self.theme_manager.jinja_env.loader
|
||||
):
|
||||
return ""
|
||||
|
||||
if asset_path.startswith("@"):
|
||||
try:
|
||||
full_asset_path = self.theme_manager.jinja_env.join_path(
|
||||
asset_path, current_template_name
|
||||
)
|
||||
_source, file_abs_path, _uptodate = (
|
||||
self.theme_manager.jinja_env.loader.get_source(
|
||||
self.theme_manager.jinja_env, full_asset_path
|
||||
)
|
||||
)
|
||||
if file_abs_path:
|
||||
logger.debug(
|
||||
f"Jinja Loader resolved asset '{asset_path}'->'{file_abs_path}'"
|
||||
)
|
||||
return Path(file_abs_path).absolute().as_uri()
|
||||
except TemplateNotFound:
|
||||
logger.warning(
|
||||
f"资源文件在命名空间中未找到: '{asset_path}'"
|
||||
f"(在模板 '{current_template_name}' 中引用)"
|
||||
)
|
||||
return ""
|
||||
|
||||
search_paths: list[tuple[str, Path]] = []
|
||||
if asset_path.startswith("./"):
|
||||
if asset_path.startswith("./") or asset_path.startswith("../"):
|
||||
relative_part = (
|
||||
asset_path[2:] if asset_path.startswith("./") else asset_path
|
||||
)
|
||||
search_paths.extend(
|
||||
self._search_paths_for_relative_asset(
|
||||
asset_path[2:], current_template_name
|
||||
relative_part, current_template_name
|
||||
)
|
||||
)
|
||||
else:
|
||||
@@ -209,6 +250,9 @@ class ThemeManager:
|
||||
|
||||
self.jinja_env.filters["md"] = self._markdown_filter
|
||||
|
||||
self._manifest_cache: dict[str, Any] = {}
|
||||
self._manifest_cache_lock = asyncio.Lock()
|
||||
|
||||
def list_available_themes(self) -> list[str]:
|
||||
"""扫描主题目录并返回所有可用的主题名称。"""
|
||||
if not THEMES_PATH.is_dir():
|
||||
@@ -377,16 +421,26 @@ class ThemeManager:
|
||||
logger.error(f"指定的模板文件路径不存在: '{component_path_base}'", e=e)
|
||||
raise e
|
||||
|
||||
entrypoint_filename = "main.html"
|
||||
manifest = await self.get_template_manifest(component_path_base)
|
||||
if manifest and manifest.entrypoint:
|
||||
entrypoint_filename = manifest.entrypoint
|
||||
base_manifest = await self.get_template_manifest(component_path_base)
|
||||
|
||||
skin_to_use = variant or (base_manifest.get("skin") if base_manifest else None)
|
||||
|
||||
final_manifest = await self.get_template_manifest(
|
||||
component_path_base, skin=skin_to_use
|
||||
)
|
||||
logger.debug(f"final_manifest: {final_manifest}")
|
||||
|
||||
entrypoint_filename = (
|
||||
final_manifest.get("entrypoint", "main.html")
|
||||
if final_manifest
|
||||
else "main.html"
|
||||
)
|
||||
|
||||
potential_paths = []
|
||||
|
||||
if variant:
|
||||
if skin_to_use:
|
||||
potential_paths.append(
|
||||
f"{component_path_base}/skins/{variant}/{entrypoint_filename}"
|
||||
f"{component_path_base}/skins/{skin_to_use}/{entrypoint_filename}"
|
||||
)
|
||||
|
||||
potential_paths.append(f"{component_path_base}/{entrypoint_filename}")
|
||||
@@ -410,28 +464,88 @@ class ThemeManager:
|
||||
logger.error(err_msg)
|
||||
raise TemplateNotFound(err_msg)
|
||||
|
||||
async def get_template_manifest(
|
||||
self, component_path: str
|
||||
) -> TemplateManifest | None:
|
||||
"""
|
||||
查找并解析组件的 manifest.json 文件。
|
||||
"""
|
||||
manifest_path_str = f"{component_path}/manifest.json"
|
||||
async def _load_single_manifest(self, path_str: str) -> dict[str, Any] | None:
|
||||
"""从指定路径加载单个 manifest.json 文件。"""
|
||||
normalized_path = path_str.replace("\\", "/")
|
||||
manifest_path_str = f"{normalized_path}/manifest.json"
|
||||
|
||||
if not self.jinja_env.loader:
|
||||
return None
|
||||
|
||||
try:
|
||||
_, full_path, _ = self.jinja_env.loader.get_source(
|
||||
source, filepath, _ = self.jinja_env.loader.get_source(
|
||||
self.jinja_env, manifest_path_str
|
||||
)
|
||||
if full_path and Path(full_path).exists():
|
||||
async with aiofiles.open(full_path, encoding="utf-8") as f:
|
||||
manifest_data = json.loads(await f.read())
|
||||
return TemplateManifest(**manifest_data)
|
||||
logger.debug(f"找到清单文件: '{manifest_path_str}' (从 '{filepath}' 加载)")
|
||||
return json.loads(source)
|
||||
except TemplateNotFound:
|
||||
logger.trace(f"未找到清单文件: '{manifest_path_str}'")
|
||||
return None
|
||||
return None
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(f"清单文件 '{manifest_path_str}' 解析失败")
|
||||
return None
|
||||
|
||||
async def _load_and_merge_manifests(
|
||||
self, component_path: Path | str, skin: str | None = None
|
||||
) -> dict[str, Any] | None:
|
||||
"""加载基础和皮肤清单并进行合并。"""
|
||||
logger.debug(f"开始加载清单: component_path='{component_path}', skin='{skin}'")
|
||||
|
||||
base_manifest = await self._load_single_manifest(str(component_path))
|
||||
|
||||
if skin:
|
||||
skin_path = Path(component_path) / "skins" / skin
|
||||
skin_manifest = await self._load_single_manifest(str(skin_path))
|
||||
|
||||
if skin_manifest:
|
||||
if base_manifest:
|
||||
merged = deep_merge_dict(base_manifest, skin_manifest)
|
||||
logger.debug(
|
||||
f"已合并基础清单和皮肤清单: '{component_path}' + skin '{skin}'"
|
||||
)
|
||||
return merged
|
||||
else:
|
||||
logger.debug(f"只找到皮肤清单: '{skin_path}'")
|
||||
return skin_manifest
|
||||
|
||||
if base_manifest:
|
||||
logger.debug(f"只找到基础清单: '{component_path}'")
|
||||
else:
|
||||
logger.debug(f"未找到任何清单: '{component_path}'")
|
||||
|
||||
return base_manifest
|
||||
|
||||
async def get_template_manifest(
|
||||
self, component_path: str, skin: str | None = None
|
||||
) -> dict[str, Any] | None:
|
||||
"""
|
||||
查找并解析组件的 manifest.json 文件。
|
||||
支持皮肤清单的继承与合并,并带有缓存。
|
||||
|
||||
Args:
|
||||
component_path: 组件路径
|
||||
skin: 皮肤名称(可选)
|
||||
|
||||
Returns:
|
||||
合并后的清单字典,如果不存在则返回 None
|
||||
"""
|
||||
cache_key = f"{component_path}:{skin or 'base'}"
|
||||
|
||||
if cache_key in self._manifest_cache:
|
||||
logger.debug(f"清单缓存命中: '{cache_key}'")
|
||||
return self._manifest_cache[cache_key]
|
||||
|
||||
async with self._manifest_cache_lock:
|
||||
if cache_key in self._manifest_cache:
|
||||
logger.debug(f"清单缓存命中(锁内): '{cache_key}'")
|
||||
return self._manifest_cache[cache_key]
|
||||
|
||||
manifest = await self._load_and_merge_manifests(component_path, skin)
|
||||
|
||||
self._manifest_cache[cache_key] = manifest
|
||||
logger.debug(f"清单已缓存: '{cache_key}'")
|
||||
|
||||
return manifest
|
||||
|
||||
async def resolve_markdown_style_path(
|
||||
self, style_name: str, context: "RenderContext"
|
||||
|
||||
@@ -126,12 +126,15 @@ class SqlUtils:
|
||||
def format_usage_for_markdown(text: str) -> str:
|
||||
"""
|
||||
智能地将Python多行字符串转换为适合Markdown渲染的格式。
|
||||
- 将单个换行符替换为Markdown的硬换行(行尾加两个空格)。
|
||||
- 在列表、标题等块级元素前自动插入换行,确保正确解析。
|
||||
- 将段落内的单个换行符替换为Markdown的硬换行(行尾加两个空格)。
|
||||
- 保留两个或更多的连续换行符,使其成为Markdown的段落分隔。
|
||||
"""
|
||||
if not text:
|
||||
return ""
|
||||
text = re.sub(r"\n{2,}", "<<PARAGRAPH_BREAK>>", text)
|
||||
text = text.replace("\n", " \n")
|
||||
text = text.replace("<<PARAGRAPH_BREAK>>", "\n\n")
|
||||
|
||||
text = re.sub(r"([^\n])\n(\s*[-*] |\s*#+\s|\s*>)", r"\1\n\n\2", text)
|
||||
|
||||
text = re.sub(r"(?<!\n)\n(?!\n)", " \n", text)
|
||||
|
||||
return text
|
||||
|
||||
@@ -263,6 +263,18 @@ class AsyncHttpx:
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
if isinstance(e, HTTPStatusError):
|
||||
status = getattr(e.response, "status_code", "?")
|
||||
try:
|
||||
body_text = getattr(e.response, "text", None)
|
||||
if body_text is not None and len(body_text) > 2000:
|
||||
body_text = body_text[:2000] + "...(truncated)"
|
||||
except Exception:
|
||||
body_text = "<unavailable>"
|
||||
logger.debug(
|
||||
f"请求失败: {url} {status} {body_text}",
|
||||
"AsyncHttpx:FallbackExecutor",
|
||||
)
|
||||
exceptions.append(e)
|
||||
if url != url_list[-1]:
|
||||
logger.warning(
|
||||
|
||||
@@ -0,0 +1,202 @@
|
||||
import copy
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from nonebot.adapters import Message, MessageSegment
|
||||
|
||||
|
||||
def _truncate_base64_string(value: str, threshold: int = 256) -> str:
|
||||
"""如果字符串是超长的base64或data URI,则截断它。"""
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
|
||||
prefixes = ("base64://", "data:image", "data:video", "data:audio")
|
||||
if value.startswith(prefixes) and len(value) > threshold:
|
||||
prefix = next((p for p in prefixes if value.startswith(p)), "base64")
|
||||
return f"[{prefix}_data_omitted_len={len(value)}]"
|
||||
return value
|
||||
|
||||
|
||||
def _sanitize_ui_html(html_string: str) -> str:
|
||||
"""
|
||||
专门用于净化UI渲染调试HTML的函数。
|
||||
它会查找所有内联的base64数据(如字体、图片)并将其截断。
|
||||
"""
|
||||
if not isinstance(html_string, str):
|
||||
return html_string
|
||||
|
||||
pattern = re.compile(r"(data:[^;]+;base64,)[A-Za-z0-9+/=\s]{100,}")
|
||||
|
||||
def replacer(match):
|
||||
prefix = match.group(1)
|
||||
original_len = len(match.group(0)) - len(prefix)
|
||||
return f"{prefix}[...base64_omitted_len={original_len}...]"
|
||||
|
||||
return pattern.sub(replacer, html_string)
|
||||
|
||||
|
||||
def _sanitize_nonebot_message(message: Message) -> Message:
|
||||
"""净化nonebot.adapter.Message对象,用于日志记录。"""
|
||||
sanitized_message = copy.deepcopy(message)
|
||||
for seg in sanitized_message:
|
||||
seg: MessageSegment
|
||||
if seg.type in ("image", "record", "video"):
|
||||
file_info = seg.data.get("file", "")
|
||||
if isinstance(file_info, str):
|
||||
seg.data["file"] = _truncate_base64_string(file_info)
|
||||
return sanitized_message
|
||||
|
||||
|
||||
def _sanitize_openai_response(response_json: dict) -> dict:
|
||||
"""净化OpenAI兼容API的响应体。"""
|
||||
try:
|
||||
sanitized_json = copy.deepcopy(response_json)
|
||||
if "choices" in sanitized_json and isinstance(sanitized_json["choices"], list):
|
||||
for choice in sanitized_json["choices"]:
|
||||
if "message" in choice and isinstance(choice["message"], dict):
|
||||
message = choice["message"]
|
||||
if "images" in message and isinstance(message["images"], list):
|
||||
for i, image_info in enumerate(message["images"]):
|
||||
if "image_url" in image_info and isinstance(
|
||||
image_info["image_url"], dict
|
||||
):
|
||||
url = image_info["image_url"].get("url", "")
|
||||
message["images"][i]["image_url"]["url"] = (
|
||||
_truncate_base64_string(url)
|
||||
)
|
||||
return sanitized_json
|
||||
except Exception:
|
||||
return response_json
|
||||
|
||||
|
||||
def _sanitize_openai_request(body: dict) -> dict:
|
||||
"""净化OpenAI兼容API的请求体,主要截断图片base64。"""
|
||||
try:
|
||||
sanitized_json = copy.deepcopy(body)
|
||||
if "messages" in sanitized_json and isinstance(
|
||||
sanitized_json["messages"], list
|
||||
):
|
||||
for message in sanitized_json["messages"]:
|
||||
if "content" in message and isinstance(message["content"], list):
|
||||
for i, part in enumerate(message["content"]):
|
||||
if part.get("type") == "image_url":
|
||||
if "image_url" in part and isinstance(
|
||||
part["image_url"], dict
|
||||
):
|
||||
url = part["image_url"].get("url", "")
|
||||
message["content"][i]["image_url"]["url"] = (
|
||||
_truncate_base64_string(url)
|
||||
)
|
||||
return sanitized_json
|
||||
except Exception:
|
||||
return body
|
||||
|
||||
|
||||
def _sanitize_gemini_response(response_json: dict) -> dict:
|
||||
"""净化Gemini API的响应体,处理文本和图片生成两种格式。"""
|
||||
try:
|
||||
sanitized_json = copy.deepcopy(response_json)
|
||||
|
||||
def _process_candidates(candidates_list: list):
|
||||
"""辅助函数,用于处理任何 candidates 列表。"""
|
||||
if not isinstance(candidates_list, list):
|
||||
return
|
||||
for candidate in candidates_list:
|
||||
if "content" in candidate and isinstance(candidate["content"], dict):
|
||||
content = candidate["content"]
|
||||
if "parts" in content and isinstance(content["parts"], list):
|
||||
for i, part in enumerate(content["parts"]):
|
||||
if "inlineData" in part and isinstance(
|
||||
part["inlineData"], dict
|
||||
):
|
||||
data = part["inlineData"].get("data", "")
|
||||
if isinstance(data, str) and len(data) > 256:
|
||||
content["parts"][i]["inlineData"]["data"] = (
|
||||
f"[base64_data_omitted_len={len(data)}]"
|
||||
)
|
||||
|
||||
if "candidates" in sanitized_json:
|
||||
_process_candidates(sanitized_json["candidates"])
|
||||
|
||||
if "image_generation" in sanitized_json and isinstance(
|
||||
sanitized_json["image_generation"], dict
|
||||
):
|
||||
if "candidates" in sanitized_json["image_generation"]:
|
||||
_process_candidates(sanitized_json["image_generation"]["candidates"])
|
||||
|
||||
return sanitized_json
|
||||
except Exception:
|
||||
return response_json
|
||||
|
||||
|
||||
def _sanitize_gemini_request(body: dict) -> dict:
|
||||
"""净化Gemini API的请求体,进行结构转换和总结。"""
|
||||
try:
|
||||
sanitized_body = copy.deepcopy(body)
|
||||
if "contents" in sanitized_body and isinstance(
|
||||
sanitized_body["contents"], list
|
||||
):
|
||||
for content_item in sanitized_body["contents"]:
|
||||
if "parts" in content_item and isinstance(content_item["parts"], list):
|
||||
media_summary = []
|
||||
new_parts = []
|
||||
for part in content_item["parts"]:
|
||||
if "inlineData" in part and isinstance(
|
||||
part["inlineData"], dict
|
||||
):
|
||||
data = part["inlineData"].get("data")
|
||||
if isinstance(data, str):
|
||||
mime_type = part["inlineData"].get(
|
||||
"mimeType", "unknown"
|
||||
)
|
||||
media_summary.append(f"{mime_type} ({len(data)} chars)")
|
||||
continue
|
||||
new_parts.append(part)
|
||||
|
||||
if media_summary:
|
||||
summary_text = (
|
||||
f"[多模态内容: {len(media_summary)}个文件 - "
|
||||
f"{', '.join(media_summary)}]"
|
||||
)
|
||||
new_parts.insert(0, {"text": summary_text})
|
||||
|
||||
content_item["parts"] = new_parts
|
||||
return sanitized_body
|
||||
except Exception:
|
||||
return body
|
||||
|
||||
|
||||
def sanitize_for_logging(data: Any, context: str | None = None) -> Any:
|
||||
"""
|
||||
统一的日志净化入口。
|
||||
|
||||
Args:
|
||||
data: 需要净化的数据 (dict, Message, etc.).
|
||||
context: 净化场景的上下文标识,例如 'gemini_request', 'openai_response'.
|
||||
|
||||
Returns:
|
||||
净化后的数据。
|
||||
"""
|
||||
if context == "nonebot_message":
|
||||
if isinstance(data, Message):
|
||||
return _sanitize_nonebot_message(data)
|
||||
elif context == "openai_response":
|
||||
if isinstance(data, dict):
|
||||
return _sanitize_openai_response(data)
|
||||
elif context == "gemini_response":
|
||||
if isinstance(data, dict):
|
||||
return _sanitize_gemini_response(data)
|
||||
elif context == "gemini_request":
|
||||
if isinstance(data, dict):
|
||||
return _sanitize_gemini_request(data)
|
||||
elif context == "openai_request":
|
||||
if isinstance(data, dict):
|
||||
return _sanitize_openai_request(data)
|
||||
elif context == "ui_html":
|
||||
if isinstance(data, str):
|
||||
return _sanitize_ui_html(data)
|
||||
else:
|
||||
if isinstance(data, str):
|
||||
return _truncate_base64_string(data)
|
||||
|
||||
return data
|
||||
@@ -5,10 +5,14 @@ Pydantic V1 & V2 兼容层模块
|
||||
包括 model_dump, model_copy, model_json_schema, parse_as 等。
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from typing import Any, TypeVar, get_args, get_origin
|
||||
|
||||
from nonebot.compat import PYDANTIC_V2, model_dump
|
||||
from pydantic import VERSION, BaseModel
|
||||
import ujson as json
|
||||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
V = TypeVar("V")
|
||||
@@ -19,6 +23,7 @@ __all__ = [
|
||||
"_dump_pydantic_obj",
|
||||
"_is_pydantic_type",
|
||||
"compat_computed_field",
|
||||
"dump_json_safely",
|
||||
"model_copy",
|
||||
"model_dump",
|
||||
"model_json_schema",
|
||||
@@ -93,3 +98,26 @@ def parse_as(type_: type[V], obj: Any) -> V:
|
||||
from pydantic import TypeAdapter # type: ignore
|
||||
|
||||
return TypeAdapter(type_).validate_python(obj)
|
||||
|
||||
|
||||
def dump_json_safely(obj: Any, **kwargs) -> str:
|
||||
"""
|
||||
安全地将可能包含 Pydantic 特定类型 (如 Enum) 的对象序列化为 JSON 字符串。
|
||||
"""
|
||||
|
||||
def default_serializer(o):
|
||||
if isinstance(o, Enum):
|
||||
return o.value
|
||||
if isinstance(o, datetime):
|
||||
return o.isoformat()
|
||||
if isinstance(o, Path):
|
||||
return str(o.as_posix())
|
||||
if isinstance(o, set):
|
||||
return list(o)
|
||||
if isinstance(o, BaseModel):
|
||||
return model_dump(o)
|
||||
raise TypeError(
|
||||
f"Object of type {o.__class__.__name__} is not JSON serializable"
|
||||
)
|
||||
|
||||
return json.dumps(obj, default=default_serializer, **kwargs)
|
||||
|
||||
Reference in New Issue
Block a user