Compare commits

...
Author SHA1 Message Date
molanp eb6d90ae88 docs(data-source): 更新插件安装函数的参数文档说明 (#2069)
检查bot是否运行正常 / bot check (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, javascript-typescript) (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
Release Drafter / Update Release Draft (push) Has been cancelled
Force Sync to Aliyun / sync (push) Has been cancelled
Update Version / update-version (push) Has been cancelled
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
修改 StoreManager 类中安装插件函数的文档字符串,更新参数列表说
明。将原有的 github_url、module_path、is_dir 参数说明替换为
plugin_info 和 source 参数说明,保持文档与实际函数签名一致。
2025-10-22 20:57:07 +08:00
molanp 4b8013d2d6 Feat: Add spaces (#2064)
检查bot是否运行正常 / bot check (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, javascript-typescript) (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
Release Drafter / Update Release Draft (push) Has been cancelled
Force Sync to Aliyun / sync (push) Has been cancelled
Update Version / update-version (push) Has been cancelled
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
2025-10-17 09:22:18 +08:00
HibiKier d528711641 🐛 fix(http_utils): 增强错误处理,记录请求失败的详细信息 (#2065)
检查bot是否运行正常 / bot check (push) Waiting to run
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, javascript-typescript) (push) Waiting to run
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Waiting to run
Sequential Lint and Type Check / ruff-call (push) Waiting to run
Sequential Lint and Type Check / pyright-call (push) Blocked by required conditions
Release Drafter / Update Release Draft (push) Waiting to run
Force Sync to Aliyun / sync (push) Waiting to run
Update Version / update-version (push) Waiting to run
* 🐛 fix(http_utils): 增强错误处理,记录请求失败的详细信息

* 🐛 fix(http_utils): 改进HTTP错误处理,记录请求失败的状态码和响应内容
2025-10-16 17:31:08 +08:00
molanp 1cc18bb195 fix(shop): 修改道具不存在时的提示信息 (#2061)
检查bot是否运行正常 / bot check (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, javascript-typescript) (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
Release Drafter / Update Release Draft (push) Has been cancelled
Force Sync to Aliyun / sync (push) Has been cancelled
Update Version / update-version (push) Has been cancelled
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
- 将道具不存在时的提示信息从具体的道具名称改为通用提示,避免暴露内部实现细节,
提升用户体验和安全性。
- resolve Bug: 使用道具功能优化
Fixes #2060
2025-10-09 09:01:20 +08:00
Rumioandwebjoin111 74a9f3a843 ✨ feat(core): 支持LLM多图片响应,增强UI主题皮肤系统及优化JSON/Markdown处理 (#2062)
- 【LLM服务】
  - `LLMResponse` 模型现在支持 `images: list[bytes]`,允许模型返回多张图片。
  - LLM适配器 (`base.py`, `gemini.py`) 和 API 层 (`api.py`, `service.py`) 已更新以处理多图片响应。
  - 响应验证逻辑已调整,以检查 `images` 列表而非单个 `image_bytes`。
- 【UI渲染服务】
  - 引入组件“皮肤”(variant)概念,允许为同一组件提供不同视觉风格。
  - 改进了 `manifest.json` 的加载、合并和缓存机制,支持基础清单与皮肤清单的递归合并。
  - `ThemeManager` 现在会缓存已加载的清单,并在主题重载时清除缓存。
  - 增强了资源解析器 (`ResourceResolver`),支持 `@` 命名空间路径和更健壮的相对路径处理。
  - 独立模板现在会继承主 Jinja 环境的过滤器。
- 【工具函数】
  - 引入 `dump_json_safely` 工具函数,用于更安全地序列化包含 Pydantic 模型、枚举等复杂类型的对象为 JSON。
  - LLM 服务中的请求体和缓存键生成已改用 `dump_json_safely`。
  - 优化了 `format_usage_for_markdown` 函数,改进了 Markdown 文本的格式化,确保块级元素前有正确换行,并正确处理段落内硬换行。

Co-authored-by: webjoin111 <455457521@qq.com>
2025-10-09 08:50:40 +08:00
HibiKierandpre-commit-ci[bot] e7f3c210df 修复并发时数据库超时 (#2063)
* 🔧 修复和优化:调整超时设置,重构检查逻辑,简化代码结构

- 在 `chkdsk_hook.py` 中重构 `check` 方法,提取公共逻辑
- 更新 `CacheManager` 中的超时设置,使用新的 `CACHE_TIMEOUT`
- 在 `utils.py` 中添加缓存逻辑,记录数据库操作的执行情况

* ✨ feat(auth): 添加并发控制,优化权限检查逻辑

* Update utils.py

* 🚨 auto fix by pre-commit hooks

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2025-10-09 08:46:08 +08:00
molanp f94121080f fix(check): 修复自检插件在ARM设备下的CPU频率获取逻辑 (#2057)
检查bot是否运行正常 / bot check (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, javascript-typescript) (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
Release Drafter / Update Release Draft (push) Has been cancelled
Force Sync to Aliyun / sync (push) Has been cancelled
Update Version / update-version (push) Has been cancelled
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
- 将插件版本从0.1更新至0.2
- 新增安全获取ARM设备CPU频率的函数get_arm_cpu_freq_safe
- 优化CPU信息采集逻辑,提高在ARM架构下的兼容性
2025-10-01 18:42:47 +08:00
HibiKier 761c8daac4 ✨ feat(configs): 优化 ConfigsManager 中的键值获取逻辑,确保未定义键时自动创建 ConfigGroup 实例 (#2058) 2025-10-01 18:42:19 +08:00
Rumioandwebjoin111 c667fc215e ✨ feat(llm): 增强LLM服务,支持图片生成、响应验证与OpenRouter集成 (#2054)
* ✨ feat(llm): 增强LLM服务,支持图片生成、响应验证与OpenRouter集成

- 【新功能】统一图片生成与编辑API `create_image`,支持文生图、图生图及多图输入
- 【新功能】引入LLM响应验证机制,通过 `validation_policy` 和 `response_validator` 确保响应内容符合预期,例如强制返回图片
- 【新功能】适配OpenRouter API,扩展LLM服务提供商支持,并添加OpenRouter特定请求头
- 【重构】将日志净化逻辑重构至 `log_sanitizer` 模块,提供统一的净化入口,并应用于NoneBot消息、LLM请求/响应日志
- 【修复】优化Gemini适配器,正确解析图片生成响应中的Base64图片数据,并更新模型能力注册表

* ✨ feat(image): 优化图片生成响应并返回完整LLMResponse

* ✨ feat(llm): 为 OpenAI 兼容请求体添加日志净化

* 🐛 fix(ui): 截断UI调试HTML日志中的长base64图片数据

---------

Co-authored-by: webjoin111 <455457521@qq.com>
2025-10-01 18:41:46 +08:00
35 changed files with 930 additions and 351 deletions
+1 -1
View File
@@ -26,7 +26,7 @@ __plugin_meta__ = PluginMetadata(
""".strip(), """.strip(),
extra=PluginExtraData( extra=PluginExtraData(
author="HibiKier", author="HibiKier",
version="0.1", version="0.2",
plugin_type=PluginType.SUPERUSER, plugin_type=PluginType.SUPERUSER,
configs=[ configs=[
RegisterConfig( RegisterConfig(
+45 -35
View File
@@ -1,3 +1,4 @@
import contextlib
from dataclasses import dataclass from dataclasses import dataclass
import os import os
from pathlib import Path from pathlib import Path
@@ -18,7 +19,47 @@ BAIDU_URL = "https://www.baidu.com/"
GOOGLE_URL = "https://www.google.com/" GOOGLE_URL = "https://www.google.com/"
VERSION_FILE = Path() / "__version__" 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 @dataclass
@@ -37,7 +78,7 @@ class CPUInfo:
if _cpu_freq := psutil.cpu_freq(): if _cpu_freq := psutil.cpu_freq():
cpu_freq = round(_cpu_freq.current / 1000, 2) cpu_freq = round(_cpu_freq.current / 1000, 2)
else: else:
cpu_freq = 0 cpu_freq = get_arm_cpu_freq_safe()
return CPUInfo(core=cpu_core, usage=cpu_usage, freq=cpu_freq) return CPUInfo(core=cpu_core, usage=cpu_usage, freq=cpu_freq)
@@ -160,44 +201,13 @@ def __get_version() -> str | None:
return 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: async def get_status_info() -> dict:
"""获取信息""" """获取信息"""
data = await __build_status() data = await __build_status()
system = platform.uname() system = platform.uname()
if system.machine == ARM_KEY and not ( data = data.get_system_info()
cpuinfo.get_cpu_info().get("brand_raw") and data.cpu.freq data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
):
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")
baidu, google = await __get_network_info() baidu, google = await __get_network_info()
data["baidu"] = "#8CC265" if baidu else "red" data["baidu"] = "#8CC265" if baidu else "red"
data["google"] = "#8CC265" if google else "red" data["google"] = "#8CC265" if google else "red"
+2 -2
View File
@@ -74,8 +74,8 @@ async def _(matcher: Matcher, message: UniMsg, session: EventSession):
message_list.append(image) message_list.append(image)
message_list.append( message_list.append(
"桀桀桀,预判到会有 '笨蛋' 把功能名称当命令用,特地前来嘲笑!" "桀桀桀,预判到会有 '笨蛋' 把功能名称当命令用,特地前来嘲笑!"
f"但还是好心来帮帮你啦!\n请at我发送 '帮助{plugin.name}' 或者" f"但还是好心来帮帮你啦!\n请at我发送 '帮助 {plugin.name}' 或者"
f" '帮助{plugin.id}' 来获取该功能帮助!" f" '帮助 {plugin.id}' 来获取该功能帮助!"
) )
logger.info("检测到功能名称当命令使用,已发送帮助信息", "功能帮助", session=session) logger.info("检测到功能名称当命令使用,已发送帮助信息", "功能帮助", session=session)
await MessageUtils.build_message(message_list).send(reply_to=True) await MessageUtils.build_message(message_list).send(reply_to=True)
@@ -58,5 +58,14 @@ Config.add_plugin_config(
type=bool, 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())) nonebot.load_plugins(str(Path(__file__).parent.resolve()))
+7 -16
View File
@@ -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}", f"查询ban记录超时: user_id={user_id}, group_id={group_id}",
LOGGER_COMMAND, LOGGER_COMMAND,
) )
# 超时时返回0,避免阻塞
return 0 return 0
# 检查记录并计算ban时间 # 检查记录并计算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检查 """用户ban检查
参数: 参数:
@@ -217,22 +216,12 @@ async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
if not time_val: if not time_val:
return return
time_str = format_time(time_val) 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 ( if (
db_plugin plugin
and not db_plugin.ignore_prompt
and time_val != -1 and time_val != -1
and ban_result 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: try:
await asyncio.wait_for( 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 检查 """权限检查 - ban 检查
参数: 参数:
@@ -289,7 +280,7 @@ async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
if entity.user_id: if entity.user_id:
try: try:
await asyncio.wait_for( await asyncio.wait_for(
user_handle(matcher.plugin_name, entity, session), user_handle(plugin, entity, session),
timeout=DB_TIMEOUT_SECONDS, timeout=DB_TIMEOUT_SECONDS,
) )
except asyncio.TimeoutError: except asyncio.TimeoutError:
@@ -1,50 +1,36 @@
import asyncio
import time import time
from nonebot_plugin_alconna import UniMsg from nonebot_plugin_alconna import UniMsg
from zhenxun.models.group_console import GroupConsole from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo 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.services.log import logger
from zhenxun.utils.utils import EntityIDs
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
from .exception import SkipPluginException 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 plugin: PluginInfo
entity: EntityIDs group: GroupConsole
message: UniMsg message: UniMsg
""" """
start_time = time.time() if not group_id:
if not entity.group_id:
return return
start_time = time.time()
try: try:
text = message.extract_plain_text() 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: if not group:
raise SkipPluginException("群组信息不存在...") raise SkipPluginException("群组信息不存在...")
if group.level < 0: if group.level < 0:
@@ -63,6 +49,5 @@ async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
logger.warning( logger.warning(
f"auth_group 耗时: {elapsed:.3f}s, plugin={plugin.module}", f"auth_group 耗时: {elapsed:.3f}s, plugin={plugin.module}",
LOGGER_COMMAND, LOGGER_COMMAND,
session=entity.user_id, group_id=group_id,
group_id=entity.group_id,
) )
@@ -6,12 +6,10 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.group_console import GroupConsole from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo 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.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.enum import BlockType from zhenxun.utils.enum import BlockType
from zhenxun.utils.utils import get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import IsSuperuserException, SkipPluginException from .exception import IsSuperuserException, SkipPluginException
@@ -20,30 +18,17 @@ from .utils import freq, is_poke, send_message
class GroupCheck: class GroupCheck:
def __init__( def __init__(
self, plugin: PluginInfo, group_id: str, session: Uninfo, is_poke: bool self, plugin: PluginInfo, group: GroupConsole, session: Uninfo, is_poke: bool
) -> None: ) -> None:
self.group_id = group_id
self.session = session self.session = session
self.is_poke = is_poke self.is_poke = is_poke
self.plugin = plugin self.plugin = plugin
self.group_dao = DataAccess(GroupConsole) self.group_data = group
self.group_data = None self.group_id = group.group_id
async def check(self): async def check(self):
start_time = time.time() start_time = time.time()
try: 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 ( if (
self.group_data self.group_data
@@ -113,12 +98,13 @@ class GroupCheck:
class PluginCheck: 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.session = session
self.is_poke = is_poke self.is_poke = is_poke
self.group_id = group_id self.group_data = group
self.group_dao = DataAccess(GroupConsole) self.group_id = None
self.group_data = None if group:
self.group_id = group.group_id
async def check_user(self, plugin: PluginInfo): async def check_user(self, plugin: PluginInfo):
"""全局私聊禁用检测 """全局私聊禁用检测
@@ -156,21 +142,8 @@ class PluginCheck:
if plugin.status or plugin.block_type != BlockType.ALL: if plugin.status or plugin.block_type != BlockType.ALL:
return return
"""全局状态""" """全局状态"""
if self.group_id: if self.group_data and self.group_data.is_super:
# 使用 DataAccess 的缓存机制 raise IsSuperuserException()
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()
sid = self.group_id or self.session.user.id sid = self.group_id or self.session.user.id
if freq.is_send_limit_message(plugin, sid, self.is_poke): 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() start_time = time.time()
try: try:
entity = get_entity_ids(session)
is_poke_event = is_poke(event) 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: tasks = []
group_check = GroupCheck(plugin, entity.group_id, session, is_poke_event) if group:
try: tasks.append(GroupCheck(plugin, group, session, is_poke_event).check())
await asyncio.wait_for(
group_check.check(), timeout=DB_TIMEOUT_SECONDS * 2
)
except asyncio.TimeoutError:
logger.error(f"群组检查超时: {entity.group_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
else: else:
try: tasks.append(user_check.check_user(plugin))
await asyncio.wait_for( tasks.append(user_check.check_global(plugin))
user_check.check_user(plugin), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error("用户检查超时", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
try: try:
await asyncio.wait_for( await asyncio.wait_for(
user_check.check_global(plugin), timeout=DB_TIMEOUT_SECONDS asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS * 2
) )
except asyncio.TimeoutError: except asyncio.TimeoutError:
logger.error("全局检查超时", LOGGER_COMMAND) logger.error("插件用户/群组/全局检查超时...", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
finally: finally:
# 记录总执行时间 # 记录总执行时间
elapsed = time.time() - start_time elapsed = time.time() - start_time
+1 -1
View File
@@ -85,7 +85,7 @@ class FreqUtils:
return False return False
if plugin.plugin_type == PluginType.DEPENDANT: if plugin.plugin_type == PluginType.DEPENDANT:
return False 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() freq = FreqUtils()
+73 -4
View File
@@ -8,6 +8,7 @@ from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo from nonebot_plugin_uninfo import Uninfo
from tortoise.exceptions import IntegrityError from tortoise.exceptions import IntegrityError
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole from zhenxun.models.user_console import UserConsole
from zhenxun.services.data_access import DataAccess from zhenxun.services.data_access import DataAccess
@@ -31,6 +32,7 @@ from .auth.exception import (
PermissionExemption, PermissionExemption,
SkipPluginException, SkipPluginException,
) )
from .auth.utils import base_config
# 超时设置(秒) # 超时设置(秒)
TIMEOUT_SECONDS = 5.0 TIMEOUT_SECONDS = 5.0
@@ -46,6 +48,16 @@ CIRCUIT_BREAKERS = {
# 熔断重置时间(秒) # 熔断重置时间(秒)
CIRCUIT_RESET_TIME = 300 # 5分钟 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): 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" 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( async def auth(
matcher: Matcher, matcher: Matcher,
event: Event, event: Event,
@@ -285,6 +321,9 @@ async def auth(
hook_times = {} hook_times = {}
hooks_time = 0 # 初始化 hooks_time 变量 hooks_time = 0 # 初始化 hooks_time 变量
# 记录是否已进入 hooks 区域(用于 finally 中释放)
entered_hooks = False
try: try:
if not module: if not module:
raise PermissionExemption("Matcher插件名称不存在...") raise PermissionExemption("Matcher插件名称不存在...")
@@ -304,6 +343,10 @@ async def auth(
) )
raise PermissionExemption("获取插件和用户数据超时,请稍后再试...") raise PermissionExemption("获取插件和用户数据超时,请稍后再试...")
# 进入 hooks 并行检查区域(会在高并发时排队)
await _enter_hooks_section()
entered_hooks = True
# 获取插件费用 # 获取插件费用
cost_start = time.time() cost_start = time.time()
try: try:
@@ -320,16 +363,32 @@ async def auth(
# 执行 bot_filter # 执行 bot_filter
bot_filter(session) 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 检查,并记录执行时间 # 并行执行所有 hook 检查,并记录执行时间
hooks_start = time.time() hooks_start = time.time()
# 创建所有 hook 任务 # 创建所有 hook 任务
hook_tasks = [ 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_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_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), time_hook(auth_limit(plugin, session), "auth_limit", hook_times),
] ]
@@ -358,7 +417,17 @@ async def auth(
logger.debug("超级用户跳过权限检测...", LOGGER_COMMAND, session=session) logger.debug("超级用户跳过权限检测...", LOGGER_COMMAND, session=session)
except PermissionExemption as e: except PermissionExemption as e:
logger.info(str(e), LOGGER_COMMAND, session=session) 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: if not ignore_flag and cost_gold > 0:
gold_start = time.time() gold_start = time.time()
+3 -32
View File
@@ -1,12 +1,12 @@
from typing import Any from typing import Any
from nonebot.adapters import Bot, Message from nonebot.adapters import Bot, Message
from nonebot.adapters.onebot.v11 import MessageSegment
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
from zhenxun.models.bot_message_store import BotMessageStore from zhenxun.models.bot_message_store import BotMessageStore
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.enum import BotSentType 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.manager.message_manager import MessageManager
from zhenxun.utils.platform import PlatformUtils from zhenxun.utils.platform import PlatformUtils
@@ -41,35 +41,6 @@ def replace_message(message: Message) -> str:
return result 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 @Bot.on_called_api
async def handle_api_result( async def handle_api_result(
bot: Bot, exception: Exception | None, api: str, data: dict[str, Any], result: Any 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: Message = data.get("message", "")
message_type = data.get("message_type") message_type = data.get("message_type")
try: try:
# 记录消息id
if user_id and message_id: if user_id and message_id:
MessageManager.add(str(user_id), str(message_id)) MessageManager.add(str(user_id), str(message_id))
logger.debug( logger.debug(
@@ -108,7 +78,8 @@ async def handle_api_result(
else replace_message(message), else replace_message(message),
platform=PlatformUtils.get_platform(bot), 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: except Exception as e:
logger.warning( logger.warning(
f"消息发送记录发生错误...data: {data}, result: {result}", f"消息发送记录发生错误...data: {data}, result: {result}",
+45 -45
View File
@@ -43,18 +43,20 @@ class BanCheckLimiter:
def check(self, key: str | float) -> bool: def check(self, key: str | float) -> bool:
if time.time() - self.mtime[key] > self.default_check_time: if time.time() - self.mtime[key] > self.default_check_time:
self.mtime[key] = time.time() return self._extracted_from_check_3(key, False)
self.mint[key] = 0
return False
if ( if (
self.mint[key] >= self.default_count self.mint[key] >= self.default_count
and time.time() - self.mtime[key] < self.default_check_time and time.time() - self.mtime[key] < self.default_check_time
): ):
self.mtime[key] = time.time() return self._extracted_from_check_3(key, True)
self.mint[key] = 0
return True
return False 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( _blmt = BanCheckLimiter(
malicious_check_time, malicious_check_time,
@@ -70,16 +72,15 @@ async def _(
module = None module = None
if plugin := matcher.plugin: if plugin := matcher.plugin:
module = plugin.module_name module = plugin.module_name
if metadata := plugin.metadata: if not (metadata := plugin.metadata):
extra = metadata.extra return
if extra.get("plugin_type") in [ extra = metadata.extra
PluginType.HIDDEN, if extra.get("plugin_type") in [
PluginType.DEPENDANT, PluginType.HIDDEN,
PluginType.ADMIN, PluginType.DEPENDANT,
PluginType.SUPERUSER, PluginType.ADMIN,
]: PluginType.SUPERUSER,
return ]:
else:
return return
if matcher.type == "notice": if matcher.type == "notice":
return return
@@ -88,32 +89,31 @@ async def _(
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME") malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
if not malicious_ban_time: if not malicious_ban_time:
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0") raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0")
if user_id: if user_id and module:
if module: if _blmt.check(f"{user_id}__{module}"):
if _blmt.check(f"{user_id}__{module}"): await BanConsole.ban(
await BanConsole.ban( user_id,
user_id, group_id,
group_id, 9,
9, "恶意触发命令检测",
"恶意触发命令检测", malicious_ban_time * 60,
malicious_ban_time * 60, bot.self_id,
bot.self_id, )
) logger.info(
logger.info( f"触发了恶意触发检测: {matcher.plugin_name}",
f"触发了恶意触发检测: {matcher.plugin_name}", "HOOK",
"HOOK", session=session,
session=session, )
) await MessageUtils.build_message(
await MessageUtils.build_message( [
[ At(flag="user", target=user_id),
At(flag="user", target=user_id), "检测到恶意触发命令,您将被封禁 30 分钟",
"检测到恶意触发命令,您将被封禁 30 分钟", ]
] ).send()
).send() logger.debug(
logger.debug( f"触发了恶意触发检测: {matcher.plugin_name}",
f"触发了恶意触发检测: {matcher.plugin_name}", "HOOK",
"HOOK", session=session,
session=session, )
) raise IgnoredException("检测到恶意触发命令")
raise IgnoredException("检测到恶意触发命令") _blmt.add(f"{user_id}__{module}")
_blmt.add(f"{user_id}__{module}")
@@ -263,10 +263,9 @@ class StoreManager:
"""安装插件 """安装插件
参数: 参数:
github_url: 仓库地址 plugin_info: 插件信息
module_path: 模块路径
is_dir: 是否是文件夹
is_external: 是否是外部仓库 is_external: 是否是外部仓库
source: 源
""" """
repo_type = RepoType.GITHUB if is_external else None repo_type = RepoType.GITHUB if is_external else None
if source == "ali": if source == "ali":
+1 -1
View File
@@ -367,7 +367,7 @@ class ShopManage:
else: else:
goods_info = await GoodsInfo.get_or_none(goods_name=goods_name) goods_info = await GoodsInfo.get_or_none(goods_name=goods_name)
if not goods_info: if not goods_info:
return f"{goods_name} 不存在..." return "对应的道具不存在..."
if goods_info.is_passive: if goods_info.is_passive:
return f"{goods_info.goods_name} 是被动道具, 无法使用..." return f"{goods_info.goods_name} 是被动道具, 无法使用..."
goods = cls.uuid2goods.get(goods_info.uuid) goods = cls.uuid2goods.get(goods_info.uuid)
+3 -1
View File
@@ -344,7 +344,9 @@ class ConfigsManager:
返回: 返回:
ConfigGroup: ConfigGroup 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): def save(self, path: str | Path | None = None, save_simple_data: bool = False):
"""保存数据 """保存数据
+2 -2
View File
@@ -98,6 +98,7 @@ from .cache_containers import CacheDict, CacheList
from .config import ( from .config import (
CACHE_KEY_PREFIX, CACHE_KEY_PREFIX,
CACHE_KEY_SEPARATOR, CACHE_KEY_SEPARATOR,
CACHE_TIMEOUT,
DEFAULT_EXPIRE, DEFAULT_EXPIRE,
LOG_COMMAND, LOG_COMMAND,
SPECIAL_KEY_FORMATS, SPECIAL_KEY_FORMATS,
@@ -551,7 +552,6 @@ class CacheManager:
返回: 返回:
Any: 缓存数据,如果不存在返回默认值 Any: 缓存数据,如果不存在返回默认值
""" """
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
# 如果缓存被禁用或缓存模式为NONE,直接返回默认值 # 如果缓存被禁用或缓存模式为NONE,直接返回默认值
if not self.enabled or cache_config.cache_mode == CacheMode.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) cache_key = self._build_key(cache_type, key)
data = await asyncio.wait_for( data = await asyncio.wait_for(
self.cache_backend.get(cache_key), # type: ignore self.cache_backend.get(cache_key), # type: ignore
timeout=DB_TIMEOUT_SECONDS, timeout=CACHE_TIMEOUT,
) )
if data is None: if data is None:
+3
View File
@@ -5,6 +5,9 @@
# 日志标识 # 日志标识
LOG_COMMAND = "CacheRoot" LOG_COMMAND = "CacheRoot"
# 缓存获取超时时间(秒)
CACHE_TIMEOUT = 10
# 默认缓存过期时间(秒) # 默认缓存过期时间(秒)
DEFAULT_EXPIRE = 600 DEFAULT_EXPIRE = 600
+4 -1
View File
@@ -27,5 +27,8 @@ async def with_db_timeout(
return result return result
except asyncio.TimeoutError: except asyncio.TimeoutError:
if operation: if operation:
logger.error(f"数据库操作超时: {operation} (>{timeout}s)", LOG_COMMAND) logger.error(
f"数据库操作超时: {operation} (>{timeout}s) 来源: {source}",
LOG_COMMAND,
)
raise raise
+2
View File
@@ -7,6 +7,7 @@ LLM 服务模块 - 公共 API 入口
from .api import ( from .api import (
chat, chat,
code, code,
create_image,
embed, embed,
generate, generate,
generate_structured, generate_structured,
@@ -74,6 +75,7 @@ __all__ = [
"chat", "chat",
"clear_model_cache", "clear_model_cache",
"code", "code",
"create_image",
"create_multimodal_message", "create_multimodal_message",
"embed", "embed",
"function_tool", "function_tool",
+44
View File
@@ -3,6 +3,9 @@ LLM 适配器基类和通用数据结构
""" """
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
import base64
import binascii
import json
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from pydantic import BaseModel from pydantic import BaseModel
@@ -32,6 +35,7 @@ class ResponseData(BaseModel):
"""响应数据封装 - 支持所有高级功能""" """响应数据封装 - 支持所有高级功能"""
text: str text: str
images: list[bytes] | None = None
usage_info: dict[str, Any] | None = None usage_info: dict[str, Any] | None = None
raw_response: dict[str, Any] | None = None raw_response: dict[str, Any] | None = None
tool_calls: list[LLMToolCall] | None = None tool_calls: list[LLMToolCall] | None = None
@@ -242,6 +246,38 @@ class BaseAdapter(ABC):
if content: if content:
content = content.strip() 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 parsed_tool_calls: list[LLMToolCall] | None = None
if message_tool_calls := message.get("tool_calls"): if message_tool_calls := message.get("tool_calls"):
from ..types.models import LLMToolFunction from ..types.models import LLMToolFunction
@@ -280,6 +316,7 @@ class BaseAdapter(ABC):
text=final_text, text=final_text,
tool_calls=parsed_tool_calls, tool_calls=parsed_tool_calls,
usage_info=usage_info, usage_info=usage_info,
images=images_bytes if images_bytes else None,
raw_response=response_json, raw_response=response_json,
) )
@@ -450,6 +487,13 @@ class OpenAICompatAdapter(BaseAdapter):
"""准备高级请求 - OpenAI兼容格式""" """准备高级请求 - OpenAI兼容格式"""
url = self.get_api_url(model, self.get_chat_endpoint(model)) url = self.get_api_url(model, self.get_chat_endpoint(model))
headers = self.get_base_headers(api_key) 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) openai_messages = self.convert_messages_to_openai_format(messages)
body = { body = {
+18 -1
View File
@@ -2,6 +2,7 @@
Gemini API 适配器 Gemini API 适配器
""" """
import base64
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from zhenxun.services.log import logger from zhenxun.services.log import logger
@@ -373,7 +374,16 @@ class GeminiAdapter(BaseAdapter):
self.validate_response(response_json) self.validate_response(response_json)
try: 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: if not candidates:
logger.debug("Gemini响应中没有candidates。") logger.debug("Gemini响应中没有candidates。")
return ResponseData(text="", raw_response=response_json) return ResponseData(text="", raw_response=response_json)
@@ -398,6 +408,7 @@ class GeminiAdapter(BaseAdapter):
parts = content_data.get("parts", []) parts = content_data.get("parts", [])
text_content = "" text_content = ""
images_bytes: list[bytes] = []
parsed_tool_calls: list["LLMToolCall"] | None = None parsed_tool_calls: list["LLMToolCall"] | None = None
thought_summary_parts = [] thought_summary_parts = []
answer_parts = [] answer_parts = []
@@ -409,6 +420,11 @@ class GeminiAdapter(BaseAdapter):
thought_summary_parts.append(part["thought"]) thought_summary_parts.append(part["thought"])
elif "thoughtSummary" in part: elif "thoughtSummary" in part:
thought_summary_parts.append(part["thoughtSummary"]) 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: elif "functionCall" in part:
if parsed_tool_calls is None: if parsed_tool_calls is None:
parsed_tool_calls = [] parsed_tool_calls = []
@@ -475,6 +491,7 @@ class GeminiAdapter(BaseAdapter):
return ResponseData( return ResponseData(
text=text_content, text=text_content,
tool_calls=parsed_tool_calls, tool_calls=parsed_tool_calls,
images=images_bytes if images_bytes else None,
usage_info=usage_info, usage_info=usage_info,
raw_response=response_json, raw_response=response_json,
grounding_metadata=grounding_metadata_obj, grounding_metadata=grounding_metadata_obj,
+8 -1
View File
@@ -21,7 +21,14 @@ class OpenAIAdapter(OpenAICompatAdapter):
@property @property
def supported_api_types(self) -> list[str]: 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: def get_chat_endpoint(self, model: "LLMModel") -> str:
"""返回聊天完成端点""" """返回聊天完成端点"""
+100 -2
View File
@@ -2,7 +2,8 @@
LLM 服务的高级 API 接口 - 便捷函数入口 (无状态) 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 nonebot_plugin_alconna.uniseg import UniMessage
from pydantic import BaseModel from pydantic import BaseModel
@@ -10,7 +11,7 @@ from pydantic import BaseModel
from zhenxun.services.log import logger from zhenxun.services.log import logger
from .config import CommonOverrides 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 .manager import get_model_instance
from .session import AI from .session import AI
from .tools.manager import tool_provider_manager from .tools.manager import tool_provider_manager
@@ -23,6 +24,7 @@ from .types import (
LLMResponse, LLMResponse,
ModelName, ModelName,
) )
from .utils import create_multimodal_message
T = TypeVar("T", bound=BaseModel) T = TypeVar("T", bound=BaseModel)
@@ -303,3 +305,99 @@ async def run_with_tools(
raise LLMException( raise LLMException(
"带工具的执行循环未能产生有效的助手回复。", code=LLMErrorCode.GENERATION_FAILED "带工具的执行循环未能产生有效的助手回复。", 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)
+12 -1
View File
@@ -2,13 +2,15 @@
LLM 生成配置相关类和函数 LLM 生成配置相关类和函数
""" """
from collections.abc import Callable
from typing import Any from typing import Any
from pydantic import BaseModel, Field from pydantic import BaseModel, ConfigDict, Field
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.pydantic_compat import model_dump from zhenxun.utils.pydantic_compat import model_dump
from ..types import LLMResponse
from ..types.enums import ResponseFormat from ..types.enums import ResponseFormat
from ..types.exceptions import LLMErrorCode, LLMException from ..types.exceptions import LLMErrorCode, LLMException
@@ -64,6 +66,15 @@ class ModelConfigOverride(BaseModel):
custom_params: dict[str, Any] | None = Field(default=None, description="自定义参数") 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]: def to_dict(self) -> dict[str, Any]:
"""转换为字典,排除None值""" """转换为字典,排除None值"""
+14 -18
View File
@@ -50,8 +50,8 @@ class LLMHttpClient:
async with self._lock: async with self._lock:
if self._client is None or self._client.is_closed: if self._client is None or self._client.is_closed:
logger.debug( logger.debug(
f"LLMHttpClient: Initializing new httpx.AsyncClient " f"LLMHttpClient: 正在初始化新的 httpx.AsyncClient "
f"with config: {self.config}" f"配置: {self.config}"
) )
headers = get_user_agent() headers = get_user_agent()
limits = httpx.Limits( limits = httpx.Limits(
@@ -92,7 +92,7 @@ class LLMHttpClient:
) )
if self._client is None: if self._client is None:
raise LLMException( raise LLMException(
"HTTP client failed to initialize.", LLMErrorCode.CONFIGURATION_ERROR "HTTP 客户端初始化失败。", LLMErrorCode.CONFIGURATION_ERROR
) )
return self._client return self._client
@@ -110,17 +110,17 @@ class LLMHttpClient:
async with self._lock: async with self._lock:
if self._client and not self._client.is_closed: if self._client and not self._client.is_closed:
logger.debug( logger.debug(
f"LLMHttpClient: Closing with config: {self.config}. " f"LLMHttpClient: 正在关闭,配置: {self.config}. "
f"Active requests: {self._active_requests}" f"活跃请求数: {self._active_requests}"
) )
if self._active_requests > 0: if self._active_requests > 0:
logger.warning( logger.warning(
f"LLMHttpClient: Closing while {self._active_requests} " f"LLMHttpClient: 关闭时仍有 {self._active_requests} "
f"requests are still active." f"个请求处于活跃状态。"
) )
await self._client.aclose() await self._client.aclose()
self._client = None self._client = None
logger.debug(f"LLMHttpClient for config {self.config} definitively closed.") logger.debug(f"配置为 {self.config} 的 LLMHttpClient 已完全关闭。")
@property @property
def is_closed(self) -> bool: def is_closed(self) -> bool:
@@ -145,20 +145,17 @@ class LLMHttpClientManager:
client = self._clients.get(key) client = self._clients.get(key)
if client and not client.is_closed: if client and not client.is_closed:
logger.debug( logger.debug(
f"LLMHttpClientManager: Reusing existing LLMHttpClient " f"LLMHttpClientManager: 复用现有的 LLMHttpClient 密钥: {key}"
f"for key: {key}"
) )
return client return client
if client and client.is_closed: if client and client.is_closed:
logger.debug( logger.debug(
f"LLMHttpClientManager: Found a closed client for key {key}. " f"LLMHttpClientManager: 发现密钥 {key} 对应的客户端已关闭。"
f"Creating a new one." f"正在创建新的客户端。"
) )
logger.debug( logger.debug(f"LLMHttpClientManager: 为密钥 {key} 创建新的 LLMHttpClient")
f"LLMHttpClientManager: Creating new LLMHttpClient for key: {key}"
)
http_client_config = HttpClientConfig( http_client_config = HttpClientConfig(
timeout=provider_config.timeout, proxy=provider_config.proxy timeout=provider_config.timeout, proxy=provider_config.proxy
) )
@@ -169,8 +166,7 @@ class LLMHttpClientManager:
async def shutdown(self): async def shutdown(self):
async with self._lock: async with self._lock:
logger.info( logger.info(
f"LLMHttpClientManager: Shutting down. " f"LLMHttpClientManager: 正在关闭。关闭 {len(self._clients)} 个客户端。"
f"Closing {len(self._clients)} client(s)."
) )
close_tasks = [ close_tasks = [
client.close() client.close()
@@ -180,7 +176,7 @@ class LLMHttpClientManager:
if close_tasks: if close_tasks:
await asyncio.gather(*close_tasks, return_exceptions=True) await asyncio.gather(*close_tasks, return_exceptions=True)
self._clients.clear() self._clients.clear()
logger.info("LLMHttpClientManager: Shutdown complete.") logger.info("LLMHttpClientManager: 关闭完成。")
http_client_manager = LLMHttpClientManager() http_client_manager = LLMHttpClientManager()
+3 -2
View File
@@ -5,12 +5,12 @@ LLM 模型管理器
""" """
import hashlib import hashlib
import json
import time import time
from typing import Any from typing import Any
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.pydantic_compat import dump_json_safely
from .config import validate_override_params from .config import validate_override_params
from .config.providers import AI_CONFIG_GROUP, PROVIDERS_CONFIG_KEY, get_ai_config from .config.providers import AI_CONFIG_GROUP, PROVIDERS_CONFIG_KEY, get_ai_config
@@ -43,7 +43,7 @@ def _make_cache_key(
) -> str: ) -> str:
"""生成缓存键""" """生成缓存键"""
config_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}" key_data = f"{provider_model_name}:{config_str}"
return hashlib.md5(key_data.encode()).hexdigest() 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", "deepseek": "https://api.deepseek.com",
"zhipu": "https://open.bigmodel.cn", "zhipu": "https://open.bigmodel.cn",
"gemini": "https://generativelanguage.googleapis.com", "gemini": "https://generativelanguage.googleapis.com",
"openrouter": "https://openrouter.ai/api",
"general_openai_compat": None, "general_openai_compat": None,
} }
+81 -9
View File
@@ -12,6 +12,8 @@ from typing import Any, TypeVar
from pydantic import BaseModel from pydantic import BaseModel
from zhenxun.services.log import logger 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 .adapters.base import RequestData
from .config import LLMGenerationConfig from .config import LLMGenerationConfig
@@ -34,7 +36,6 @@ from .types import (
ToolExecutable, ToolExecutable,
) )
from .types.capabilities import ModelCapabilities, ModelModality from .types.capabilities import ModelCapabilities, ModelModality
from .utils import _sanitize_request_body_for_logging
T = TypeVar("T", bound=BaseModel) T = TypeVar("T", bound=BaseModel)
@@ -187,21 +188,32 @@ class LLMModel(LLMModelBase):
logger.debug(f"🔑 API密钥: {masked_key}") logger.debug(f"🔑 API密钥: {masked_key}")
logger.debug(f"📋 请求头: {dict(request_data.headers)}") logger.debug(f"📋 请求头: {dict(request_data.headers)}")
sanitized_body = _sanitize_request_body_for_logging(request_data.body) sanitizer_req_context_map = {"gemini": "gemini_request"}
request_body_str = json.dumps(sanitized_body, ensure_ascii=False, indent=2) 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}") logger.debug(f"📦 请求体: {request_body_str}")
http_response = await http_client.post( http_response = await http_client.post(
request_data.url, request_data.url,
headers=request_data.headers, 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"📥 响应状态码: {http_response.status_code}")
logger.debug(f"📄 响应头: {dict(http_response.headers)}") 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: if http_response.status_code != 200:
error_text = http_response.text error_text = response_bytes.decode("utf-8", errors="ignore")
logger.error( logger.error(
f"❌ HTTP请求失败: {http_response.status_code} - {error_text} " f"❌ HTTP请求失败: {http_response.status_code} - {error_text} "
f"[{log_context}]" f"[{log_context}]"
@@ -232,13 +244,22 @@ class LLMModel(LLMModelBase):
) )
try: 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_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}") logger.debug(f"📋 响应JSON: {response_json_str}")
parsed_data = parse_response_func(response_json) parsed_data = parse_response_func(response_json)
except Exception as e: except Exception as e:
logger.error(f"解析 {log_context} 响应失败: {e}", e=e) logger.error(f"解析 {log_context} 响应失败: {e}", e=e)
await self.key_store.record_failure(api_key, None, str(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) adapter.validate_embedding_response(response_json)
return adapter.parse_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, prepare_request_func=prepare_request,
parse_response_func=parse_response, parse_response_func=parse_response,
http_client=http_client, http_client=http_client,
@@ -376,6 +397,7 @@ class LLMModel(LLMModelBase):
return LLMResponse( return LLMResponse(
text=response_data.text, text=response_data.text,
usage_info=response_data.usage_info, usage_info=response_data.usage_info,
images=response_data.images,
raw_response=response_data.raw_response, raw_response=response_data.raw_response,
tool_calls=response_tool_calls if response_tool_calls else None, tool_calls=response_tool_calls if response_tool_calls else None,
code_executions=response_data.code_executions, code_executions=response_data.code_executions,
@@ -390,6 +412,56 @@ class LLMModel(LLMModelBase):
failed_keys=failed_keys, failed_keys=failed_keys,
log_context="Generation", 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 return parsed_data, api_key_used
async def close(self): async def close(self):
@@ -44,6 +44,13 @@ GEMINI_CAPABILITIES = ModelCapabilities(
supports_tool_calling=True, 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( DOUBAO_ADVANCED_MULTIMODAL_CAPABILITIES = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.IMAGE, ModelModality.VIDEO}, input_modalities={ModelModality.TEXT, ModelModality.IMAGE, ModelModality.VIDEO},
output_modalities={ModelModality.TEXT}, output_modalities={ModelModality.TEXT},
@@ -83,6 +90,7 @@ MODEL_CAPABILITIES_REGISTRY: dict[str, ModelCapabilities] = {
output_modalities={ModelModality.EMBEDDING}, output_modalities={ModelModality.EMBEDDING},
is_embedding_model=True, is_embedding_model=True,
), ),
"*gemini-*-image-preview*": GEMINI_IMAGE_GEN_CAPABILITIES,
"gemini-2.5-pro*": GEMINI_CAPABILITIES, "gemini-2.5-pro*": GEMINI_CAPABILITIES,
"gemini-1.5-pro*": GEMINI_CAPABILITIES, "gemini-1.5-pro*": GEMINI_CAPABILITIES,
"gemini-2.5-flash*": GEMINI_CAPABILITIES, "gemini-2.5-flash*": GEMINI_CAPABILITIES,
+1
View File
@@ -425,6 +425,7 @@ class LLMResponse(BaseModel):
"""LLM 响应""" """LLM 响应"""
text: str text: str
images: list[bytes] | None = None
usage_info: dict[str, Any] | None = None usage_info: dict[str, Any] | None = None
raw_response: dict[str, Any] | None = None raw_response: dict[str, Any] | None = None
tool_calls: list[Any] | None = None tool_calls: list[Any] | None = None
-48
View File
@@ -273,54 +273,6 @@ def message_to_unimessage(message: PlatformMessage) -> UniMessage:
return UniMessage(uni_segments) 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: def sanitize_schema_for_llm(schema: Any, api_type: str) -> Any:
""" """
递归地净化 JSON Schema,移除特定 LLM API 不支持的关键字。 递归地净化 JSON Schema,移除特定 LLM API 不支持的关键字。
+18 -12
View File
@@ -22,6 +22,7 @@ from zhenxun.configs.config import Config
from zhenxun.configs.path_config import THEMES_PATH, UI_CACHE_PATH from zhenxun.configs.path_config import THEMES_PATH, UI_CACHE_PATH
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.exception import RenderingError 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 zhenxun.utils.pydantic_compat import _dump_pydantic_obj
from .config import RESERVED_TEMPLATE_KEYS from .config import RESERVED_TEMPLATE_KEYS
@@ -216,16 +217,17 @@ class RendererService:
context.processed_components.add(component_id) context.processed_components.add(component_id)
component_path_base = str(component.template_name) component_path_base = str(component.template_name)
variant = getattr(component, "variant", None)
manifest = await context.theme_manager.get_template_manifest( manifest = await context.theme_manager.get_template_manifest(
component_path_base component_path_base, skin=variant
) )
style_paths_to_load = [] style_paths_to_load = []
if manifest and manifest.styles: if manifest and "styles" in manifest:
styles = ( styles = (
[manifest.styles] [manifest["styles"]]
if isinstance(manifest.styles, str) if isinstance(manifest["styles"], str)
else manifest.styles else manifest["styles"]
) )
for style_path in styles: for style_path in styles:
full_style_path = str(Path(component_path_base) / style_path).replace( 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.globals.update(context.theme_manager.jinja_env.globals)
temp_env.filters.update(context.theme_manager.jinja_env.filters)
temp_env.globals["asset"] = ( temp_env.globals["asset"] = (
context.theme_manager._create_standalone_asset_loader(template_dir) context.theme_manager._create_standalone_asset_loader(template_dir)
) )
@@ -430,10 +433,11 @@ class RendererService:
component_render_options = {} component_render_options = {}
manifest_options = {} manifest_options = {}
variant = getattr(component, "variant", None)
if manifest := await context.theme_manager.get_template_manifest( 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 = component_render_options.copy()
final_render_options.update(manifest_options) final_render_options.update(manifest_options)
@@ -470,10 +474,7 @@ class RendererService:
) from e ) from e
async def render( async def render(
self, self, component: Renderable, use_cache: bool = False, **render_options
component: Renderable,
use_cache: bool = False,
**render_options,
) -> bytes: ) -> bytes:
""" """
统一的、多态的渲染入口,直接返回图片字节。 统一的、多态的渲染入口,直接返回图片字节。
@@ -504,9 +505,12 @@ class RendererService:
) )
result = await self._render_component(context) result = await self._render_component(context)
if Config.get_config("UI", "DEBUG_MODE") and result.html_content: if Config.get_config("UI", "DEBUG_MODE") and result.html_content:
sanitized_html = sanitize_for_logging(
result.html_content, context="ui_html"
)
logger.info( logger.info(
f"--- [UI DEBUG] HTML for {component.__class__.__name__} ---\n" f"--- [UI DEBUG] HTML for {component.__class__.__name__} ---\n"
f"{result.html_content}\n" f"{sanitized_html}\n"
f"--- [UI DEBUG] End of HTML ---" f"--- [UI DEBUG] End of HTML ---"
) )
if result.image_bytes is None: if result.image_bytes is None:
@@ -556,6 +560,8 @@ class RendererService:
await self.initialize() await self.initialize()
assert self._theme_manager is not None, "ThemeManager 未初始化" 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") current_theme_name = Config.get_config("UI", "THEME", "default")
await self._theme_manager.load_theme(current_theme_name) await self._theme_manager.load_theme(current_theme_name)
logger.info(f"主题 '{current_theme_name}' 已成功重载。") logger.info(f"主题 '{current_theme_name}' 已成功重载。")
+138 -24
View File
@@ -1,11 +1,11 @@
from __future__ import annotations from __future__ import annotations
import asyncio
from collections.abc import Callable from collections.abc import Callable
import os import os
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
import aiofiles
from jinja2 import ( from jinja2 import (
ChoiceLoader, ChoiceLoader,
Environment, Environment,
@@ -21,7 +21,6 @@ import ujson as json
from zhenxun.configs.path_config import THEMES_PATH from zhenxun.configs.path_config import THEMES_PATH
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.services.renderer.models import TemplateManifest
from zhenxun.services.renderer.protocols import Renderable from zhenxun.services.renderer.protocols import Renderable
from zhenxun.services.renderer.registry import asset_registry from zhenxun.services.renderer.registry import asset_registry
from zhenxun.utils.pydantic_compat import model_dump from zhenxun.utils.pydantic_compat import model_dump
@@ -32,6 +31,20 @@ if TYPE_CHECKING:
from .config import RESERVED_TEMPLATE_KEYS 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): class RelativePathEnvironment(Environment):
""" """
一个自定义的 Jinja2 环境,重写了 join_path 方法以支持模板间的相对路径引用。 一个自定义的 Jinja2 环境,重写了 join_path 方法以支持模板间的相对路径引用。
@@ -151,14 +164,42 @@ class ResourceResolver:
def resolve_asset_uri(self, asset_path: str, current_template_name: str) -> str: def resolve_asset_uri(self, asset_path: str, current_template_name: str) -> str:
"""解析资源路径,实现完整的回退逻辑,并返回可用的URI。""" """解析资源路径,实现完整的回退逻辑,并返回可用的URI。"""
if not self.theme_manager.current_theme: if (
not self.theme_manager.current_theme
or not self.theme_manager.jinja_env.loader
):
return "" 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]] = [] 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( search_paths.extend(
self._search_paths_for_relative_asset( self._search_paths_for_relative_asset(
asset_path[2:], current_template_name relative_part, current_template_name
) )
) )
else: else:
@@ -209,6 +250,9 @@ class ThemeManager:
self.jinja_env.filters["md"] = self._markdown_filter 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]: def list_available_themes(self) -> list[str]:
"""扫描主题目录并返回所有可用的主题名称。""" """扫描主题目录并返回所有可用的主题名称。"""
if not THEMES_PATH.is_dir(): if not THEMES_PATH.is_dir():
@@ -377,16 +421,26 @@ class ThemeManager:
logger.error(f"指定的模板文件路径不存在: '{component_path_base}'", e=e) logger.error(f"指定的模板文件路径不存在: '{component_path_base}'", e=e)
raise e raise e
entrypoint_filename = "main.html" base_manifest = await self.get_template_manifest(component_path_base)
manifest = await self.get_template_manifest(component_path_base)
if manifest and manifest.entrypoint: skin_to_use = variant or (base_manifest.get("skin") if base_manifest else None)
entrypoint_filename = manifest.entrypoint
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 = [] potential_paths = []
if variant: if skin_to_use:
potential_paths.append( 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}") potential_paths.append(f"{component_path_base}/{entrypoint_filename}")
@@ -410,28 +464,88 @@ class ThemeManager:
logger.error(err_msg) logger.error(err_msg)
raise TemplateNotFound(err_msg) raise TemplateNotFound(err_msg)
async def get_template_manifest( async def _load_single_manifest(self, path_str: str) -> dict[str, Any] | None:
self, component_path: str """从指定路径加载单个 manifest.json 文件。"""
) -> TemplateManifest | None: normalized_path = path_str.replace("\\", "/")
""" manifest_path_str = f"{normalized_path}/manifest.json"
查找并解析组件的 manifest.json 文件。
"""
manifest_path_str = f"{component_path}/manifest.json"
if not self.jinja_env.loader: if not self.jinja_env.loader:
return None return None
try: try:
_, full_path, _ = self.jinja_env.loader.get_source( source, filepath, _ = self.jinja_env.loader.get_source(
self.jinja_env, manifest_path_str self.jinja_env, manifest_path_str
) )
if full_path and Path(full_path).exists(): logger.debug(f"找到清单文件: '{manifest_path_str}' (从 '{filepath}' 加载)")
async with aiofiles.open(full_path, encoding="utf-8") as f: return json.loads(source)
manifest_data = json.loads(await f.read())
return TemplateManifest(**manifest_data)
except TemplateNotFound: except TemplateNotFound:
logger.trace(f"未找到清单文件: '{manifest_path_str}'")
return None 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( async def resolve_markdown_style_path(
self, style_name: str, context: "RenderContext" self, style_name: str, context: "RenderContext"
+7 -4
View File
@@ -126,12 +126,15 @@ class SqlUtils:
def format_usage_for_markdown(text: str) -> str: def format_usage_for_markdown(text: str) -> str:
""" """
智能地将Python多行字符串转换为适合Markdown渲染的格式。 智能地将Python多行字符串转换为适合Markdown渲染的格式。
- 将单个换行符替换为Markdown的硬换行(行尾加两个空格)。 - 在列表、标题等块级元素前自动插入换行,确保正确解析。
- 将段落内的单个换行符替换为Markdown的硬换行(行尾加两个空格)。
- 保留两个或更多的连续换行符,使其成为Markdown的段落分隔。 - 保留两个或更多的连续换行符,使其成为Markdown的段落分隔。
""" """
if not text: if not text:
return "" return ""
text = re.sub(r"\n{2,}", "<<PARAGRAPH_BREAK>>", text)
text = text.replace("\n", " \n") text = re.sub(r"([^\n])\n(\s*[-*] |\s*#+\s|\s*>)", r"\1\n\n\2", text)
text = text.replace("<<PARAGRAPH_BREAK>>", "\n\n")
text = re.sub(r"(?<!\n)\n(?!\n)", " \n", text)
return text return text
+12
View File
@@ -263,6 +263,18 @@ class AsyncHttpx:
) )
return result return result
except Exception as e: 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) exceptions.append(e)
if url != url_list[-1]: if url != url_list[-1]:
logger.warning( logger.warning(
+202
View File
@@ -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
+28
View File
@@ -5,10 +5,14 @@ Pydantic V1 & V2 兼容层模块
包括 model_dump, model_copy, model_json_schema, parse_as 等。 包括 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 typing import Any, TypeVar, get_args, get_origin
from nonebot.compat import PYDANTIC_V2, model_dump from nonebot.compat import PYDANTIC_V2, model_dump
from pydantic import VERSION, BaseModel from pydantic import VERSION, BaseModel
import ujson as json
T = TypeVar("T", bound=BaseModel) T = TypeVar("T", bound=BaseModel)
V = TypeVar("V") V = TypeVar("V")
@@ -19,6 +23,7 @@ __all__ = [
"_dump_pydantic_obj", "_dump_pydantic_obj",
"_is_pydantic_type", "_is_pydantic_type",
"compat_computed_field", "compat_computed_field",
"dump_json_safely",
"model_copy", "model_copy",
"model_dump", "model_dump",
"model_json_schema", "model_json_schema",
@@ -93,3 +98,26 @@ def parse_as(type_: type[V], obj: Any) -> V:
from pydantic import TypeAdapter # type: ignore from pydantic import TypeAdapter # type: ignore
return TypeAdapter(type_).validate_python(obj) 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)