mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-08 21:30:01 +08:00
性能优化 (#2126)
* 性能优化 * 代码改进 * 优化浏览器代际切换逻辑 * 统一缓存与生命周期 * 添加aiomysql依赖 * 优化插件路径处理逻辑,简化条件判断;在虚拟环境包管理器中添加编码和错误处理参数以增强稳定性 * 🚨 auto fix by pre-commit hooks * 优化Windows下的关闭逻辑 * 代码优化 * bugfix:修复配置重载问题 * bugfix:修复插件加载启动竞态问题 * 收敛事件入口和权限上下文 * 优化 Windows launcher 关闭重启兜底 --------- Co-authored-by: HibiKier <775757368@qq.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
HibiKier
pre-commit-ci[bot]
parent
24c316cd2c
commit
5d92ccd3b0
@@ -1,10 +1,9 @@
|
||||
from nonebot.adapters import Bot
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.services.cache import CacheRoot
|
||||
from zhenxun.services.cache.runtime_cache import GroupMemoryCache
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.enum import BlockType, CacheType
|
||||
from zhenxun.utils.enum import BlockType
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
from .strategy import get_strategy
|
||||
@@ -134,7 +133,6 @@ class PluginManager:
|
||||
await GroupConsole.bulk_update(
|
||||
update_list, [norm_field, su_field], batch_size=500
|
||||
)
|
||||
await CacheRoot.clear(CacheType.GROUPS)
|
||||
for group in update_list:
|
||||
await GroupMemoryCache.upsert_from_model(group)
|
||||
|
||||
@@ -318,7 +316,6 @@ class PluginManager:
|
||||
status=False
|
||||
)
|
||||
|
||||
await CacheRoot.clear(CacheType.GROUPS)
|
||||
await GroupMemoryCache.refresh()
|
||||
|
||||
action_str = "醒来" if status else "休眠"
|
||||
|
||||
@@ -4,12 +4,11 @@ from typing import Any, cast
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.task_info import TaskInfo
|
||||
from zhenxun.services.cache import CacheRoot
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
PluginInfoMemoryCache,
|
||||
TaskInfoMemoryCache,
|
||||
)
|
||||
from zhenxun.utils.enum import BlockType, CacheType, PluginType
|
||||
from zhenxun.utils.enum import BlockType, PluginType
|
||||
|
||||
|
||||
class SwitchStrategy(ABC):
|
||||
@@ -135,7 +134,6 @@ class PluginStrategy(SwitchStrategy):
|
||||
await self.refresh_cache()
|
||||
|
||||
async def refresh_cache(self) -> None:
|
||||
await CacheRoot.invalidate_cache(CacheType.PLUGINS)
|
||||
await PluginInfoMemoryCache.refresh()
|
||||
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
|
||||
from zhenxun.models.chat_history import ChatHistory
|
||||
from zhenxun.services.db_context import with_db_timeout
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import is_overloaded, should_pause_tasks
|
||||
from zhenxun.utils.enum import PluginType
|
||||
@@ -47,6 +48,9 @@ _HISTORY_QUEUE: asyncio.Queue[ChatHistory] = asyncio.Queue(maxsize=5000)
|
||||
_DROP_COUNT = 0
|
||||
_LAST_DROP_LOG = 0.0
|
||||
_DROP_LOG_INTERVAL = 10.0
|
||||
_FLUSH_BATCH_SIZE = 200
|
||||
_FLUSH_MAX_PER_TICK = 1000
|
||||
_FLUSH_DB_TIMEOUT = 5.0
|
||||
|
||||
|
||||
@chat_history.handle()
|
||||
@@ -80,19 +84,34 @@ async def _(message: UniMsg, session: Uninfo):
|
||||
@scheduler.scheduled_job(
|
||||
"interval",
|
||||
minutes=1,
|
||||
max_instances=1,
|
||||
coalesce=True,
|
||||
)
|
||||
async def _():
|
||||
try:
|
||||
if should_pause_tasks():
|
||||
return
|
||||
message_list: list[ChatHistory] = []
|
||||
while True:
|
||||
try:
|
||||
message_list.append(_HISTORY_QUEUE.get_nowait())
|
||||
except asyncio.QueueEmpty:
|
||||
flushed = 0
|
||||
while flushed < _FLUSH_MAX_PER_TICK:
|
||||
message_list: list[ChatHistory] = []
|
||||
limit = min(_FLUSH_BATCH_SIZE, _FLUSH_MAX_PER_TICK - flushed)
|
||||
for _ in range(limit):
|
||||
try:
|
||||
message_list.append(_HISTORY_QUEUE.get_nowait())
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
if not message_list:
|
||||
break
|
||||
if message_list:
|
||||
await ChatHistory.bulk_create(message_list)
|
||||
logger.debug(f"批量添加聊天记录 {len(message_list)} 条", "定时任务")
|
||||
await with_db_timeout(
|
||||
ChatHistory.bulk_create(message_list, _FLUSH_BATCH_SIZE),
|
||||
timeout=_FLUSH_DB_TIMEOUT,
|
||||
operation=f"ChatHistory.bulk_create[{len(message_list)}]",
|
||||
source="chat_history",
|
||||
)
|
||||
flushed += len(message_list)
|
||||
if flushed:
|
||||
backlog = _HISTORY_QUEUE.qsize()
|
||||
suffix = f",剩余队列 {backlog} 条" if backlog else ""
|
||||
logger.debug(f"批量添加聊天记录 {flushed} 条{suffix}", "定时任务")
|
||||
except Exception as e:
|
||||
logger.warning("存储聊天记录失败", "chat_history", e=e)
|
||||
|
||||
@@ -7,9 +7,10 @@ from zhenxun.models.level_user import LevelUser
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.cache.runtime_cache import LevelUserMemoryCache, LevelUserSnapshot
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
@@ -20,6 +21,9 @@ async def auth_admin(
|
||||
LevelUser | LevelUserSnapshot | None, LevelUser | LevelUserSnapshot | None
|
||||
]
|
||||
| None = None,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
entity: EntityIDs | None = None,
|
||||
):
|
||||
"""管理员命令 个人权限
|
||||
|
||||
@@ -33,7 +37,12 @@ async def auth_admin(
|
||||
return
|
||||
|
||||
try:
|
||||
entity = get_entity_ids(session)
|
||||
if context is not None:
|
||||
entity = context.entity
|
||||
if cached_levels is None:
|
||||
cached_levels = context.admin_levels
|
||||
if entity is None:
|
||||
entity = get_entity_ids(session)
|
||||
|
||||
global_user: LevelUser | LevelUserSnapshot | None = None
|
||||
group_users: LevelUser | LevelUserSnapshot | None = None
|
||||
@@ -42,7 +51,7 @@ async def auth_admin(
|
||||
global_user, group_users = cached_levels
|
||||
else:
|
||||
global_user, group_users = await LevelUserMemoryCache.get_levels(
|
||||
session.user.id, entity.group_id
|
||||
entity.user_id, entity.group_id
|
||||
)
|
||||
|
||||
user_level = global_user.user_level if global_user else 0
|
||||
@@ -53,7 +62,7 @@ async def auth_admin(
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 管理员权限不足...",
|
||||
tip_message=[
|
||||
At(flag="user", target=session.user.id),
|
||||
At(flag="user", target=entity.user_id),
|
||||
f"你的权限不足喔,该功能需要的权限等级: {plugin.admin_level}",
|
||||
],
|
||||
tip_check_tag=entity.user_id,
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from nonebot.matcher import Matcher
|
||||
@@ -8,7 +7,6 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.ban_console import BanConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.cache.cache_containers import CacheDict
|
||||
from zhenxun.services.cache.runtime_cache import BanMemoryCache
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
@@ -16,6 +14,7 @@ from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .exception import SkipPluginException
|
||||
from .utils import freq
|
||||
|
||||
@@ -25,37 +24,6 @@ Config.add_plugin_config(
|
||||
"才不会给你发消息.",
|
||||
help="对被ban用户发送的消息",
|
||||
)
|
||||
BAN_CACHE_TTL = 2
|
||||
BAN_CACHE_TTL_POSITIVE = 30
|
||||
BAN_CACHE_TTL_NEGATIVE = 5
|
||||
|
||||
BAN_CACHE = (
|
||||
CacheDict("AUTH_BAN_CACHE", expire=0)
|
||||
if max(BAN_CACHE_TTL_POSITIVE, BAN_CACHE_TTL_NEGATIVE) > 0
|
||||
else None
|
||||
)
|
||||
|
||||
|
||||
def _ban_cache_key(user_id: str | None, group_id: str | None) -> str:
|
||||
return f"{user_id or ''}:{group_id or ''}"
|
||||
|
||||
|
||||
def _ban_cache_get(key: str) -> int | None:
|
||||
if not BAN_CACHE:
|
||||
return None
|
||||
try:
|
||||
return BAN_CACHE[key]
|
||||
except KeyError:
|
||||
return None
|
||||
|
||||
|
||||
def _ban_cache_set(key: str, value: int) -> None:
|
||||
if not BAN_CACHE:
|
||||
return
|
||||
ttl = BAN_CACHE_TTL_POSITIVE if value else BAN_CACHE_TTL_NEGATIVE
|
||||
if ttl <= 0:
|
||||
return
|
||||
BAN_CACHE.set(key, value, expire=ttl)
|
||||
|
||||
|
||||
async def calculate_ban_time(ban_record: BanConsole | None) -> int:
|
||||
@@ -214,6 +182,7 @@ async def auth_ban(
|
||||
session: Uninfo,
|
||||
plugin: PluginInfo,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
entity: EntityIDs | None = None,
|
||||
is_superuser: bool = False,
|
||||
) -> None:
|
||||
@@ -229,28 +198,18 @@ async def auth_ban(
|
||||
return
|
||||
if not matcher.plugin_name:
|
||||
return
|
||||
if context is not None:
|
||||
entity = context.entity
|
||||
is_superuser = context.is_superuser
|
||||
if entity is None:
|
||||
entity = get_entity_ids(session)
|
||||
if is_superuser:
|
||||
return
|
||||
if entity.group_id:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"群组ban检查超时: {entity.group_id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
await group_handle(entity.group_id)
|
||||
|
||||
if entity.user_id:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_handle(plugin, entity, session),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
await user_handle(plugin, entity, session)
|
||||
finally:
|
||||
# 记录总执行时间
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
@@ -7,6 +7,7 @@ from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
@@ -16,6 +17,8 @@ async def auth_bot(
|
||||
bot_data: BotConsole | BotSnapshot | None = None,
|
||||
skip_fetch: bool = False,
|
||||
allow_sleep_bypass: bool = False,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
):
|
||||
"""bot层面的权限检查
|
||||
|
||||
@@ -30,6 +33,9 @@ async def auth_bot(
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
if context is not None:
|
||||
bot_id = context.event.bot_id
|
||||
bot_data = context.bot_data
|
||||
bot: BotConsole | BotSnapshot | None = bot_data
|
||||
if bot is None and not skip_fetch:
|
||||
bot = await BotMemoryCache.get(bot_id)
|
||||
|
||||
@@ -7,13 +7,18 @@ from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .exception import SkipPluginException
|
||||
|
||||
DEFAULT_GOLD = 100
|
||||
|
||||
|
||||
async def auth_cost(
|
||||
user: UserConsole | None, plugin: PluginInfo, session: Uninfo
|
||||
user: UserConsole | None,
|
||||
plugin: PluginInfo,
|
||||
session: Uninfo,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
) -> int:
|
||||
"""检测是否满足金币条件
|
||||
|
||||
@@ -28,6 +33,8 @@ async def auth_cost(
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
if context is not None and user is None:
|
||||
user = context.user
|
||||
user_gold = user.gold if user else DEFAULT_GOLD
|
||||
if user_gold < plugin.cost_gold:
|
||||
"""插件消耗金币不足"""
|
||||
|
||||
@@ -7,6 +7,7 @@ from zhenxun.services.cache.runtime_cache import GroupSnapshot
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
|
||||
from .context import PermissionContext
|
||||
from .exception import SkipPluginException
|
||||
|
||||
_GROUP_WAKE_PATTERN = re.compile(r"^醒来$", re.IGNORECASE)
|
||||
@@ -34,6 +35,8 @@ async def auth_group(
|
||||
group: GroupConsole | GroupSnapshot | None,
|
||||
text: str | None,
|
||||
group_id: str | None,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
):
|
||||
"""群黑名单检测 群总开关检测
|
||||
|
||||
@@ -42,6 +45,11 @@ async def auth_group(
|
||||
group: GroupConsole
|
||||
message: UniMsg
|
||||
"""
|
||||
if context is not None:
|
||||
group = context.group or group
|
||||
text = context.plain_text
|
||||
group_id = context.group_id
|
||||
|
||||
if not group_id:
|
||||
return
|
||||
|
||||
|
||||
@@ -19,9 +19,10 @@ from zhenxun.utils.limiters import CountLimiter, FreqLimiter, UserBlockLimiter
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
from zhenxun.utils.time_utils import TimeUtils
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .exception import SkipPluginException
|
||||
|
||||
driver = nonebot.get_driver()
|
||||
@@ -31,7 +32,7 @@ _LIMIT_NOTICE_LIMITER = FreqLimiter(_LIMIT_NOTICE_CD)
|
||||
_LIMIT_NOTICE_TASKS: set[asyncio.Task] = set()
|
||||
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=5)
|
||||
@PriorityLifecycle.on_startup(priority=7)
|
||||
async def _():
|
||||
"""初始化限制"""
|
||||
await LimitManager.init_limit()
|
||||
@@ -117,6 +118,7 @@ class LimitManager:
|
||||
cls.cd_limit = {}
|
||||
cls.block_limit = {}
|
||||
cls.count_limit = {}
|
||||
cls.module_limit_cache.clear()
|
||||
# 添加新数据
|
||||
for limit in limit_list:
|
||||
cls.add_limit(limit)
|
||||
@@ -137,22 +139,22 @@ class LimitManager:
|
||||
"""
|
||||
if limit.module not in cls.add_module:
|
||||
cls.add_module.append(limit.module)
|
||||
if limit.limit_type == PluginLimitType.BLOCK:
|
||||
cls.block_limit[limit.module] = Limit(
|
||||
limit=limit, limiter=UserBlockLimiter()
|
||||
)
|
||||
elif limit.limit_type == PluginLimitType.CD:
|
||||
cd_value = int(limit.cd or 0)
|
||||
cls.cd_limit[limit.module] = Limit(
|
||||
limit=limit, limiter=FreqLimiter(cd_value)
|
||||
)
|
||||
elif limit.limit_type == PluginLimitType.COUNT:
|
||||
max_count = int(limit.max_count or 0)
|
||||
if max_count <= 0:
|
||||
return
|
||||
cls.count_limit[limit.module] = Limit(
|
||||
limit=limit, limiter=CountLimiter(max_count)
|
||||
)
|
||||
if limit.limit_type == PluginLimitType.BLOCK:
|
||||
cls.block_limit[limit.module] = Limit(
|
||||
limit=limit, limiter=UserBlockLimiter()
|
||||
)
|
||||
elif limit.limit_type == PluginLimitType.CD:
|
||||
cd_value = int(limit.cd or 0)
|
||||
cls.cd_limit[limit.module] = Limit(
|
||||
limit=limit, limiter=FreqLimiter(cd_value)
|
||||
)
|
||||
elif limit.limit_type == PluginLimitType.COUNT:
|
||||
max_count = int(limit.max_count or 0)
|
||||
if max_count <= 0:
|
||||
return
|
||||
cls.count_limit[limit.module] = Limit(
|
||||
limit=limit, limiter=CountLimiter(max_count)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def unblock(
|
||||
@@ -322,14 +324,23 @@ class LimitManager:
|
||||
limiter.increase(key_type)
|
||||
|
||||
|
||||
async def auth_limit(plugin: PluginInfo, session: Uninfo):
|
||||
async def auth_limit(
|
||||
plugin: PluginInfo,
|
||||
session: Uninfo,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
entity: EntityIDs | None = None,
|
||||
):
|
||||
"""插件限制
|
||||
|
||||
参数:
|
||||
plugin: PluginInfo
|
||||
session: Uninfo
|
||||
"""
|
||||
entity = get_entity_ids(session)
|
||||
if context is not None:
|
||||
entity = context.entity
|
||||
if entity is None:
|
||||
entity = get_entity_ids(session)
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
LimitManager.check(
|
||||
|
||||
@@ -10,6 +10,7 @@ from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import BlockType
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .exception import IsSuperuserException, SkipPluginException
|
||||
from .utils import freq, is_poke
|
||||
|
||||
@@ -107,11 +108,16 @@ class GroupCheck:
|
||||
|
||||
class PluginCheck:
|
||||
def __init__(
|
||||
self, group: GroupConsole | GroupSnapshot | None, session: Uninfo, is_poke: bool
|
||||
self,
|
||||
group: GroupConsole | GroupSnapshot | None,
|
||||
session: Uninfo,
|
||||
is_poke: bool,
|
||||
user_id: str | None,
|
||||
):
|
||||
self.session = session
|
||||
self.is_poke = is_poke
|
||||
self.group_data = group
|
||||
self.user_id = user_id or session.user.id
|
||||
self.group_id = None
|
||||
if group:
|
||||
self.group_id = group.group_id
|
||||
@@ -126,13 +132,11 @@ class PluginCheck:
|
||||
IgnoredException: 忽略插件
|
||||
"""
|
||||
if plugin.block_type == BlockType.PRIVATE:
|
||||
should_tip = freq.is_send_limit_message(
|
||||
plugin, self.session.user.id, self.is_poke
|
||||
)
|
||||
should_tip = freq.is_send_limit_message(plugin, self.user_id, self.is_poke)
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 该插件在私聊中已被禁用...",
|
||||
tip_message="该功能在私聊中已被禁用..." if should_tip else None,
|
||||
tip_check_tag=self.session.user.id if should_tip else None,
|
||||
tip_check_tag=self.user_id if should_tip else None,
|
||||
tip_background=should_tip,
|
||||
)
|
||||
|
||||
@@ -153,7 +157,7 @@ class PluginCheck:
|
||||
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.user_id
|
||||
should_tip = freq.is_send_limit_message(plugin, sid, self.is_poke)
|
||||
raise SkipPluginException(
|
||||
f"{plugin.name}({plugin.module}) 全局未开启此功能...",
|
||||
@@ -176,7 +180,9 @@ async def auth_plugin(
|
||||
session: Uninfo,
|
||||
event: Event,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
skip_group_block: bool = False,
|
||||
user_id: str | None = None,
|
||||
):
|
||||
"""插件状态
|
||||
|
||||
@@ -187,8 +193,11 @@ async def auth_plugin(
|
||||
"""
|
||||
start_time = time.time()
|
||||
try:
|
||||
if context is not None:
|
||||
group = context.group or group
|
||||
user_id = context.user_id
|
||||
is_poke_event = is_poke(event)
|
||||
user_check = PluginCheck(group, session, is_poke_event)
|
||||
user_check = PluginCheck(group, session, is_poke_event, user_id)
|
||||
|
||||
if group:
|
||||
block_set, super_block_set = _get_group_block_sets(group)
|
||||
|
||||
@@ -3,6 +3,7 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
|
||||
from .context import PermissionContext
|
||||
from .exception import SkipPluginException
|
||||
|
||||
Config.add_plugin_config(
|
||||
@@ -15,7 +16,12 @@ Config.add_plugin_config(
|
||||
)
|
||||
|
||||
|
||||
def bot_filter(session: Uninfo):
|
||||
def bot_filter(
|
||||
session: Uninfo,
|
||||
*,
|
||||
context: PermissionContext | None = None,
|
||||
user_id: str | None = None,
|
||||
):
|
||||
"""过滤bot调用bot
|
||||
|
||||
参数:
|
||||
@@ -26,10 +32,13 @@ def bot_filter(session: Uninfo):
|
||||
"""
|
||||
if not Config.get_config("hook", "FILTER_BOT"):
|
||||
return
|
||||
if context is not None:
|
||||
user_id = context.user_id
|
||||
bot_ids = list(nonebot.get_bots().keys())
|
||||
if session.user.id == session.self_id:
|
||||
checked_user_id = user_id or session.user.id
|
||||
if checked_user_id == session.self_id:
|
||||
return
|
||||
if session.user.id in bot_ids:
|
||||
if checked_user_id in bot_ids:
|
||||
raise SkipPluginException(
|
||||
f"bot:{session.self_id} 尝试调用 bot:{session.user.id}"
|
||||
f"bot:{session.self_id} 尝试调用 bot:{checked_user_id}"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,321 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from nonebot.adapters import Bot, Event
|
||||
from nonebot_plugin_alconna import UniMsg
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.services.cache.cache_containers import CacheDict
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
AUTH_EVENT_CACHE_TTL = 5
|
||||
|
||||
STATE_EVENT_CONTEXT = "_zx_event_context"
|
||||
STATE_PERMISSION_CONTEXT = "_zx_permission_context"
|
||||
STATE_ENTITY = "_zx_entity"
|
||||
STATE_EVENT_CACHE = "_zx_event_cache"
|
||||
STATE_PLAIN_TEXT = "_zx_plain_text"
|
||||
STATE_ROUTE_MODULES = "_zx_route_modules"
|
||||
STATE_IS_SUPERUSER = "_zx_is_superuser"
|
||||
STATE_PERMISSION_SIDE_EFFECTS = "_zx_permission_side_effects"
|
||||
EVENT_CACHE_PERMISSION_SIDE_EFFECTS = "permission_side_effects"
|
||||
|
||||
EVENT_CACHE = (
|
||||
CacheDict("AUTH_EVENT_CACHE", expire=AUTH_EVENT_CACHE_TTL)
|
||||
if AUTH_EVENT_CACHE_TTL > 0
|
||||
else None
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class EventContext:
|
||||
bot_id: str
|
||||
platform: str
|
||||
event_type: str
|
||||
message_id: str | int | None
|
||||
entity: EntityIDs
|
||||
plain_text: str = ""
|
||||
route_modules: set[str] = field(default_factory=set)
|
||||
route_modules_loaded: bool = False
|
||||
is_superuser: bool = False
|
||||
event_cache: dict[str, Any] | None = None
|
||||
|
||||
@property
|
||||
def user_id(self) -> str:
|
||||
return self.entity.user_id
|
||||
|
||||
@property
|
||||
def group_id(self) -> str | None:
|
||||
return self.entity.group_id
|
||||
|
||||
@property
|
||||
def channel_id(self) -> str | None:
|
||||
return self.entity.channel_id
|
||||
|
||||
|
||||
@dataclass
|
||||
class PermissionSideEffectCache:
|
||||
auth_results: dict[str, tuple[bool, str | None]] = field(default_factory=dict)
|
||||
module_locks: dict[str, asyncio.Lock] = field(default_factory=dict)
|
||||
|
||||
def lock_for(self, module: str) -> asyncio.Lock:
|
||||
lock = self.module_locks.get(module)
|
||||
if lock is None:
|
||||
lock = asyncio.Lock()
|
||||
self.module_locks[module] = lock
|
||||
return lock
|
||||
|
||||
|
||||
@dataclass
|
||||
class PermissionContext:
|
||||
event: EventContext
|
||||
module: str
|
||||
plugin: Any = None
|
||||
user: Any = None
|
||||
group: Any = None
|
||||
bot_data: Any = None
|
||||
admin_levels: Any = None
|
||||
|
||||
@property
|
||||
def entity(self) -> EntityIDs:
|
||||
return self.event.entity
|
||||
|
||||
@property
|
||||
def user_id(self) -> str:
|
||||
return self.event.user_id
|
||||
|
||||
@property
|
||||
def group_id(self) -> str | None:
|
||||
return self.event.group_id
|
||||
|
||||
@property
|
||||
def channel_id(self) -> str | None:
|
||||
return self.event.channel_id
|
||||
|
||||
@property
|
||||
def plain_text(self) -> str:
|
||||
return self.event.plain_text
|
||||
|
||||
@property
|
||||
def is_superuser(self) -> bool:
|
||||
return self.event.is_superuser
|
||||
|
||||
|
||||
def resolve_actor_user_id(event: Event, fallback_user_id: str | None) -> str:
|
||||
"""优先使用事件发起者 ID,避免 notice 场景 session.user 指向 bot 自身。"""
|
||||
event_user_id = getattr(event, "user_id", None)
|
||||
if event_user_id is None:
|
||||
return fallback_user_id or ""
|
||||
resolved = str(event_user_id)
|
||||
return resolved or fallback_user_id or ""
|
||||
|
||||
|
||||
def resolve_event_group_id(event: Event, fallback_group_id: str | None) -> str | None:
|
||||
"""notice 场景 session.group 可能缺失,回退到事件上的 group_id。"""
|
||||
event_group_id = getattr(event, "group_id", None)
|
||||
if event_group_id is None:
|
||||
return fallback_group_id
|
||||
resolved = str(event_group_id)
|
||||
return resolved or fallback_group_id
|
||||
|
||||
|
||||
def resolve_event_channel_id(
|
||||
event: Event, fallback_channel_id: str | None
|
||||
) -> str | None:
|
||||
"""频道场景回退到事件上的 channel_id。"""
|
||||
event_channel_id = getattr(event, "channel_id", None)
|
||||
if event_channel_id is None:
|
||||
return fallback_channel_id
|
||||
resolved = str(event_channel_id)
|
||||
return resolved or fallback_channel_id
|
||||
|
||||
|
||||
def resolve_entity_ids(event: Event, session: Uninfo) -> EntityIDs:
|
||||
entity = get_entity_ids(session)
|
||||
entity.user_id = resolve_actor_user_id(event, entity.user_id)
|
||||
entity.group_id = resolve_event_group_id(event, entity.group_id)
|
||||
entity.channel_id = resolve_event_channel_id(event, entity.channel_id)
|
||||
return entity
|
||||
|
||||
|
||||
def extract_plain_text(message: UniMsg | None, event: Event) -> str:
|
||||
if message is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
return message.extract_plain_text()
|
||||
with contextlib.suppress(Exception):
|
||||
plain = event.get_plaintext()
|
||||
if plain:
|
||||
return plain.strip()
|
||||
return ""
|
||||
|
||||
|
||||
def _event_message_id(event: Event) -> str | int | None:
|
||||
msg_id = getattr(event, "message_id", None)
|
||||
if msg_id is None:
|
||||
msg_id = getattr(event, "id", None)
|
||||
return msg_id
|
||||
|
||||
|
||||
def event_cache_key(
|
||||
event: Event,
|
||||
*,
|
||||
bot_id: str,
|
||||
platform: str,
|
||||
entity: EntityIDs,
|
||||
) -> str:
|
||||
msg_id = _event_message_id(event)
|
||||
if msg_id is None:
|
||||
msg_id = id(event)
|
||||
group_id = entity.group_id or ""
|
||||
channel_id = entity.channel_id or ""
|
||||
return f"{platform}:{bot_id}:{entity.user_id}:{group_id}:{channel_id}:{msg_id}"
|
||||
|
||||
|
||||
def get_event_cache(
|
||||
event: Event,
|
||||
*,
|
||||
bot_id: str,
|
||||
platform: str,
|
||||
entity: EntityIDs,
|
||||
) -> dict[str, Any] | None:
|
||||
if not EVENT_CACHE:
|
||||
return None
|
||||
key = event_cache_key(event, bot_id=bot_id, platform=platform, entity=entity)
|
||||
try:
|
||||
return EVENT_CACHE[key]
|
||||
except KeyError:
|
||||
cache: dict[str, Any] = {}
|
||||
EVENT_CACHE[key] = cache
|
||||
return cache
|
||||
|
||||
|
||||
def _sync_context_state(state: dict[str, Any], context: EventContext) -> None:
|
||||
state[STATE_EVENT_CONTEXT] = context
|
||||
state[STATE_ENTITY] = context.entity
|
||||
state[STATE_EVENT_CACHE] = context.event_cache
|
||||
state[STATE_PLAIN_TEXT] = context.plain_text
|
||||
state[STATE_ROUTE_MODULES] = context.route_modules
|
||||
state[STATE_IS_SUPERUSER] = context.is_superuser
|
||||
get_permission_side_effect_cache(state=state, event_cache=context.event_cache)
|
||||
|
||||
|
||||
def get_permission_side_effect_cache(
|
||||
*,
|
||||
state: dict[str, Any] | None = None,
|
||||
event_cache: dict[str, Any] | None = None,
|
||||
) -> PermissionSideEffectCache:
|
||||
side_effects = None
|
||||
if state is not None:
|
||||
side_effects = state.get(STATE_PERMISSION_SIDE_EFFECTS)
|
||||
if (
|
||||
not isinstance(side_effects, PermissionSideEffectCache)
|
||||
and event_cache is not None
|
||||
):
|
||||
side_effects = event_cache.get(EVENT_CACHE_PERMISSION_SIDE_EFFECTS)
|
||||
if not isinstance(side_effects, PermissionSideEffectCache):
|
||||
side_effects = PermissionSideEffectCache()
|
||||
if state is not None:
|
||||
state[STATE_PERMISSION_SIDE_EFFECTS] = side_effects
|
||||
if event_cache is not None:
|
||||
event_cache[EVENT_CACHE_PERMISSION_SIDE_EFFECTS] = side_effects
|
||||
return side_effects
|
||||
|
||||
|
||||
def get_event_context(state: dict[str, Any] | None) -> EventContext | None:
|
||||
if state is None:
|
||||
return None
|
||||
context = state.get(STATE_EVENT_CONTEXT)
|
||||
return context if isinstance(context, EventContext) else None
|
||||
|
||||
|
||||
def get_or_create_event_context(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
state: dict[str, Any],
|
||||
*,
|
||||
message: UniMsg | None = None,
|
||||
) -> EventContext:
|
||||
context = get_event_context(state)
|
||||
if context is not None:
|
||||
_sync_context_state(state, context)
|
||||
return context
|
||||
|
||||
entity = state.get(STATE_ENTITY)
|
||||
if not isinstance(entity, EntityIDs):
|
||||
entity = resolve_entity_ids(event, session)
|
||||
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
bot_id = str(bot.self_id)
|
||||
event_cache = state.get(STATE_EVENT_CACHE)
|
||||
if not isinstance(event_cache, dict):
|
||||
event_cache = get_event_cache(
|
||||
event,
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
entity=entity,
|
||||
)
|
||||
|
||||
text = state.get(STATE_PLAIN_TEXT)
|
||||
if not isinstance(text, str):
|
||||
cached_text = event_cache.get("plain_text") if event_cache is not None else None
|
||||
text = (
|
||||
cached_text
|
||||
if isinstance(cached_text, str)
|
||||
else extract_plain_text(message, event)
|
||||
)
|
||||
if event_cache is not None:
|
||||
event_cache["plain_text"] = text
|
||||
|
||||
route_modules_loaded = STATE_ROUTE_MODULES in state
|
||||
route_modules = state.get(STATE_ROUTE_MODULES)
|
||||
if not isinstance(route_modules, set):
|
||||
cached_routes = (
|
||||
event_cache.get("route_modules") if event_cache is not None else None
|
||||
)
|
||||
route_modules = cached_routes if isinstance(cached_routes, set) else set()
|
||||
route_modules_loaded = isinstance(cached_routes, set)
|
||||
|
||||
is_superuser = state.get(STATE_IS_SUPERUSER)
|
||||
if not isinstance(is_superuser, bool):
|
||||
is_superuser = entity.user_id in bot.config.superusers
|
||||
|
||||
context = EventContext(
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
event_type=event.get_type(),
|
||||
message_id=_event_message_id(event),
|
||||
entity=entity,
|
||||
plain_text=text,
|
||||
route_modules=route_modules,
|
||||
route_modules_loaded=route_modules_loaded,
|
||||
is_superuser=is_superuser,
|
||||
event_cache=event_cache,
|
||||
)
|
||||
_sync_context_state(state, context)
|
||||
return context
|
||||
|
||||
|
||||
def set_route_modules(
|
||||
state: dict[str, Any] | None,
|
||||
context: EventContext,
|
||||
route_modules: set[str],
|
||||
) -> None:
|
||||
context.route_modules = route_modules
|
||||
context.route_modules_loaded = True
|
||||
if context.event_cache is not None:
|
||||
context.event_cache["route_modules"] = route_modules
|
||||
if state is not None:
|
||||
_sync_context_state(state, context)
|
||||
|
||||
|
||||
def store_permission_context(
|
||||
state: dict[str, Any] | None, context: PermissionContext
|
||||
) -> None:
|
||||
if state is not None:
|
||||
state[STATE_PERMISSION_CONTEXT] = context
|
||||
@@ -1,7 +1,7 @@
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable
|
||||
import contextlib
|
||||
import os
|
||||
import importlib
|
||||
import re
|
||||
import time
|
||||
from typing import cast
|
||||
@@ -11,7 +11,6 @@ from nonebot.adapters import Bot, Event
|
||||
from nonebot.exception import IgnoredException
|
||||
from nonebot.matcher import Matcher
|
||||
import nonebot.message as nb_message
|
||||
from nonebot_plugin_alconna import UniMsg
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
@@ -33,7 +32,6 @@ from zhenxun.services.message_load import is_overloaded
|
||||
from zhenxun.utils.enum import BlockType, GoldHandle, PluginType
|
||||
from zhenxun.utils.exception import InsufficientGold
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
from .auth.auth_admin import auth_admin
|
||||
from .auth.auth_ban import auth_ban
|
||||
@@ -44,6 +42,16 @@ from .auth.auth_limit import LimitManager, auth_limit
|
||||
from .auth.auth_plugin import auth_plugin
|
||||
from .auth.bot_filter import bot_filter
|
||||
from .auth.config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .auth.context import (
|
||||
EVENT_CACHE,
|
||||
STATE_PLAIN_TEXT,
|
||||
EventContext,
|
||||
PermissionContext,
|
||||
get_event_context,
|
||||
get_permission_side_effect_cache,
|
||||
set_route_modules,
|
||||
store_permission_context,
|
||||
)
|
||||
from .auth.exception import (
|
||||
IsSuperuserException,
|
||||
PermissionExemption,
|
||||
@@ -53,7 +61,6 @@ from .auth.utils import send_message
|
||||
|
||||
AUTH_HOOKS_CONCURRENCY_LIMIT = 5
|
||||
AUTH_DB_CONCURRENCY_LIMIT = 6
|
||||
AUTH_EVENT_CACHE_TTL = 5 # 增加到5秒,减少缓存抖动
|
||||
|
||||
|
||||
# 超时设置(秒)
|
||||
@@ -74,13 +81,6 @@ CIRCUIT_RESET_TIME = 300 # 5分钟
|
||||
HOOKS_CONCURRENCY_LIMIT = AUTH_HOOKS_CONCURRENCY_LIMIT
|
||||
DB_CONCURRENCY_LIMIT = AUTH_DB_CONCURRENCY_LIMIT
|
||||
|
||||
EVENT_CACHE_TTL = AUTH_EVENT_CACHE_TTL
|
||||
EVENT_CACHE = (
|
||||
CacheDict("AUTH_EVENT_CACHE", expire=EVENT_CACHE_TTL)
|
||||
if EVENT_CACHE_TTL > 0
|
||||
else None
|
||||
)
|
||||
|
||||
# 路由索引缓存
|
||||
_ROUTE_INDEX_LOCK = asyncio.Lock()
|
||||
_ROUTE_INDEX_READY = False
|
||||
@@ -91,21 +91,17 @@ MATCHER_ROUTE_PREFILTER_TTL = 2
|
||||
PREFILTER_STATS_LOG_INTERVAL = 10.0
|
||||
CACHE_SWEEP_INTERVAL = 1.0
|
||||
|
||||
CPU_COUNT = os.cpu_count() or 4
|
||||
COMMAND_MATCHER_CONCURRENCY = max(8, min(48, CPU_COUNT * 4))
|
||||
HEAVY_COMMAND_CONCURRENCY = max(1, min(3, CPU_COUNT // 2))
|
||||
HEAVY_COMMAND_MODULES = frozenset({"shop", "sign_in"})
|
||||
|
||||
# 全局信号量与计数器
|
||||
HOOKS_ACTIVE_COUNT = 0
|
||||
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT)
|
||||
COMMAND_MATCHER_SEMAPHORE = asyncio.Semaphore(COMMAND_MATCHER_CONCURRENCY)
|
||||
HEAVY_COMMAND_SEMAPHORE = asyncio.Semaphore(HEAVY_COMMAND_CONCURRENCY)
|
||||
|
||||
DB_SEMAPHORE = asyncio.Semaphore(DB_CONCURRENCY_LIMIT)
|
||||
DB_ACTIVE_COUNT = 0
|
||||
_CHECK_MATCHER_PATCHED = False
|
||||
_ORIGINAL_CHECK_AND_RUN_MATCHER: Callable[..., Awaitable[None]] | None = None
|
||||
_HANDLE_EVENT_PATCHED = False
|
||||
_ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None
|
||||
_ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {}
|
||||
_MATCHER_COMMAND_TYPE_CACHE: dict[type[Matcher], bool] = {}
|
||||
_MATCHER_COMMAND_LITERAL_CACHE: dict[type[Matcher], tuple[str, ...] | None] = {}
|
||||
_MATCHER_ALCONNA_SHORTCUT_CACHE: dict[type[Matcher], bool] = {}
|
||||
@@ -115,6 +111,10 @@ _CHECK_MATCHER_ROUTE_CACHE = CacheDict(
|
||||
_PREFILTER_STATS = {
|
||||
"checked": 0,
|
||||
"skipped": 0,
|
||||
"before_task_checked": 0,
|
||||
"before_task_skipped": 0,
|
||||
"inside_task_checked": 0,
|
||||
"inside_task_skipped": 0,
|
||||
"type_miss": 0,
|
||||
"route_miss": 0,
|
||||
"command_miss": 0,
|
||||
@@ -163,33 +163,6 @@ def _debug_log(message: str, *args, **kwargs) -> None:
|
||||
logger.debug(message, *args, **kwargs)
|
||||
|
||||
|
||||
def _event_cache_key(event: Event, session: Uninfo, entity) -> str:
|
||||
msg_id = getattr(event, "message_id", None)
|
||||
if msg_id is None:
|
||||
msg_id = getattr(event, "id", None)
|
||||
if msg_id is None:
|
||||
msg_id = id(event)
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
group_id = entity.group_id or ""
|
||||
channel_id = entity.channel_id or ""
|
||||
return (
|
||||
f"{platform}:{session.self_id}:{entity.user_id}:"
|
||||
f"{group_id}:{channel_id}:{msg_id}"
|
||||
)
|
||||
|
||||
|
||||
def _get_event_cache(event: Event, session: Uninfo, entity):
|
||||
if not EVENT_CACHE:
|
||||
return None
|
||||
key = _event_cache_key(event, session, entity)
|
||||
try:
|
||||
return EVENT_CACHE[key]
|
||||
except KeyError:
|
||||
cache = {}
|
||||
EVENT_CACHE[key] = cache
|
||||
return cache
|
||||
|
||||
|
||||
def _normalize_command(command: str) -> str:
|
||||
text = command.strip()
|
||||
if not text:
|
||||
@@ -488,6 +461,9 @@ def _event_plain_text(event: Event) -> str:
|
||||
def _state_plain_text(state: dict | None) -> str:
|
||||
if state is None:
|
||||
return ""
|
||||
context = get_event_context(state)
|
||||
if context is not None:
|
||||
return context.plain_text.strip()
|
||||
text = state.get("_zx_plain_text")
|
||||
if isinstance(text, str):
|
||||
return text.strip()
|
||||
@@ -496,6 +472,9 @@ def _state_plain_text(state: dict | None) -> str:
|
||||
|
||||
def _get_route_modules_for_event(event: Event, state: dict | None = None) -> set[str]:
|
||||
if state is not None:
|
||||
context = get_event_context(state)
|
||||
if context is not None and context.route_modules_loaded:
|
||||
return context.route_modules
|
||||
route_modules = state.get("_zx_route_modules")
|
||||
if isinstance(route_modules, set):
|
||||
return route_modules
|
||||
@@ -506,15 +485,67 @@ def _get_route_modules_for_event(event: Event, state: dict | None = None) -> set
|
||||
route_modules = _match_route_modules(_event_plain_text(event))
|
||||
_CHECK_MATCHER_ROUTE_CACHE[key] = route_modules
|
||||
if state is not None:
|
||||
state["_zx_route_modules"] = route_modules
|
||||
context = get_event_context(state)
|
||||
if context is not None:
|
||||
set_route_modules(state, context, route_modules)
|
||||
else:
|
||||
state["_zx_route_modules"] = route_modules
|
||||
return route_modules
|
||||
|
||||
|
||||
def _record_prefilter_stats(skipped: bool, reason: str | None) -> None:
|
||||
def _prepare_handle_event_state(event: Event, state: dict) -> None:
|
||||
get_permission_side_effect_cache(state=state)
|
||||
if event.get_type() != "message":
|
||||
return
|
||||
if _state_plain_text(state):
|
||||
return
|
||||
text = _event_plain_text(event)
|
||||
if text:
|
||||
state[STATE_PLAIN_TEXT] = text
|
||||
|
||||
|
||||
def _build_matcher_state(base_state: dict) -> dict:
|
||||
get_permission_side_effect_cache(state=base_state)
|
||||
matcher_state = base_state.copy()
|
||||
get_permission_side_effect_cache(state=matcher_state)
|
||||
return matcher_state
|
||||
|
||||
|
||||
async def _run_selected_matcher(
|
||||
matcher: type[Matcher],
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
state: dict,
|
||||
stack,
|
||||
dependency_cache,
|
||||
) -> None:
|
||||
await nb_message.check_and_run_matcher(
|
||||
matcher,
|
||||
bot,
|
||||
event,
|
||||
state,
|
||||
stack,
|
||||
dependency_cache,
|
||||
)
|
||||
|
||||
|
||||
def _record_prefilter_stats(
|
||||
skipped: bool,
|
||||
reason: str | None,
|
||||
stage: str = "inside_task",
|
||||
) -> None:
|
||||
global _PREFILTER_LAST_LOG
|
||||
_PREFILTER_STATS["checked"] += 1
|
||||
if skipped:
|
||||
_PREFILTER_STATS["skipped"] += 1
|
||||
if stage == "before_task":
|
||||
_PREFILTER_STATS["before_task_checked"] += 1
|
||||
if skipped:
|
||||
_PREFILTER_STATS["before_task_skipped"] += 1
|
||||
else:
|
||||
_PREFILTER_STATS["inside_task_checked"] += 1
|
||||
if skipped:
|
||||
_PREFILTER_STATS["inside_task_skipped"] += 1
|
||||
if reason == "type_miss":
|
||||
_PREFILTER_STATS["type_miss"] += 1
|
||||
elif reason == "route_miss":
|
||||
@@ -537,6 +568,10 @@ def _record_prefilter_stats(skipped: bool, reason: str | None) -> None:
|
||||
"matcher prefilter stats: "
|
||||
f"checked={_PREFILTER_STATS['checked']} "
|
||||
f"skipped={_PREFILTER_STATS['skipped']} "
|
||||
f"before_task={_PREFILTER_STATS['before_task_skipped']}/"
|
||||
f"{_PREFILTER_STATS['before_task_checked']} "
|
||||
f"inside_task={_PREFILTER_STATS['inside_task_skipped']}/"
|
||||
f"{_PREFILTER_STATS['inside_task_checked']} "
|
||||
f"type_miss={_PREFILTER_STATS['type_miss']} "
|
||||
f"route_miss={_PREFILTER_STATS['route_miss']} "
|
||||
f"command_miss={_PREFILTER_STATS['command_miss']} "
|
||||
@@ -643,15 +678,6 @@ def _matcher_has_alconna_shortcuts(matcher_cls: type[Matcher]) -> bool:
|
||||
return has_shortcuts
|
||||
|
||||
|
||||
def _is_heavy_command_module(module: str) -> bool:
|
||||
normalized = module.strip().lower()
|
||||
if not normalized:
|
||||
return False
|
||||
if normalized in HEAVY_COMMAND_MODULES:
|
||||
return True
|
||||
return any(normalized.endswith(f".{name}") for name in HEAVY_COMMAND_MODULES)
|
||||
|
||||
|
||||
async def _check_matcher_prefilter(
|
||||
matcher_cls: type[Matcher], event: Event, state: dict | None = None
|
||||
) -> tuple[bool, str | None]:
|
||||
@@ -686,6 +712,17 @@ async def _check_matcher_prefilter(
|
||||
if not module:
|
||||
return False, None
|
||||
|
||||
command_matched = False
|
||||
matcher_commands = _extract_matcher_command_literals(matcher_cls)
|
||||
if matcher_commands:
|
||||
for command in matcher_commands:
|
||||
if _command_matches(text, command):
|
||||
command_matched = True
|
||||
break
|
||||
else:
|
||||
if not _matcher_has_alconna_shortcuts(matcher_cls):
|
||||
return True, "command_miss"
|
||||
|
||||
ai_route_modules = _collect_ai_route_modules(event, state)
|
||||
ai_route_heads = _collect_ai_route_heads(event, state)
|
||||
if ai_route_modules and module not in ai_route_modules:
|
||||
@@ -696,25 +733,85 @@ async def _check_matcher_prefilter(
|
||||
await _ensure_route_index()
|
||||
|
||||
if module not in _ROUTE_MODULES_WITH_COMMANDS:
|
||||
matcher_commands = _extract_matcher_command_literals(matcher_cls)
|
||||
if matcher_commands:
|
||||
for command in matcher_commands:
|
||||
if _command_matches(text, command):
|
||||
return False, None
|
||||
if _matcher_has_alconna_shortcuts(matcher_cls):
|
||||
return False, None
|
||||
return True, "command_miss"
|
||||
return False, None
|
||||
|
||||
route_modules = _get_route_modules_for_event(event, state)
|
||||
if module not in route_modules:
|
||||
if command_matched:
|
||||
return False, None
|
||||
if _matcher_has_alconna_shortcuts(matcher_cls):
|
||||
return False, None
|
||||
return True, "route_miss"
|
||||
return False, None
|
||||
|
||||
|
||||
_MATCHER_SEMAPHORE_TIMEOUT = 8.0
|
||||
def _check_matcher_prefilter_before_task(
|
||||
matcher_cls: type[Matcher], event: Event, state: dict | None = None
|
||||
) -> tuple[bool, str | None]:
|
||||
"""Conservative selector before creating matcher task.
|
||||
|
||||
This mirrors the async matcher prefilter but never performs IO or route-index
|
||||
rebuild. If anything is uncertain, let the existing check_and_run_matcher
|
||||
patch handle it inside the task.
|
||||
"""
|
||||
event_type = event.get_type()
|
||||
matcher_type = getattr(matcher_cls, "type", "") or ""
|
||||
if isinstance(matcher_type, str) and matcher_type and matcher_type != event_type:
|
||||
return True, "type_miss"
|
||||
|
||||
if event_type != "message":
|
||||
return False, None
|
||||
|
||||
if getattr(matcher_cls, "temp", False):
|
||||
return False, None
|
||||
|
||||
if not _is_command_matcher_class(matcher_cls):
|
||||
return False, None
|
||||
|
||||
text = _state_plain_text(state)
|
||||
if not text:
|
||||
text = _event_plain_text(event)
|
||||
if state is not None and text:
|
||||
state["_zx_plain_text"] = text
|
||||
if not text:
|
||||
return True, "empty_text"
|
||||
|
||||
module = _matcher_module_name(matcher_cls)
|
||||
if not module:
|
||||
return False, None
|
||||
|
||||
command_matched = False
|
||||
matcher_commands = _extract_matcher_command_literals(matcher_cls)
|
||||
has_alconna_shortcuts = _matcher_has_alconna_shortcuts(matcher_cls)
|
||||
if matcher_commands:
|
||||
for command in matcher_commands:
|
||||
if _command_matches(text, command):
|
||||
command_matched = True
|
||||
break
|
||||
else:
|
||||
if not has_alconna_shortcuts:
|
||||
return True, "command_miss"
|
||||
|
||||
ai_route_modules = _collect_ai_route_modules(event, state)
|
||||
ai_route_heads = _collect_ai_route_heads(event, state)
|
||||
if ai_route_modules and module not in ai_route_modules:
|
||||
if not _matcher_matches_ai_route_heads(matcher_cls, ai_route_heads):
|
||||
return True, "route_miss"
|
||||
|
||||
if not _ROUTE_INDEX_READY:
|
||||
return False, None
|
||||
|
||||
if module not in _ROUTE_MODULES_WITH_COMMANDS:
|
||||
return False, None
|
||||
|
||||
route_modules = _get_route_modules_for_event(event, state)
|
||||
if module not in route_modules:
|
||||
if command_matched or has_alconna_shortcuts:
|
||||
return False, None
|
||||
return True, "route_miss"
|
||||
return False, None
|
||||
|
||||
|
||||
_MAX_MATCHER_CACHE = 512
|
||||
|
||||
|
||||
@@ -729,7 +826,7 @@ async def _patched_check_and_run_matcher(
|
||||
skip, reason = await _check_matcher_prefilter(
|
||||
Matcher, event, state if isinstance(state, dict) else None
|
||||
)
|
||||
_record_prefilter_stats(skip, reason)
|
||||
_record_prefilter_stats(skip, reason, "inside_task")
|
||||
if skip:
|
||||
return
|
||||
|
||||
@@ -744,28 +841,6 @@ async def _patched_check_and_run_matcher(
|
||||
"stack": stack,
|
||||
"dependency_cache": dependency_cache,
|
||||
}
|
||||
if _is_command_matcher_class(Matcher):
|
||||
module = _matcher_module_name(Matcher)
|
||||
sem = (
|
||||
HEAVY_COMMAND_SEMAPHORE
|
||||
if _is_heavy_command_module(module)
|
||||
else COMMAND_MATCHER_SEMAPHORE
|
||||
)
|
||||
try:
|
||||
await asyncio.wait_for(sem.acquire(), timeout=_MATCHER_SEMAPHORE_TIMEOUT)
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(
|
||||
f"matcher semaphore acquire timeout for {module}, "
|
||||
"executing without concurrency limit",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
await original(**kwargs)
|
||||
return
|
||||
try:
|
||||
await original(**kwargs)
|
||||
finally:
|
||||
sem.release()
|
||||
return
|
||||
await original(**kwargs)
|
||||
|
||||
|
||||
@@ -788,27 +863,137 @@ def _uninstall_matcher_prefilter() -> None:
|
||||
_ORIGINAL_CHECK_AND_RUN_MATCHER = None
|
||||
|
||||
|
||||
def _get_message_text(
|
||||
message: UniMsg | None,
|
||||
event_cache: dict | None,
|
||||
event: Event | None = None,
|
||||
) -> str:
|
||||
if event_cache is not None:
|
||||
cached = event_cache.get("plain_text")
|
||||
if isinstance(cached, str):
|
||||
return cached
|
||||
async def _patched_handle_event(bot: Bot, event: Event) -> None:
|
||||
show_log = True
|
||||
escape_tag = getattr(nb_message, "escape_tag")
|
||||
logger_ = getattr(nb_message, "logger")
|
||||
no_log_exception = getattr(nb_message, "NoLogException")
|
||||
|
||||
text = ""
|
||||
if message is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
text = message.extract_plain_text()
|
||||
if not text and event is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
text = (event.get_plaintext() or "").strip()
|
||||
log_msg = f"<m>{escape_tag(bot.type)} {escape_tag(bot.self_id)}</m> | "
|
||||
try:
|
||||
log_msg += event.get_log_string()
|
||||
except no_log_exception:
|
||||
show_log = False
|
||||
if show_log:
|
||||
logger_.opt(colors=True).success(log_msg)
|
||||
|
||||
if event_cache is not None:
|
||||
event_cache["plain_text"] = text
|
||||
return text
|
||||
state = {}
|
||||
dependency_cache = {}
|
||||
async_exit_stack = getattr(nb_message, "AsyncExitStack")
|
||||
apply_event_preprocessors = getattr(nb_message, "_apply_event_preprocessors")
|
||||
apply_event_postprocessors = getattr(nb_message, "_apply_event_postprocessors")
|
||||
trie_rule = getattr(nb_message, "TrieRule")
|
||||
matchers = getattr(nb_message, "matchers")
|
||||
catch = getattr(nb_message, "catch")
|
||||
stop_propagation = getattr(nb_message, "StopPropagation")
|
||||
handle_exception = getattr(nb_message, "_handle_exception")
|
||||
anyio_mod = getattr(nb_message, "anyio")
|
||||
run_coro_with_shield = getattr(nb_message, "run_coro_with_shield")
|
||||
|
||||
async with async_exit_stack() as stack:
|
||||
if not await apply_event_preprocessors(
|
||||
bot=bot,
|
||||
event=event,
|
||||
state=state,
|
||||
stack=stack,
|
||||
dependency_cache=dependency_cache,
|
||||
):
|
||||
return
|
||||
|
||||
try:
|
||||
trie_rule.get_value(bot, event, state)
|
||||
except Exception as e:
|
||||
logger_.opt(colors=True, exception=e).warning(
|
||||
"Error while parsing command for event"
|
||||
)
|
||||
_prepare_handle_event_state(event, state)
|
||||
|
||||
break_flag = False
|
||||
|
||||
def _handle_stop_propagation(_exc_group) -> None:
|
||||
nonlocal break_flag
|
||||
break_flag = True
|
||||
logger_.debug("Stop event propagation")
|
||||
|
||||
for priority in sorted(matchers.keys()):
|
||||
if break_flag:
|
||||
break
|
||||
|
||||
if show_log:
|
||||
logger_.debug(f"Checking for matchers in priority {priority}...")
|
||||
|
||||
if not (priority_matchers := matchers[priority]):
|
||||
continue
|
||||
|
||||
with catch(
|
||||
{
|
||||
stop_propagation: _handle_stop_propagation,
|
||||
Exception: handle_exception(
|
||||
"<r><bg #f8bbd0>Error when checking Matcher.</bg #f8bbd0></r>"
|
||||
),
|
||||
}
|
||||
):
|
||||
async with anyio_mod.create_task_group() as tg:
|
||||
for matcher in priority_matchers:
|
||||
skip, reason = _check_matcher_prefilter_before_task(
|
||||
matcher,
|
||||
event,
|
||||
state,
|
||||
)
|
||||
_record_prefilter_stats(skip, reason, "before_task")
|
||||
if skip:
|
||||
continue
|
||||
matcher_state = _build_matcher_state(state)
|
||||
tg.start_soon(
|
||||
run_coro_with_shield,
|
||||
_run_selected_matcher(
|
||||
matcher,
|
||||
bot,
|
||||
event,
|
||||
matcher_state,
|
||||
stack,
|
||||
dependency_cache,
|
||||
),
|
||||
)
|
||||
|
||||
if show_log:
|
||||
logger_.debug("Checking for matchers completed")
|
||||
|
||||
await apply_event_postprocessors(bot, event, state, stack, dependency_cache)
|
||||
|
||||
|
||||
def _install_handle_event_selector() -> None:
|
||||
global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT
|
||||
if _HANDLE_EVENT_PATCHED:
|
||||
return
|
||||
_ORIGINAL_HANDLE_EVENT = nb_message.handle_event
|
||||
nb_message.handle_event = _patched_handle_event # type: ignore[assignment]
|
||||
for module_name in (
|
||||
"nonebot.adapters.onebot.v11.bot",
|
||||
"nonebot.adapters.onebot.v12.bot",
|
||||
"onebug.mixin.process",
|
||||
):
|
||||
with contextlib.suppress(Exception):
|
||||
module = importlib.import_module(module_name)
|
||||
current = getattr(module, "handle_event", None)
|
||||
if current is not None:
|
||||
_ORIGINAL_ADAPTER_HANDLE_EVENTS[module] = current
|
||||
setattr(module, "handle_event", _patched_handle_event)
|
||||
_HANDLE_EVENT_PATCHED = True
|
||||
|
||||
|
||||
def _uninstall_handle_event_selector() -> None:
|
||||
global _HANDLE_EVENT_PATCHED, _ORIGINAL_HANDLE_EVENT
|
||||
if not _HANDLE_EVENT_PATCHED:
|
||||
return
|
||||
if _ORIGINAL_HANDLE_EVENT is not None:
|
||||
nb_message.handle_event = _ORIGINAL_HANDLE_EVENT # type: ignore[assignment]
|
||||
for module, original in list(_ORIGINAL_ADAPTER_HANDLE_EVENTS.items()):
|
||||
with contextlib.suppress(Exception):
|
||||
setattr(module, "handle_event", original)
|
||||
_ORIGINAL_ADAPTER_HANDLE_EVENTS.clear()
|
||||
_HANDLE_EVENT_PATCHED = False
|
||||
_ORIGINAL_HANDLE_EVENT = None
|
||||
|
||||
|
||||
async def _get_route_context(text: str, event_cache: dict | None) -> set[str]:
|
||||
@@ -843,12 +1028,14 @@ async def start_auth_runtime_tasks() -> None:
|
||||
global _CACHE_SWEEP_TASK
|
||||
await _ensure_route_index()
|
||||
_install_matcher_prefilter()
|
||||
_install_handle_event_selector()
|
||||
if _CACHE_SWEEP_TASK is None or _CACHE_SWEEP_TASK.done():
|
||||
_CACHE_SWEEP_TASK = asyncio.create_task(_cache_sweep_loop())
|
||||
|
||||
|
||||
async def stop_auth_runtime_tasks() -> None:
|
||||
global _CACHE_SWEEP_TASK
|
||||
_uninstall_handle_event_selector()
|
||||
_uninstall_matcher_prefilter()
|
||||
task = _CACHE_SWEEP_TASK
|
||||
_CACHE_SWEEP_TASK = None
|
||||
@@ -873,6 +1060,12 @@ async def _has_limits_cached(module: str, event_cache: dict | None) -> bool:
|
||||
@contextlib.asynccontextmanager
|
||||
async def _db_section():
|
||||
global DB_ACTIVE_COUNT
|
||||
if DB_SEMAPHORE.locked():
|
||||
logger.warning(
|
||||
"db semaphore saturated, allowing permission check to continue",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
raise PermissionExemption("db semaphore saturated, allow pass")
|
||||
await DB_SEMAPHORE.acquire()
|
||||
DB_ACTIVE_COUNT += 1
|
||||
try:
|
||||
@@ -918,7 +1111,9 @@ def _group_has_plugin_block(group, module: str) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _needs_auth_plugin(plugin: PluginInfo, group, entity) -> bool:
|
||||
def _needs_auth_plugin(plugin: PluginInfo, context: PermissionContext) -> bool:
|
||||
group = context.group
|
||||
entity = context.entity
|
||||
if plugin.block_type == BlockType.ALL and not plugin.status:
|
||||
if group and getattr(group, "is_super", False):
|
||||
return False
|
||||
@@ -953,11 +1148,11 @@ async def _get_bot_data_cached(
|
||||
|
||||
|
||||
async def _get_admin_levels_cached(
|
||||
session: Uninfo, entity, event_cache
|
||||
entity, event_cache
|
||||
) -> tuple[tuple[LevelUserSnapshot | None, LevelUserSnapshot | None] | None, bool]:
|
||||
if event_cache is not None and "admin_levels" in event_cache:
|
||||
return event_cache.get("admin_levels"), event_cache.get("admin_timeout", False)
|
||||
levels = await LevelUserMemoryCache.get_levels(session.user.id, entity.group_id)
|
||||
levels = await LevelUserMemoryCache.get_levels(entity.user_id, entity.group_id)
|
||||
if event_cache is not None:
|
||||
event_cache["admin_levels"] = levels
|
||||
event_cache["admin_timeout"] = False
|
||||
@@ -1074,12 +1269,18 @@ async def get_plugin_and_user(
|
||||
if user_id in user_cache:
|
||||
user = user_cache[user_id]
|
||||
else:
|
||||
async with _db_section():
|
||||
user = await _fetch_user_readonly(user_dao, user_id)
|
||||
try:
|
||||
async with _db_section():
|
||||
user = await _fetch_user_readonly(user_dao, user_id)
|
||||
except PermissionExemption:
|
||||
user = None
|
||||
user_cache[user_id] = user
|
||||
else:
|
||||
async with _db_section():
|
||||
user = await _fetch_user_readonly(user_dao, user_id)
|
||||
try:
|
||||
async with _db_section():
|
||||
user = await _fetch_user_readonly(user_dao, user_id)
|
||||
except PermissionExemption:
|
||||
user = None
|
||||
|
||||
return plugin, user
|
||||
|
||||
@@ -1089,7 +1290,7 @@ async def get_plugin_cost(
|
||||
plugin: PluginInfo,
|
||||
session: Uninfo,
|
||||
*,
|
||||
is_superuser: bool = False,
|
||||
context: PermissionContext | None = None,
|
||||
) -> int:
|
||||
"""获取插件费用
|
||||
|
||||
@@ -1106,7 +1307,10 @@ async def get_plugin_cost(
|
||||
返回:
|
||||
int: 调用插件金币费用
|
||||
"""
|
||||
cost_gold = await with_timeout(auth_cost(user, plugin, session), name="auth_cost")
|
||||
cost_gold = await with_timeout(
|
||||
auth_cost(user, plugin, session, context=context), name="auth_cost"
|
||||
)
|
||||
is_superuser = context.is_superuser if context is not None else False
|
||||
if is_superuser:
|
||||
if plugin.plugin_type == PluginType.SUPERUSER:
|
||||
raise IsSuperuserException()
|
||||
@@ -1124,7 +1328,7 @@ async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo
|
||||
cost_gold: 消耗金币
|
||||
session: Uninfo
|
||||
"""
|
||||
user_dao = DataAccess(UserConsole)
|
||||
should_clear_cache = False
|
||||
try:
|
||||
await with_timeout(
|
||||
UserConsole.reduce_gold(
|
||||
@@ -1141,14 +1345,16 @@ async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo
|
||||
u.gold = 0
|
||||
await u.save(update_fields=["gold"])
|
||||
except asyncio.TimeoutError:
|
||||
should_clear_cache = True
|
||||
logger.error(
|
||||
f"扣除金币超时,用户: {user_id}, 金币: {cost_gold}",
|
||||
LOGGER_COMMAND,
|
||||
session=session,
|
||||
)
|
||||
|
||||
# 清除缓存,使下次查询时从数据库获取最新数据
|
||||
await user_dao.clear_cache(user_id=user_id)
|
||||
# 正常写入路径由 UserConsole.save() 统一失效缓存;超时状态不确定时兜底清理。
|
||||
if should_clear_cache:
|
||||
await DataAccess(UserConsole).clear_cache(user_id=user_id)
|
||||
logger.debug(f"调用功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session)
|
||||
|
||||
|
||||
@@ -1174,16 +1380,15 @@ async def time_hook(coro, name, recorder: HookTraceRecorder | None = None):
|
||||
|
||||
|
||||
async def _enter_hooks_section():
|
||||
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。"""
|
||||
"""尝试获取全局信号量并更新计数器,饱和时快速放行。"""
|
||||
global HOOKS_ACTIVE_COUNT
|
||||
try:
|
||||
await asyncio.wait_for(HOOKS_SEMAPHORE.acquire(), timeout=TIMEOUT_SECONDS)
|
||||
except asyncio.TimeoutError:
|
||||
if HOOKS_SEMAPHORE.locked():
|
||||
logger.warning(
|
||||
"hooks semaphore acquire timeout, allowing pass",
|
||||
"hooks semaphore saturated, allowing pass",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
raise PermissionExemption("hooks semaphore timeout, allow pass")
|
||||
raise PermissionExemption("hooks semaphore saturated, allow pass")
|
||||
await HOOKS_SEMAPHORE.acquire()
|
||||
HOOKS_ACTIVE_COUNT += 1
|
||||
|
||||
|
||||
@@ -1197,14 +1402,7 @@ async def _leave_hooks_section():
|
||||
|
||||
async def route_precheck(
|
||||
matcher: Matcher,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
message: UniMsg | None,
|
||||
*,
|
||||
entity=None,
|
||||
event_cache: dict | None = None,
|
||||
text: str | None = None,
|
||||
route_modules: set[str] | None = None,
|
||||
context: EventContext,
|
||||
) -> bool:
|
||||
module = matcher.plugin_name or ""
|
||||
if not module:
|
||||
@@ -1213,19 +1411,20 @@ async def route_precheck(
|
||||
return False
|
||||
if not _is_command_matcher_class(type(matcher)):
|
||||
return False
|
||||
if entity is None:
|
||||
entity = get_entity_ids(session)
|
||||
if event_cache is None:
|
||||
event_cache = _get_event_cache(event, session, entity)
|
||||
if text is None:
|
||||
text = _get_message_text(message, event_cache, event)
|
||||
|
||||
route_modules = context.route_modules if context.route_modules_loaded else None
|
||||
if route_modules is None:
|
||||
route_modules = await _get_route_context(text, event_cache)
|
||||
route_modules = await _get_route_context(
|
||||
context.plain_text,
|
||||
context.event_cache,
|
||||
)
|
||||
set_route_modules(None, context, route_modules)
|
||||
|
||||
if module in _ROUTE_MODULES_WITH_COMMANDS and module not in route_modules:
|
||||
if _matcher_has_alconna_shortcuts(type(matcher)):
|
||||
return False
|
||||
if event_cache is not None:
|
||||
event_cache["route_skip"] = True
|
||||
if context.event_cache is not None:
|
||||
context.event_cache["route_skip"] = True
|
||||
return True
|
||||
return False
|
||||
|
||||
@@ -1235,14 +1434,10 @@ async def auth(
|
||||
event: Event,
|
||||
bot: Bot,
|
||||
session: Uninfo,
|
||||
message: UniMsg | None,
|
||||
*,
|
||||
context: EventContext,
|
||||
skip_ban: bool = False,
|
||||
entity=None,
|
||||
event_cache: dict | None = None,
|
||||
text: str | None = None,
|
||||
route_modules: set[str] | None = None,
|
||||
is_superuser: bool = False,
|
||||
state: dict | None = None,
|
||||
):
|
||||
"""权限检查
|
||||
|
||||
@@ -1251,20 +1446,28 @@ async def auth(
|
||||
event: Event
|
||||
bot: bot
|
||||
session: Uninfo
|
||||
message: UniMsg
|
||||
context: EventContext
|
||||
"""
|
||||
start_time = time.time()
|
||||
cost_gold = 0
|
||||
ignore_flag = False
|
||||
if entity is None:
|
||||
entity = get_entity_ids(session)
|
||||
entity = context.entity
|
||||
event_cache = context.event_cache
|
||||
text = context.plain_text
|
||||
is_superuser = context.is_superuser
|
||||
route_modules = context.route_modules if context.route_modules_loaded else None
|
||||
module = matcher.plugin_name or ""
|
||||
is_command_matcher = _is_command_matcher_class(type(matcher))
|
||||
if event_cache is None:
|
||||
event_cache = _get_event_cache(event, session, entity)
|
||||
auth_allowed = None
|
||||
auth_result_cache = None
|
||||
admin_checked_pre = False
|
||||
permission_context: PermissionContext | None = None
|
||||
side_effect_cache = get_permission_side_effect_cache(
|
||||
state=state,
|
||||
event_cache=event_cache,
|
||||
)
|
||||
side_effect_lock = None
|
||||
entered_side_effect_lock = False
|
||||
|
||||
# 仅在慢请求时记录 hook 明细,避免热路径高频构造字符串
|
||||
hook_recorder = HookTraceRecorder(start_time)
|
||||
@@ -1278,14 +1481,17 @@ async def auth(
|
||||
auth_allowed = True
|
||||
return
|
||||
|
||||
if event_cache is not None:
|
||||
auth_result_cache = event_cache.setdefault("auth_result", {})
|
||||
cached_result = auth_result_cache.get(module)
|
||||
if cached_result is not None:
|
||||
allowed, reason = cached_result
|
||||
if not allowed:
|
||||
raise SkipPluginException(reason or "auth cached skip")
|
||||
return
|
||||
side_effect_lock = side_effect_cache.lock_for(module)
|
||||
await side_effect_lock.acquire()
|
||||
entered_side_effect_lock = True
|
||||
|
||||
auth_result_cache = side_effect_cache.auth_results
|
||||
cached_result = auth_result_cache.get(module)
|
||||
if cached_result is not None:
|
||||
allowed, reason = cached_result
|
||||
if not allowed:
|
||||
raise SkipPluginException(reason or "auth cached skip")
|
||||
return
|
||||
|
||||
if _is_hidden_plugin(matcher):
|
||||
auth_allowed = True
|
||||
@@ -1293,10 +1499,9 @@ async def auth(
|
||||
if event_cache is not None and event_cache.get("ban_state") is True:
|
||||
raise SkipPluginException("user or group banned (cached)")
|
||||
|
||||
if text is None:
|
||||
text = _get_message_text(message, event_cache, event)
|
||||
if route_modules is None:
|
||||
route_modules = await _get_route_context(text, event_cache)
|
||||
set_route_modules(state, context, route_modules)
|
||||
route_skip_checks = (
|
||||
is_command_matcher
|
||||
and module in _ROUTE_MODULES_WITH_COMMANDS
|
||||
@@ -1310,7 +1515,7 @@ async def auth(
|
||||
auth_allowed = True
|
||||
return
|
||||
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
platform = context.platform
|
||||
# 获取插件和用户数据
|
||||
plugin_user_start = time.time()
|
||||
try:
|
||||
@@ -1336,6 +1541,14 @@ async def auth(
|
||||
auth_allowed = True
|
||||
return
|
||||
|
||||
permission_context = PermissionContext(
|
||||
event=context,
|
||||
module=module,
|
||||
plugin=plugin,
|
||||
user=user,
|
||||
)
|
||||
store_permission_context(state, permission_context)
|
||||
|
||||
if not route_skip_checks and _needs_admin_check(plugin):
|
||||
if plugin.plugin_type in {
|
||||
PluginType.SUPERUSER,
|
||||
@@ -1356,13 +1569,18 @@ async def auth(
|
||||
admin_timeout = False
|
||||
if event_cache is not None:
|
||||
admin_levels, admin_timeout = await _get_admin_levels_cached(
|
||||
session, entity, event_cache
|
||||
entity, event_cache
|
||||
)
|
||||
permission_context.admin_levels = admin_levels
|
||||
if admin_timeout:
|
||||
hook_recorder.set("auth_admin", "timeout")
|
||||
else:
|
||||
admin_start = time.time()
|
||||
await auth_admin(plugin, session, cached_levels=admin_levels)
|
||||
await auth_admin(
|
||||
plugin,
|
||||
session,
|
||||
context=permission_context,
|
||||
)
|
||||
hook_recorder.set(
|
||||
"auth_admin", f"{time.time() - admin_start:.3f}s(pre)"
|
||||
)
|
||||
@@ -1386,8 +1604,7 @@ async def auth(
|
||||
matcher,
|
||||
session,
|
||||
plugin,
|
||||
entity=entity,
|
||||
is_superuser=is_superuser,
|
||||
context=permission_context,
|
||||
)
|
||||
hook_recorder.set("auth_ban", f"{time.time() - ban_start:.3f}s")
|
||||
if event_cache is not None:
|
||||
@@ -1407,7 +1624,7 @@ async def auth(
|
||||
user,
|
||||
plugin,
|
||||
session,
|
||||
is_superuser=is_superuser,
|
||||
context=permission_context,
|
||||
),
|
||||
name="get_plugin_cost",
|
||||
)
|
||||
@@ -1421,7 +1638,7 @@ async def auth(
|
||||
hook_recorder.set("cost_gold", "skipped")
|
||||
|
||||
# 执行 bot_filter
|
||||
bot_filter(session)
|
||||
bot_filter(session, context=permission_context)
|
||||
|
||||
group = await _get_group_cached(entity, event_cache)
|
||||
|
||||
@@ -1439,13 +1656,23 @@ async def auth(
|
||||
and not route_skip_checks
|
||||
):
|
||||
admin_levels, admin_timeout = await _get_admin_levels_cached(
|
||||
session, entity, event_cache
|
||||
entity, event_cache
|
||||
)
|
||||
|
||||
permission_context.group = group
|
||||
permission_context.bot_data = bot_data
|
||||
if admin_levels is not None:
|
||||
permission_context.admin_levels = admin_levels
|
||||
store_permission_context(state, permission_context)
|
||||
|
||||
# 并行执行所有 hook 检查,并记录执行时间
|
||||
hooks_start = time.time()
|
||||
allow_sleep_bypass = _is_bot_wake_command(module, text)
|
||||
|
||||
# 先进入 hooks 并行检查区域;饱和时快速放行,避免创建并积压协程。
|
||||
await _enter_hooks_section()
|
||||
entered_hooks = True
|
||||
|
||||
# 创建所有 hook 任务
|
||||
hook_tasks = []
|
||||
if event_cache is None:
|
||||
@@ -1455,6 +1682,7 @@ async def auth(
|
||||
plugin,
|
||||
bot.self_id,
|
||||
allow_sleep_bypass=allow_sleep_bypass,
|
||||
context=permission_context,
|
||||
),
|
||||
"auth_bot",
|
||||
hook_recorder,
|
||||
@@ -1472,6 +1700,7 @@ async def auth(
|
||||
bot_data=bot_data,
|
||||
skip_fetch=True,
|
||||
allow_sleep_bypass=allow_sleep_bypass,
|
||||
context=permission_context,
|
||||
),
|
||||
"auth_bot",
|
||||
hook_recorder,
|
||||
@@ -1483,7 +1712,13 @@ async def auth(
|
||||
else:
|
||||
hook_tasks.append(
|
||||
time_hook(
|
||||
auth_group(plugin, group, text, entity.group_id),
|
||||
auth_group(
|
||||
plugin,
|
||||
group,
|
||||
text,
|
||||
entity.group_id,
|
||||
context=permission_context,
|
||||
),
|
||||
"auth_group",
|
||||
hook_recorder,
|
||||
)
|
||||
@@ -1492,7 +1727,11 @@ async def auth(
|
||||
if not route_skip_checks and plugin.admin_level and not admin_checked_pre:
|
||||
if event_cache is None:
|
||||
hook_tasks.append(
|
||||
time_hook(auth_admin(plugin, session), "auth_admin", hook_recorder)
|
||||
time_hook(
|
||||
auth_admin(plugin, session, context=permission_context),
|
||||
"auth_admin",
|
||||
hook_recorder,
|
||||
)
|
||||
)
|
||||
else:
|
||||
if admin_timeout:
|
||||
@@ -1500,7 +1739,11 @@ async def auth(
|
||||
else:
|
||||
hook_tasks.append(
|
||||
time_hook(
|
||||
auth_admin(plugin, session, cached_levels=admin_levels),
|
||||
auth_admin(
|
||||
plugin,
|
||||
session,
|
||||
context=permission_context,
|
||||
),
|
||||
"auth_admin",
|
||||
hook_recorder,
|
||||
)
|
||||
@@ -1510,7 +1753,7 @@ async def auth(
|
||||
|
||||
if is_superuser:
|
||||
hook_recorder.set("auth_plugin", "superuser")
|
||||
elif not route_skip_checks and _needs_auth_plugin(plugin, group, entity):
|
||||
elif not route_skip_checks and _needs_auth_plugin(plugin, permission_context):
|
||||
hook_tasks.append(
|
||||
time_hook(
|
||||
auth_plugin(
|
||||
@@ -1518,6 +1761,7 @@ async def auth(
|
||||
group,
|
||||
session,
|
||||
event,
|
||||
context=permission_context,
|
||||
skip_group_block=is_superuser,
|
||||
),
|
||||
"auth_plugin",
|
||||
@@ -1531,18 +1775,17 @@ async def auth(
|
||||
has_limits = await _has_limits_cached(module, event_cache)
|
||||
if has_limits:
|
||||
hook_tasks.append(
|
||||
time_hook(auth_limit(plugin, session), "auth_limit", hook_recorder)
|
||||
time_hook(
|
||||
auth_limit(plugin, session, context=permission_context),
|
||||
"auth_limit",
|
||||
hook_recorder,
|
||||
)
|
||||
)
|
||||
else:
|
||||
hook_recorder.set("auth_limit", "skipped")
|
||||
else:
|
||||
hook_recorder.set("auth_limit", "skipped")
|
||||
|
||||
if hook_tasks:
|
||||
# 进入 hooks 并行检查区域(会在高并发时排队)
|
||||
await _enter_hooks_section()
|
||||
entered_hooks = True
|
||||
|
||||
# 使用 gather 并行执行所有 hook,但添加总体超时控制
|
||||
try:
|
||||
await with_timeout(
|
||||
@@ -1599,6 +1842,9 @@ async def auth(
|
||||
)
|
||||
if auth_result_cache is not None and auth_allowed is not None:
|
||||
auth_result_cache[module] = (auth_allowed, None)
|
||||
if entered_side_effect_lock and side_effect_lock is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
side_effect_lock.release()
|
||||
# 扣除金币
|
||||
if not ignore_flag and cost_gold > 0:
|
||||
gold_start = time.time()
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import contextlib
|
||||
import time
|
||||
|
||||
from nonebot import get_driver
|
||||
@@ -12,14 +11,20 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import is_cache_ready
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import is_overloaded
|
||||
from zhenxun.services.message_load import is_overloaded, mark_activity
|
||||
from zhenxun.services.runtime_bootstrap import register_runtime_bootstrap
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
from .auth.config import LOGGER_COMMAND
|
||||
from .auth.context import (
|
||||
get_event_context,
|
||||
get_or_create_event_context,
|
||||
resolve_actor_user_id,
|
||||
resolve_event_channel_id,
|
||||
resolve_event_group_id,
|
||||
set_route_modules,
|
||||
)
|
||||
from .auth_checker import (
|
||||
LimitManager,
|
||||
_get_event_cache,
|
||||
_get_route_context,
|
||||
auth,
|
||||
route_precheck,
|
||||
@@ -41,17 +46,6 @@ async def _mark_bot_connected(bot: Bot):
|
||||
_BOT_CONNECT_TS = time.time()
|
||||
|
||||
|
||||
def _extract_plain_text(message: UniMsg | None, event: Event) -> str:
|
||||
if message is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
return message.extract_plain_text()
|
||||
with contextlib.suppress(Exception):
|
||||
plain = event.get_plaintext()
|
||||
if plain:
|
||||
return plain.strip()
|
||||
return ""
|
||||
|
||||
|
||||
@driver.on_startup
|
||||
async def _start_auth_runtime_tasks():
|
||||
await start_auth_runtime_tasks()
|
||||
@@ -72,37 +66,9 @@ def _skip_auth_for_plugin(matcher: Matcher) -> bool:
|
||||
return "chat_history" in module_name
|
||||
|
||||
|
||||
def _resolve_actor_user_id(event: Event, fallback_user_id: str) -> str:
|
||||
"""优先使用事件发起者ID,避免 notice 场景 session.user 指向 bot 自身。"""
|
||||
event_user_id = getattr(event, "user_id", None)
|
||||
if event_user_id is None:
|
||||
return fallback_user_id
|
||||
event_user_id = str(event_user_id)
|
||||
return event_user_id or fallback_user_id
|
||||
|
||||
|
||||
def _resolve_event_group_id(event: Event, fallback_group_id: str | None) -> str | None:
|
||||
"""notice 场景 session.group 可能缺失,回退到事件上的 group_id。"""
|
||||
event_group_id = getattr(event, "group_id", None)
|
||||
if event_group_id is None:
|
||||
return fallback_group_id
|
||||
resolved = str(event_group_id)
|
||||
return resolved or fallback_group_id
|
||||
|
||||
|
||||
def _resolve_event_channel_id(
|
||||
event: Event, fallback_channel_id: str | None
|
||||
) -> str | None:
|
||||
"""频道场景回退到事件上的 channel_id。"""
|
||||
event_channel_id = getattr(event, "channel_id", None)
|
||||
if event_channel_id is None:
|
||||
return fallback_channel_id
|
||||
resolved = str(event_channel_id)
|
||||
return resolved or fallback_channel_id
|
||||
|
||||
|
||||
@event_preprocessor
|
||||
async def _drop_message_before_cache_ready(event: Event):
|
||||
mark_activity()
|
||||
if event.get_type() != "message":
|
||||
return
|
||||
if not is_cache_ready():
|
||||
@@ -130,46 +96,22 @@ async def _auth_preprocessor(
|
||||
return
|
||||
|
||||
start_time = time.time()
|
||||
entity = state.get("_zx_entity")
|
||||
if entity is None:
|
||||
entity = get_entity_ids(session)
|
||||
entity.user_id = _resolve_actor_user_id(event, entity.user_id)
|
||||
entity.group_id = _resolve_event_group_id(event, entity.group_id)
|
||||
entity.channel_id = _resolve_event_channel_id(event, entity.channel_id)
|
||||
state["_zx_entity"] = entity
|
||||
|
||||
event_cache = state.get("_zx_event_cache")
|
||||
if event_cache is None:
|
||||
event_cache = _get_event_cache(event, session, entity)
|
||||
state["_zx_event_cache"] = event_cache
|
||||
|
||||
text = state.get("_zx_plain_text")
|
||||
if text is None:
|
||||
text = _extract_plain_text(message, event)
|
||||
state["_zx_plain_text"] = text
|
||||
if event_cache is not None:
|
||||
event_cache["plain_text"] = text
|
||||
|
||||
route_modules = state.get("_zx_route_modules")
|
||||
if route_modules is None:
|
||||
route_modules = await _get_route_context(text, event_cache)
|
||||
state["_zx_route_modules"] = route_modules
|
||||
|
||||
is_superuser = state.get("_zx_is_superuser")
|
||||
if is_superuser is None:
|
||||
is_superuser = entity.user_id in bot.config.superusers
|
||||
state["_zx_is_superuser"] = is_superuser
|
||||
|
||||
if await route_precheck(
|
||||
matcher,
|
||||
event_context = get_or_create_event_context(
|
||||
bot,
|
||||
event,
|
||||
session,
|
||||
message,
|
||||
entity=entity,
|
||||
event_cache=event_cache,
|
||||
text=text,
|
||||
route_modules=route_modules,
|
||||
):
|
||||
state,
|
||||
message=message,
|
||||
)
|
||||
|
||||
if not event_context.route_modules_loaded:
|
||||
route_modules = await _get_route_context(
|
||||
event_context.plain_text,
|
||||
event_context.event_cache,
|
||||
)
|
||||
set_route_modules(state, event_context, route_modules)
|
||||
|
||||
if await route_precheck(matcher, event_context):
|
||||
return
|
||||
|
||||
try:
|
||||
@@ -178,13 +120,9 @@ async def _auth_preprocessor(
|
||||
event,
|
||||
bot,
|
||||
session,
|
||||
message,
|
||||
context=event_context,
|
||||
skip_ban=False,
|
||||
entity=entity,
|
||||
event_cache=event_cache,
|
||||
text=text,
|
||||
route_modules=route_modules,
|
||||
is_superuser=is_superuser,
|
||||
state=state,
|
||||
)
|
||||
except IgnoredException:
|
||||
raise
|
||||
@@ -203,16 +141,27 @@ async def _auth_preprocessor(
|
||||
|
||||
|
||||
@run_postprocessor
|
||||
async def _unblock_after_matcher(matcher: Matcher, session: Uninfo, event: Event):
|
||||
user_id = _resolve_actor_user_id(event, session.user.id)
|
||||
group_id = _resolve_event_group_id(event, None)
|
||||
channel_id = _resolve_event_channel_id(event, None)
|
||||
if session.group:
|
||||
if session.group.parent:
|
||||
group_id = session.group.parent.id
|
||||
channel_id = session.group.id
|
||||
else:
|
||||
group_id = session.group.id
|
||||
async def _unblock_after_matcher(
|
||||
matcher: Matcher,
|
||||
session: Uninfo,
|
||||
event: Event,
|
||||
state: T_State,
|
||||
):
|
||||
context = get_event_context(state)
|
||||
if context is not None:
|
||||
user_id = context.user_id
|
||||
group_id = context.group_id
|
||||
channel_id = context.channel_id
|
||||
else:
|
||||
user_id = resolve_actor_user_id(event, session.user.id)
|
||||
group_id = resolve_event_group_id(event, None)
|
||||
channel_id = resolve_event_channel_id(event, None)
|
||||
if session.group:
|
||||
if session.group.parent:
|
||||
group_id = session.group.parent.id
|
||||
channel_id = session.group.id
|
||||
else:
|
||||
group_id = session.group.id
|
||||
if user_id and matcher.plugin:
|
||||
module = matcher.plugin.name
|
||||
LimitManager.unblock(module, user_id, group_id, channel_id)
|
||||
|
||||
@@ -16,13 +16,7 @@ from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
malicious_check_time = Config.get_config("hook", "MALICIOUS_CHECK_TIME")
|
||||
malicious_ban_count = Config.get_config("hook", "MALICIOUS_BAN_COUNT")
|
||||
|
||||
if not malicious_check_time:
|
||||
raise ValueError("模块: [hook], 配置项: [MALICIOUS_CHECK_TIME] 为空或小于0")
|
||||
if not malicious_ban_count:
|
||||
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_COUNT] 为空或小于0")
|
||||
from .auth.context import resolve_actor_user_id, resolve_event_group_id
|
||||
|
||||
|
||||
class BanCheckLimiter:
|
||||
@@ -36,6 +30,10 @@ class BanCheckLimiter:
|
||||
self.default_check_time = default_check_time
|
||||
self.default_count = default_count
|
||||
|
||||
def configure(self, check_time: float, count: int) -> None:
|
||||
self.default_check_time = check_time
|
||||
self.default_count = count
|
||||
|
||||
def add(self, key: str | float):
|
||||
if self.mint[key] == 1:
|
||||
self.mtime[key] = time.time()
|
||||
@@ -59,11 +57,22 @@ class BanCheckLimiter:
|
||||
|
||||
|
||||
_blmt = BanCheckLimiter(
|
||||
malicious_check_time,
|
||||
malicious_ban_count,
|
||||
5,
|
||||
4,
|
||||
)
|
||||
|
||||
|
||||
def _get_positive_config(key: str, cast_type: type[int] | type[float]) -> int | float:
|
||||
value = Config.get_config("hook", key)
|
||||
try:
|
||||
parsed_value = cast_type(value)
|
||||
except (TypeError, ValueError) as e:
|
||||
raise ValueError(f"模块: [hook], 配置项: [{key}] 不是有效数字") from e
|
||||
if parsed_value <= 0:
|
||||
raise ValueError(f"模块: [hook], 配置项: [{key}] 为空或小于0")
|
||||
return parsed_value
|
||||
|
||||
|
||||
# 恶意触发命令检测
|
||||
@run_preprocessor
|
||||
async def _(
|
||||
@@ -88,11 +97,12 @@ async def _(
|
||||
else:
|
||||
return
|
||||
|
||||
user_id = session.id1
|
||||
group_id = session.id3 or session.id2
|
||||
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
|
||||
if not malicious_ban_time:
|
||||
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0")
|
||||
user_id = resolve_actor_user_id(event, session.id1)
|
||||
group_id = resolve_event_group_id(event, session.id3 or session.id2)
|
||||
malicious_check_time = float(_get_positive_config("MALICIOUS_CHECK_TIME", float))
|
||||
malicious_ban_count = int(_get_positive_config("MALICIOUS_BAN_COUNT", int))
|
||||
malicious_ban_time = int(_get_positive_config("MALICIOUS_BAN_TIME", int))
|
||||
_blmt.configure(malicious_check_time, malicious_ban_count)
|
||||
if user_id and module:
|
||||
if _blmt.check(f"{user_id}__{module}"):
|
||||
await BanConsole.ban(
|
||||
|
||||
@@ -29,7 +29,6 @@ def register_cache_types():
|
||||
GroupPluginSetting,
|
||||
key_format="{group_id}_{plugin_name}_{key}",
|
||||
)
|
||||
CacheRegistry.register(CacheType.GROUP_PLUGIN_SETTINGS_VIEW, dict)
|
||||
CacheRegistry.register(
|
||||
CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}"
|
||||
)
|
||||
|
||||
@@ -88,7 +88,7 @@ async def _handle_setting(
|
||||
)
|
||||
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=5)
|
||||
@PriorityLifecycle.on_startup(priority=4)
|
||||
async def _():
|
||||
"""
|
||||
初始化插件数据配置
|
||||
|
||||
@@ -3,12 +3,12 @@ from pathlib import Path
|
||||
import random
|
||||
import shutil
|
||||
|
||||
from aiocache import cached
|
||||
import ujson as json
|
||||
|
||||
from zhenxun.builtin_plugins.plugin_store.models import StorePluginInfo
|
||||
from zhenxun.configs.path_config import TEMP_PATH
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.cache.bounded_ttl import BoundedTTLCache
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.plugin_init import PluginInitManager
|
||||
from zhenxun.utils.enum import PluginType
|
||||
@@ -26,6 +26,14 @@ from .config import (
|
||||
)
|
||||
from .exceptions import PluginStoreException
|
||||
|
||||
_PLUGIN_STORE_DATA_CACHE = BoundedTTLCache[
|
||||
str, tuple[list[StorePluginInfo], list[StorePluginInfo]]
|
||||
](
|
||||
"PLUGIN_STORE_DATA",
|
||||
ttl_seconds=60,
|
||||
max_items=1,
|
||||
)
|
||||
|
||||
|
||||
def row_style(column: str, text: str) -> RowStyle:
|
||||
"""被动技能文本风格
|
||||
@@ -56,20 +64,9 @@ class StoreManager:
|
||||
relative_parts = [part for part in plugin_info.module_path.split(".") if part]
|
||||
relative_path = Path(*relative_parts) if relative_parts else Path(plugin_name)
|
||||
path = BASE_PATH.parent / relative_path
|
||||
if plugin_info.is_dir:
|
||||
return path
|
||||
return path.parent / f"{plugin_name}.py"
|
||||
return path if plugin_info.is_dir else path.parent / f"{plugin_name}.py"
|
||||
|
||||
@classmethod
|
||||
def _is_plugin_installed(
|
||||
cls, plugin_info: StorePluginInfo, *, is_external: bool
|
||||
) -> bool:
|
||||
return cls._resolve_local_plugin_path(
|
||||
plugin_info, is_external=is_external
|
||||
).exists()
|
||||
|
||||
@classmethod
|
||||
@cached(60)
|
||||
async def get_data(cls) -> tuple[list[StorePluginInfo], list[StorePluginInfo]]:
|
||||
"""获取插件信息数据
|
||||
|
||||
@@ -77,15 +74,22 @@ class StoreManager:
|
||||
tuple[list[StorePluginInfo], list[StorePluginInfo]]:
|
||||
原生插件信息数据,第三方插件信息数据
|
||||
"""
|
||||
cache_key = "plugins_json"
|
||||
if cached_data := await _PLUGIN_STORE_DATA_CACHE.get(cache_key):
|
||||
return cached_data
|
||||
|
||||
plugins = await RepoFileManager.get_file_content(
|
||||
DEFAULT_GITHUB_URL, "plugins.json"
|
||||
)
|
||||
extra_plugins = await RepoFileManager.get_file_content(
|
||||
EXTRA_GITHUB_URL, "plugins.json", "index"
|
||||
)
|
||||
return [StorePluginInfo(**plugin) for plugin in json.loads(plugins)], [
|
||||
StorePluginInfo(**plugin) for plugin in json.loads(extra_plugins)
|
||||
]
|
||||
result = (
|
||||
[StorePluginInfo(**plugin) for plugin in json.loads(plugins)],
|
||||
[StorePluginInfo(**plugin) for plugin in json.loads(extra_plugins)],
|
||||
)
|
||||
await _PLUGIN_STORE_DATA_CACHE.set(cache_key, result)
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def version_check(cls, plugin_info: StorePluginInfo, suc_plugin: dict[str, str]):
|
||||
@@ -330,13 +334,12 @@ class StoreManager:
|
||||
source: 源
|
||||
"""
|
||||
repo_type = RepoType.GITHUB if is_external else None
|
||||
if source == "ali":
|
||||
if (
|
||||
source != "ali" and source != "git" and plugin_info.ali_url
|
||||
) or source == "ali":
|
||||
repo_type = RepoType.ALIYUN
|
||||
elif source == "git":
|
||||
repo_type = RepoType.GITHUB
|
||||
else:
|
||||
if plugin_info.ali_url:
|
||||
repo_type = RepoType.ALIYUN
|
||||
module_path = plugin_info.module_path
|
||||
is_dir = plugin_info.is_dir
|
||||
github_url = plugin_info.github_url
|
||||
@@ -380,7 +383,7 @@ class StoreManager:
|
||||
requirement_file = target_dir / requirement_path.path
|
||||
if requirement_file.exists():
|
||||
is_install_req = True
|
||||
await VirtualEnvPackageManager.install_requirement(requirement_file)
|
||||
await VirtualEnvPackageManager.add_requirement(requirement_file)
|
||||
|
||||
if not is_install_req:
|
||||
# 从仓库根目录查找文件
|
||||
@@ -401,13 +404,13 @@ class StoreManager:
|
||||
f"开始安装插件 {module_path} 依赖文件: {requirement_path}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
await VirtualEnvPackageManager.install_requirement(requirement_path)
|
||||
await VirtualEnvPackageManager.add_requirement(requirement_path)
|
||||
if requirements_path.exists():
|
||||
logger.info(
|
||||
f"开始安装插件 {module_path} 依赖文件: {requirements_path}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
await VirtualEnvPackageManager.install_requirement(requirements_path)
|
||||
await VirtualEnvPackageManager.add_requirement(requirements_path)
|
||||
|
||||
@classmethod
|
||||
async def remove_plugin(cls, index_or_module: str) -> str:
|
||||
|
||||
@@ -19,9 +19,9 @@ from zhenxun.models.friend_user import FriendUser
|
||||
from zhenxun.models.goods_info import GoodsInfo
|
||||
from zhenxun.models.group_member_info import GroupInfoUser
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.models.user_gold_log import UserGoldLog
|
||||
from zhenxun.models.user_props_log import UserPropsLog
|
||||
from zhenxun.services import avatar_service
|
||||
from zhenxun.services.buffered_writers import append_user_gold_log
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.ui.models import ImageCell, TextCell
|
||||
from zhenxun.utils.enum import GoldHandle, PropHandle
|
||||
@@ -480,10 +480,6 @@ class ShopManage:
|
||||
).count()
|
||||
if goods.daily_limit and count >= goods.daily_limit:
|
||||
return "今天的购买已达限制了喔!"
|
||||
await UserGoldLog.create(user_id=user_id, gold=price, handle=GoldHandle.BUY)
|
||||
await UserPropsLog.create(
|
||||
user_id=user_id, uuid=goods.uuid, gold=price, num=num, handle=PropHandle.BUY
|
||||
)
|
||||
logger.info(
|
||||
f"花费 {price} 金币购买 {goods.goods_name} ×{num} 成功!",
|
||||
"购买道具",
|
||||
@@ -494,6 +490,12 @@ class ShopManage:
|
||||
user.props[goods.uuid] = 0
|
||||
user.props[goods.uuid] += num
|
||||
await user.save(update_fields=["gold", "props"])
|
||||
await append_user_gold_log(
|
||||
user_id=user_id, gold=int(price), handle=GoldHandle.BUY
|
||||
)
|
||||
await UserPropsLog.create(
|
||||
user_id=user_id, uuid=goods.uuid, gold=price, num=num, handle=PropHandle.BUY
|
||||
)
|
||||
return f"花费 {price} 金币购买 {goods.goods_name} ×{num} 成功!"
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -9,13 +9,17 @@ import pytz
|
||||
from zhenxun import ui
|
||||
from zhenxun.configs.path_config import IMAGE_PATH
|
||||
from zhenxun.models.friend_user import FriendUser
|
||||
from zhenxun.models.goods_info import GoodsInfo
|
||||
from zhenxun.models.group_member_info import GroupInfoUser
|
||||
from zhenxun.models.sign_log import SignLog
|
||||
from zhenxun.models.sign_user import SignUser
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services.avatar_service import avatar_service
|
||||
from zhenxun.services.buffered_writers import append_user_gold_log
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.ui.models import ImageCell, TextCell
|
||||
from zhenxun.utils.enum import GoldHandle
|
||||
from zhenxun.utils.exception import GoodsNotFound
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
from ._random_event import random_event
|
||||
@@ -182,11 +186,32 @@ class SignManage:
|
||||
gift = random_event(float(user.impression))
|
||||
if isinstance(gift, int):
|
||||
gold += gift
|
||||
await UserConsole.add_gold(user.user_id, gold + gift, "sign_in", platform)
|
||||
user_console = await UserConsole.get_user(user.user_id, platform)
|
||||
user_console.gold += gold
|
||||
await user_console.save(update_fields=["gold"])
|
||||
await append_user_gold_log(
|
||||
user_id=user.user_id,
|
||||
gold=gold,
|
||||
handle=GoldHandle.GET,
|
||||
source="sign_in",
|
||||
)
|
||||
gift = f"额外金币 +{gift}"
|
||||
else:
|
||||
await UserConsole.add_gold(user.user_id, gold, "sign_in", platform)
|
||||
await UserConsole.add_props_by_name(user.user_id, gift, 1, platform)
|
||||
goods = await GoodsInfo.get_or_none(goods_name=gift)
|
||||
if not goods:
|
||||
raise GoodsNotFound("未找到商品...")
|
||||
user_console = await UserConsole.get_user(user.user_id, platform)
|
||||
user_console.gold += gold
|
||||
if goods.uuid not in user_console.props:
|
||||
user_console.props[goods.uuid] = 0
|
||||
user_console.props[goods.uuid] += 1
|
||||
await user_console.save(update_fields=["gold", "props"])
|
||||
await append_user_gold_log(
|
||||
user_id=user.user_id,
|
||||
gold=gold,
|
||||
handle=GoldHandle.GET,
|
||||
source="sign_in",
|
||||
)
|
||||
gift += " + 1"
|
||||
logger.info(
|
||||
f"签到成功. score: {user.impression:.2f} "
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
|
||||
from nonebot import get_driver
|
||||
from nonebot.adapters import Bot, Event
|
||||
from nonebot.adapters.onebot.v11 import PokeNotifyEvent
|
||||
from nonebot.matcher import Matcher
|
||||
@@ -25,7 +27,35 @@ __plugin_meta__ = PluginMetadata(
|
||||
).to_dict(),
|
||||
)
|
||||
|
||||
TEMP_LIST = []
|
||||
STATS_BUFFER_FLUSH_SIZE = 5000
|
||||
STATS_BUFFER_MAX_RETAIN = 10000
|
||||
TEMP_LIST: list[Statistics] = []
|
||||
_STATS_FLUSH_LOCK = asyncio.Lock()
|
||||
driver = get_driver()
|
||||
|
||||
|
||||
async def _flush_statistics_buffer(reason: str) -> int:
|
||||
async with _STATS_FLUSH_LOCK:
|
||||
call_list = TEMP_LIST.copy()
|
||||
TEMP_LIST.clear()
|
||||
if not call_list:
|
||||
return 0
|
||||
try:
|
||||
await Statistics.bulk_create(call_list)
|
||||
except Exception as e:
|
||||
logger.error(f"{reason}批量添加调用记录失败", "定时任务", e=e)
|
||||
retain_count = max(STATS_BUFFER_MAX_RETAIN - len(TEMP_LIST), 0)
|
||||
if retain_count:
|
||||
TEMP_LIST[:0] = call_list[-retain_count:]
|
||||
return 0
|
||||
logger.debug(f"{reason}批量添加调用记录 {len(call_list)} 条", "定时任务")
|
||||
return len(call_list)
|
||||
|
||||
|
||||
async def _append_statistics(record: Statistics) -> None:
|
||||
TEMP_LIST.append(record)
|
||||
if len(TEMP_LIST) >= STATS_BUFFER_FLUSH_SIZE and not _STATS_FLUSH_LOCK.locked():
|
||||
await _flush_statistics_buffer("缓冲区触发")
|
||||
|
||||
|
||||
@run_postprocessor
|
||||
@@ -50,7 +80,7 @@ async def _(
|
||||
if plugin_type == PluginType.NORMAL:
|
||||
entity = get_entity_ids(session)
|
||||
logger.debug(f"提交调用记录: {matcher.plugin_name}...", session=session)
|
||||
TEMP_LIST.append(
|
||||
await _append_statistics(
|
||||
Statistics(
|
||||
user_id=entity.user_id,
|
||||
group_id=entity.group_id,
|
||||
@@ -66,10 +96,11 @@ async def _():
|
||||
try:
|
||||
if should_pause_tasks():
|
||||
return
|
||||
call_list = TEMP_LIST.copy()
|
||||
TEMP_LIST.clear()
|
||||
if call_list:
|
||||
await Statistics.bulk_create(call_list)
|
||||
logger.debug(f"批量添加调用记录 {len(call_list)} 条", "定时任务")
|
||||
await _flush_statistics_buffer("定时")
|
||||
except Exception as e:
|
||||
logger.error("定时批量添加调用记录", "定时任务", e=e)
|
||||
|
||||
|
||||
@driver.on_shutdown
|
||||
async def _flush_statistics_on_shutdown():
|
||||
await _flush_statistics_buffer("关闭")
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import contextlib
|
||||
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot.rule import to_me
|
||||
@@ -11,8 +13,11 @@ from zhenxun.services.llm.config.providers import get_llm_config
|
||||
from zhenxun.services.llm.manager import clear_model_cache
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
AUTO_RELOAD_JOB_ID = "zhenxun.reload_setting.auto_reload"
|
||||
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="重载配置",
|
||||
description="重新加载config.yaml",
|
||||
@@ -53,22 +58,75 @@ _matcher = on_alconna(
|
||||
)
|
||||
|
||||
|
||||
@_matcher.handle()
|
||||
async def _(session: EventSession, arparma: Arparma):
|
||||
def _get_auto_reload_interval() -> int:
|
||||
value = Config.get_config("reload_setting", "AUTO_RELOAD_TIME", 180)
|
||||
try:
|
||||
seconds = int(value)
|
||||
except (TypeError, ValueError):
|
||||
logger.warning(
|
||||
f"AUTO_RELOAD_TIME 配置无效: {value!r},已使用默认值 180 秒",
|
||||
"重载配置",
|
||||
)
|
||||
return 180
|
||||
if seconds <= 0:
|
||||
logger.warning(
|
||||
f"AUTO_RELOAD_TIME 配置小于等于 0: {seconds},已使用默认值 180 秒",
|
||||
"重载配置",
|
||||
)
|
||||
return 180
|
||||
return seconds
|
||||
|
||||
|
||||
def _reschedule_auto_reload_job() -> None:
|
||||
seconds = _get_auto_reload_interval()
|
||||
if scheduler.get_job(AUTO_RELOAD_JOB_ID):
|
||||
scheduler.reschedule_job(
|
||||
AUTO_RELOAD_JOB_ID,
|
||||
trigger="interval",
|
||||
seconds=seconds,
|
||||
)
|
||||
else:
|
||||
scheduler.add_job(
|
||||
_auto_reload_config,
|
||||
"interval",
|
||||
seconds=seconds,
|
||||
id=AUTO_RELOAD_JOB_ID,
|
||||
replace_existing=True,
|
||||
)
|
||||
logger.debug(f"自动重载配置任务间隔已设置为 {seconds} 秒", "重载配置")
|
||||
|
||||
|
||||
async def _reload_plugin_limit_config() -> None:
|
||||
from zhenxun.builtin_plugins.hooks.auth.auth_limit import LimitManager
|
||||
from zhenxun.builtin_plugins.init.manager import manager
|
||||
|
||||
manager.init()
|
||||
await manager.load_to_db()
|
||||
await LimitManager.update_limits()
|
||||
|
||||
|
||||
async def _reload_runtime_config() -> None:
|
||||
Config.reload()
|
||||
get_llm_config.cache_clear()
|
||||
clear_model_cache()
|
||||
await _reload_plugin_limit_config()
|
||||
with contextlib.suppress(Exception):
|
||||
_reschedule_auto_reload_job()
|
||||
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=1)
|
||||
def _init_auto_reload_job() -> None:
|
||||
_reschedule_auto_reload_job()
|
||||
|
||||
|
||||
@_matcher.handle()
|
||||
async def _(session: EventSession, arparma: Arparma):
|
||||
await _reload_runtime_config()
|
||||
logger.debug("自动重载配置文件", arparma.header_result, session=session)
|
||||
await MessageUtils.build_message("重载完成!").send(reply_to=True)
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
"interval",
|
||||
seconds=Config.get_config("reload_setting", "AUTO_RELOAD_TIME", 180),
|
||||
)
|
||||
async def _():
|
||||
async def _auto_reload_config() -> None:
|
||||
if Config.get_config("reload_setting", "AUTO_RELOAD"):
|
||||
Config.reload()
|
||||
get_llm_config.cache_clear()
|
||||
clear_model_cache()
|
||||
await _reload_runtime_config()
|
||||
logger.debug("已自动重载配置文件...")
|
||||
|
||||
@@ -1,20 +1,17 @@
|
||||
import asyncio
|
||||
import secrets
|
||||
|
||||
from fastapi import APIRouter, FastAPI
|
||||
import nonebot
|
||||
from nonebot.log import default_filter, default_format
|
||||
from nonebot.plugin import PluginMetadata
|
||||
|
||||
from zhenxun.configs.config import Config as gConfig
|
||||
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
|
||||
from zhenxun.services.log import logger, logger_
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
|
||||
from .api.configure import router as configure_router
|
||||
from .api.logs import router as ws_log_routes
|
||||
from .api.logs.log_manager import LOG_STORAGE
|
||||
from .api.menu import router as menu_router
|
||||
from .api.tabs.dashboard import router as dashboard_router
|
||||
from .api.tabs.database import router as database_router
|
||||
@@ -95,25 +92,6 @@ WsApiRouter.include_router(chat_routes)
|
||||
@PriorityLifecycle.on_startup(priority=0)
|
||||
async def _():
|
||||
try:
|
||||
# 存储任务引用的列表,防止任务被垃圾回收
|
||||
_tasks = []
|
||||
|
||||
async def log_sink(message: str):
|
||||
loop = None
|
||||
if not loop:
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except Exception as e:
|
||||
logger.warning("Web Ui log_sink", e=e)
|
||||
if not loop:
|
||||
loop = asyncio.new_event_loop()
|
||||
# 存储任务引用到外部列表中
|
||||
_tasks.append(loop.create_task(LOG_STORAGE.add(message.rstrip("\n"))))
|
||||
|
||||
logger_.add(
|
||||
log_sink, colorize=True, filter=default_filter, format=default_format
|
||||
)
|
||||
|
||||
app: FastAPI = nonebot.get_app()
|
||||
app.include_router(BaseApiRouter)
|
||||
app.include_router(WsApiRouter)
|
||||
|
||||
@@ -1,33 +1,97 @@
|
||||
import asyncio
|
||||
from collections import deque
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Generic, TypeVar
|
||||
import contextlib
|
||||
|
||||
_T = TypeVar("_T")
|
||||
LogListener = Callable[[_T], Awaitable[None]]
|
||||
from nonebot.log import default_filter, default_format
|
||||
|
||||
from zhenxun.services.log import logger_
|
||||
|
||||
LogListener = Callable[[str], Awaitable[None]]
|
||||
DEFAULT_MAX_LOGS = 1000
|
||||
DEFAULT_MAX_LISTENERS = 16
|
||||
|
||||
|
||||
class LogStorage(Generic[_T]):
|
||||
class LogStorage:
|
||||
"""
|
||||
日志存储
|
||||
"""
|
||||
|
||||
def __init__(self, rotation: float = 5 * 60):
|
||||
def __init__(
|
||||
self,
|
||||
rotation: float = 5 * 60,
|
||||
max_logs: int = DEFAULT_MAX_LOGS,
|
||||
max_listeners: int = DEFAULT_MAX_LISTENERS,
|
||||
):
|
||||
self.count, self.rotation = 0, rotation
|
||||
self.max_logs = max_logs
|
||||
self.max_listeners = max_listeners
|
||||
self.logs: dict[int, str] = {}
|
||||
self.listeners: set[LogListener[str]] = set()
|
||||
self._order: deque[int] = deque()
|
||||
self.listeners: set[LogListener] = set()
|
||||
|
||||
async def add(self, log: str):
|
||||
seq = self.count = self.count + 1
|
||||
self.logs[seq] = log
|
||||
self._order.append(seq)
|
||||
self._trim()
|
||||
asyncio.get_running_loop().call_later(self.rotation, self.remove, seq)
|
||||
await asyncio.gather(
|
||||
*(listener(log) for listener in self.listeners),
|
||||
return_exceptions=True,
|
||||
)
|
||||
listeners = tuple(self.listeners)
|
||||
if listeners:
|
||||
results = await asyncio.gather(
|
||||
*(listener(log) for listener in listeners),
|
||||
return_exceptions=True,
|
||||
)
|
||||
for listener, result in zip(listeners, results, strict=False):
|
||||
if isinstance(result, BaseException):
|
||||
self.listeners.discard(listener)
|
||||
return seq
|
||||
|
||||
def add_listener(self, listener: LogListener) -> bool:
|
||||
if len(self.listeners) >= self.max_listeners:
|
||||
return False
|
||||
self.listeners.add(listener)
|
||||
return True
|
||||
|
||||
def remove_listener(self, listener: LogListener) -> None:
|
||||
self.listeners.discard(listener)
|
||||
|
||||
def remove(self, seq: int):
|
||||
del self.logs[seq]
|
||||
self.logs.pop(seq, None)
|
||||
with contextlib.suppress(ValueError):
|
||||
self._order.remove(seq)
|
||||
|
||||
def _trim(self) -> None:
|
||||
while self._order and self._order[0] not in self.logs:
|
||||
self._order.popleft()
|
||||
while len(self.logs) > self.max_logs and self._order:
|
||||
self.logs.pop(self._order.popleft(), None)
|
||||
|
||||
|
||||
LOG_STORAGE: LogStorage[str] = LogStorage[str]()
|
||||
LOG_STORAGE = LogStorage()
|
||||
|
||||
_LOG_SINK_ID: int | None = None
|
||||
|
||||
|
||||
async def ensure_log_sink_started() -> None:
|
||||
global _LOG_SINK_ID
|
||||
if _LOG_SINK_ID is not None:
|
||||
return
|
||||
|
||||
async def log_sink(message: str) -> None:
|
||||
await LOG_STORAGE.add(message.rstrip("\n"))
|
||||
|
||||
_LOG_SINK_ID = logger_.add(
|
||||
log_sink,
|
||||
colorize=True,
|
||||
filter=default_filter,
|
||||
format=default_format,
|
||||
)
|
||||
|
||||
|
||||
def stop_log_sink_if_idle() -> None:
|
||||
global _LOG_SINK_ID
|
||||
if LOG_STORAGE.listeners or _LOG_SINK_ID is None:
|
||||
return
|
||||
logger_.remove(_LOG_SINK_ID)
|
||||
_LOG_SINK_ID = None
|
||||
|
||||
@@ -3,7 +3,7 @@ from loguru import logger
|
||||
from nonebot.utils import escape_tag
|
||||
from starlette.websockets import WebSocket, WebSocketDisconnect, WebSocketState
|
||||
|
||||
from .log_manager import LOG_STORAGE
|
||||
from .log_manager import LOG_STORAGE, ensure_log_sink_started, stop_log_sink_if_idle
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -11,11 +11,16 @@ router = APIRouter()
|
||||
@router.websocket("/logs")
|
||||
async def system_logs_realtime(websocket: WebSocket):
|
||||
await websocket.accept()
|
||||
await ensure_log_sink_started()
|
||||
|
||||
async def log_listener(log: str):
|
||||
await websocket.send_text(log)
|
||||
|
||||
LOG_STORAGE.listeners.add(log_listener)
|
||||
if not LOG_STORAGE.add_listener(log_listener):
|
||||
await websocket.send_text("日志连接数已达上限,请稍后再试。")
|
||||
await websocket.close()
|
||||
stop_log_sink_if_idle()
|
||||
return
|
||||
try:
|
||||
while websocket.client_state == WebSocketState.CONNECTED:
|
||||
recv = await websocket.receive()
|
||||
@@ -26,4 +31,5 @@ async def system_logs_realtime(websocket: WebSocket):
|
||||
except WebSocketDisconnect:
|
||||
pass
|
||||
finally:
|
||||
LOG_STORAGE.listeners.remove(log_listener)
|
||||
LOG_STORAGE.remove_listener(log_listener)
|
||||
stop_log_sink_if_idle()
|
||||
|
||||
@@ -34,6 +34,31 @@ run_time = time.time()
|
||||
|
||||
ws_router = APIRouter()
|
||||
router = APIRouter(prefix="/main")
|
||||
_SYSTEM_STATUS_CONNECTIONS: set[WebSocket] = set()
|
||||
_SYSTEM_STATUS_STOPPING = False
|
||||
|
||||
|
||||
async def _close_system_status_websocket(websocket: WebSocket) -> None:
|
||||
with contextlib.suppress(Exception):
|
||||
if websocket.client_state == WebSocketState.CONNECTED:
|
||||
await asyncio.wait_for(
|
||||
websocket.close(code=1001, reason="server shutdown"),
|
||||
timeout=2,
|
||||
)
|
||||
|
||||
|
||||
@driver.on_shutdown
|
||||
async def _close_system_status_websockets() -> None:
|
||||
global _SYSTEM_STATUS_STOPPING
|
||||
_SYSTEM_STATUS_STOPPING = True
|
||||
websockets = list(_SYSTEM_STATUS_CONNECTIONS)
|
||||
if not websockets:
|
||||
return
|
||||
await asyncio.gather(
|
||||
*(_close_system_status_websocket(websocket) for websocket in websockets),
|
||||
return_exceptions=True,
|
||||
)
|
||||
_SYSTEM_STATUS_CONNECTIONS.clear()
|
||||
|
||||
|
||||
@router.get(
|
||||
@@ -243,11 +268,39 @@ async def _(param: BotManageUpdateParam):
|
||||
@ws_router.websocket("/system_status")
|
||||
async def system_logs_realtime(websocket: WebSocket, sleep: int = 5):
|
||||
await websocket.accept()
|
||||
_SYSTEM_STATUS_CONNECTIONS.add(websocket)
|
||||
logger.debug("ws system_status is connect")
|
||||
with contextlib.suppress(
|
||||
WebSocketDisconnect, ConnectionClosedError, ConnectionClosedOK
|
||||
):
|
||||
while websocket.client_state == WebSocketState.CONNECTED:
|
||||
|
||||
disconnect_event = asyncio.Event()
|
||||
|
||||
async def _watch_disconnect() -> None:
|
||||
try:
|
||||
while websocket.client_state == WebSocketState.CONNECTED:
|
||||
await websocket.receive()
|
||||
except (WebSocketDisconnect, ConnectionClosedError, ConnectionClosedOK):
|
||||
pass
|
||||
except Exception as e:
|
||||
logger.debug(f"ws system_status receive stopped: {type(e).__name__}")
|
||||
finally:
|
||||
disconnect_event.set()
|
||||
|
||||
receive_task = asyncio.create_task(_watch_disconnect())
|
||||
try:
|
||||
while (
|
||||
websocket.client_state == WebSocketState.CONNECTED
|
||||
and not _SYSTEM_STATUS_STOPPING
|
||||
):
|
||||
system_status = await get_system_status()
|
||||
await websocket.send_text(system_status.json())
|
||||
await asyncio.sleep(sleep)
|
||||
await asyncio.wait_for(websocket.send_text(system_status.json()), timeout=5)
|
||||
try:
|
||||
await asyncio.wait_for(disconnect_event.wait(), timeout=max(sleep, 1))
|
||||
except TimeoutError:
|
||||
pass
|
||||
except (WebSocketDisconnect, ConnectionClosedError, ConnectionClosedOK):
|
||||
pass
|
||||
finally:
|
||||
_SYSTEM_STATUS_CONNECTIONS.discard(websocket)
|
||||
receive_task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await receive_task
|
||||
await _close_system_status_websocket(websocket)
|
||||
|
||||
@@ -52,34 +52,20 @@ async def _(
|
||||
async def _() -> Result[PluginCount]:
|
||||
try:
|
||||
plugin_count = PluginCount()
|
||||
plugin_count.normal = len(
|
||||
await DbPluginInfo.get_plugins(
|
||||
plugin_type=PluginType.NORMAL,
|
||||
load_status=True,
|
||||
filter_parent=False,
|
||||
)
|
||||
)
|
||||
plugin_count.admin = len(
|
||||
await DbPluginInfo.get_plugins(
|
||||
plugin_type__in=[PluginType.ADMIN, PluginType.SUPER_AND_ADMIN],
|
||||
load_status=True,
|
||||
filter_parent=False,
|
||||
)
|
||||
)
|
||||
plugin_count.superuser = len(
|
||||
await DbPluginInfo.get_plugins(
|
||||
plugin_type__in=[PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN],
|
||||
load_status=True,
|
||||
filter_parent=False,
|
||||
)
|
||||
)
|
||||
plugin_count.other = len(
|
||||
await DbPluginInfo.get_plugins(
|
||||
plugin_type__in=[PluginType.HIDDEN, PluginType.DEPENDANT],
|
||||
load_status=True,
|
||||
filter_parent=False,
|
||||
)
|
||||
plugins = await DbPluginInfo.get_plugins(
|
||||
load_status=True,
|
||||
filter_parent=False,
|
||||
)
|
||||
for plugin in plugins:
|
||||
plugin_type = plugin.plugin_type
|
||||
if plugin_type == PluginType.NORMAL:
|
||||
plugin_count.normal += 1
|
||||
if plugin_type in {PluginType.ADMIN, PluginType.SUPER_AND_ADMIN}:
|
||||
plugin_count.admin += 1
|
||||
if plugin_type in {PluginType.SUPERUSER, PluginType.SUPER_AND_ADMIN}:
|
||||
plugin_count.superuser += 1
|
||||
if plugin_type in {PluginType.HIDDEN, PluginType.DEPENDANT}:
|
||||
plugin_count.other += 1
|
||||
return Result.ok(plugin_count, "拿到信息啦!")
|
||||
except Exception as e:
|
||||
logger.error(f"{router.prefix}/get_plugin_count 调用错误", "WebUi", e=e)
|
||||
|
||||
@@ -111,10 +111,19 @@ class ApiDataSource:
|
||||
other_update_fields = set()
|
||||
updated_count = 0
|
||||
errors = []
|
||||
modules = [item.module for item in params.updates]
|
||||
plugin_records = await DbPluginInfo.get_plugins(
|
||||
module__in=modules,
|
||||
load_status=None,
|
||||
filter_parent=False,
|
||||
)
|
||||
plugin_map = {plugin.module: plugin for plugin in plugin_records}
|
||||
|
||||
for item in params.updates:
|
||||
try:
|
||||
db_plugin = await DbPluginInfo.get(module=item.module)
|
||||
db_plugin = plugin_map.get(item.module)
|
||||
if db_plugin is None:
|
||||
raise DoesNotExist()
|
||||
plugin_changed_other = False
|
||||
plugin_changed_block = False
|
||||
|
||||
|
||||
Reference in New Issue
Block a user