mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-28 16:20:56 +08:00
bugfix:更换数据库初始化超时路径以修复连接超时问题 (#2137)
* bugfix:更换数据库初始化超时路径以修复连接超时问题 * 移除部分观测链路 * 细节修改 * 权限检查细节修改2 * bugfix:修复金币懒加载造成插件金币消耗不了的问题 * bugfix:整理鉴权逻辑 * 完善缓存系统 * 优化官端使用 * bugfix:增加官机和频道的前导统一剥离,新增获取群列表自愈 * bugfix:修复导入问题
This commit is contained in:
@@ -109,7 +109,7 @@ class MemberUpdateManage:
|
||||
members = await interface.get_members(SceneType.GROUP, group_scene.id)
|
||||
|
||||
try:
|
||||
group_console, _ = await GroupConsole.get_or_create(
|
||||
group_console, _ = await GroupConsole.get_or_create_root_group(
|
||||
group_id=group_id, defaults={"platform": platform}
|
||||
)
|
||||
group_console.member_count = len(members)
|
||||
|
||||
@@ -39,6 +39,11 @@ _matcher = on_alconna(
|
||||
async def _(bot: Bot, session: EventSession, arparma: Arparma):
|
||||
logger.info("更新群组信息", arparma.header_result, session=session)
|
||||
try:
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
await MessageUtils.build_message(
|
||||
"当前平台不支持旧群组信息同步,仅 OneBot 协议端可用。"
|
||||
).send(reply_to=True)
|
||||
return
|
||||
await PlatformUtils.update_group(bot)
|
||||
await MessageUtils.build_message("已经成功更新了群组信息!").send(reply_to=True)
|
||||
except Exception:
|
||||
|
||||
@@ -87,7 +87,7 @@ class PluginManager:
|
||||
|
||||
for gid in groups_to_open | groups_to_close:
|
||||
platform = bot.adapter.get_name() if bot else "qq"
|
||||
await GroupConsole.get_or_create(
|
||||
await GroupConsole.get_or_create_root_group(
|
||||
group_id=gid, defaults={"platform": platform}
|
||||
)
|
||||
|
||||
|
||||
@@ -173,10 +173,12 @@ class TaskStrategy(SwitchStrategy):
|
||||
|
||||
async def set_all_default_status(self, status: bool) -> None:
|
||||
await TaskInfo.all().update(default_status=status)
|
||||
# Bulk updates bypass model save hooks; keep TaskInfoMemoryCache in sync.
|
||||
await self.refresh_cache()
|
||||
|
||||
async def set_all_global_status(self, status: bool) -> None:
|
||||
await TaskInfo.all().update(status=status)
|
||||
# Bulk updates bypass model save hooks; keep TaskInfoMemoryCache in sync.
|
||||
await self.refresh_cache()
|
||||
|
||||
async def refresh_cache(self) -> None:
|
||||
|
||||
@@ -4,6 +4,7 @@ from nonebot.adapters import Bot
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
Config.add_plugin_config(
|
||||
"catchphrase",
|
||||
@@ -16,6 +17,8 @@ Config.add_plugin_config(
|
||||
|
||||
@Bot.on_calling_api
|
||||
async def handle_api_call(bot: Bot, api: str, data: dict[str, Any]):
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
return
|
||||
if api == "send_msg":
|
||||
catchphrase = Config.get_config("catchphrase", "CATCHPHRASE")
|
||||
if catchphrase and (message := data.get("message")):
|
||||
|
||||
@@ -49,14 +49,4 @@ Config.add_plugin_config(
|
||||
type=bool,
|
||||
)
|
||||
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
"RECORD_BOT_SENT_MESSAGES",
|
||||
True,
|
||||
help="记录bot消息发送",
|
||||
default_value=True,
|
||||
type=bool,
|
||||
)
|
||||
|
||||
|
||||
nonebot.load_plugins(str(Path(__file__).parent.resolve()))
|
||||
|
||||
@@ -5,12 +5,12 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
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 EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER, LevelUserSnapshot
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
@@ -50,7 +50,10 @@ async def auth_admin(
|
||||
if cached_levels is not None:
|
||||
global_user, group_users = cached_levels
|
||||
else:
|
||||
global_user, group_users = await LevelUserMemoryCache.get_levels(
|
||||
(
|
||||
global_user,
|
||||
group_users,
|
||||
) = await DEFAULT_PERMISSION_DATA_PROVIDER.get_admin_levels(
|
||||
entity.user_id, entity.group_id
|
||||
)
|
||||
|
||||
|
||||
@@ -7,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.runtime_cache import BanMemoryCache
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
@@ -15,6 +14,7 @@ from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
from .exception import SkipPluginException
|
||||
from .utils import freq
|
||||
|
||||
@@ -60,9 +60,10 @@ async def is_ban(user_id: str | None, group_id: str | None) -> int:
|
||||
"""
|
||||
if not user_id and not group_id:
|
||||
return 0
|
||||
if not BanMemoryCache.is_loaded():
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
if not provider.ban_cache_loaded():
|
||||
return 0
|
||||
return BanMemoryCache.remaining_time(user_id, group_id)
|
||||
return provider.get_ban_remaining_time(user_id, group_id)
|
||||
|
||||
|
||||
def check_plugin_type(matcher: Matcher) -> bool:
|
||||
|
||||
@@ -2,12 +2,12 @@ import time
|
||||
|
||||
from zhenxun.models.bot_console import BotConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.cache.runtime_cache import BotMemoryCache, BotSnapshot
|
||||
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 .data_provider import DEFAULT_PERMISSION_DATA_PROVIDER, BotSnapshot
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
@@ -33,12 +33,13 @@ async def auth_bot(
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
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)
|
||||
bot = await provider.get_bot(bot_id)
|
||||
|
||||
if bot is None:
|
||||
raise SkipPluginException("Bot不存在,阻断权限检测...")
|
||||
|
||||
@@ -10,10 +10,6 @@ from pydantic import BaseModel
|
||||
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.plugin_limit import PluginLimit
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
PluginLimitMemoryCache,
|
||||
PluginLimitSnapshot,
|
||||
)
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import LimitWatchType, PluginLimitType
|
||||
@@ -25,6 +21,10 @@ from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .context import PermissionContext
|
||||
from .data_provider import (
|
||||
DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
PluginLimitSnapshot,
|
||||
)
|
||||
from .exception import SkipPluginException
|
||||
|
||||
driver = nonebot.get_driver()
|
||||
@@ -106,11 +106,10 @@ class LimitManager:
|
||||
block_limit: ClassVar[dict[str, Limit]] = {}
|
||||
count_limit: ClassVar[dict[str, Limit]] = {}
|
||||
|
||||
# 模块限制缓存,避免频繁查询数据库
|
||||
module_limit_cache: ClassVar[
|
||||
dict[str, tuple[float, list[PluginLimitSnapshot], bool]]
|
||||
# 只缓存异常短路结果;正常 limit 列表统一从 PluginLimitMemoryCache 读取。
|
||||
module_limit_error_cache: ClassVar[
|
||||
dict[str, tuple[float, list[PluginLimitSnapshot]]]
|
||||
] = {}
|
||||
module_cache_ttl: ClassVar[float] = 60 # 模块缓存有效期(秒)
|
||||
module_cache_error_ttl: ClassVar[float] = 5 # 超时缓存有效期(秒)
|
||||
|
||||
@classmethod
|
||||
@@ -132,15 +131,16 @@ class LimitManager:
|
||||
cls.is_updating = True
|
||||
try:
|
||||
start_time = time.time()
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
limit_list = PluginLimitMemoryCache.get_all_limits()
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
await provider.ensure_module_limits_loaded()
|
||||
limit_list = await provider.get_all_module_limits()
|
||||
|
||||
# 清空旧数据
|
||||
cls.add_module = []
|
||||
cls.cd_limit = {}
|
||||
cls.block_limit = {}
|
||||
cls.count_limit = {}
|
||||
cls.module_limit_cache.clear()
|
||||
cls.module_limit_error_cache.clear()
|
||||
# 添加新数据
|
||||
for limit in limit_list:
|
||||
cls.add_limit(limit)
|
||||
@@ -216,22 +216,21 @@ class LimitManager:
|
||||
"""
|
||||
current_time = time.time()
|
||||
|
||||
# 检查缓存
|
||||
if module in cls.module_limit_cache:
|
||||
cache_time, limits, is_error = cls.module_limit_cache[module]
|
||||
ttl = cls.module_cache_error_ttl if is_error else cls.module_cache_ttl
|
||||
if current_time - cache_time < ttl:
|
||||
# 正常路径不再二次缓存列表,避免与 PluginLimitMemoryCache 形成双真源。
|
||||
if module in cls.module_limit_error_cache:
|
||||
cache_time, limits = cls.module_limit_error_cache[module]
|
||||
if current_time - cache_time < cls.module_cache_error_ttl:
|
||||
return limits
|
||||
cls.module_limit_error_cache.pop(module, None)
|
||||
|
||||
# 缓存不存在或已过期,从内存缓存获取
|
||||
try:
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
limits = await PluginLimitMemoryCache.get_limits(module)
|
||||
cls.module_limit_cache[module] = (current_time, limits, False)
|
||||
return limits
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
await provider.ensure_module_limits_loaded()
|
||||
return await provider.get_module_limits(module)
|
||||
except Exception as exc:
|
||||
logger.error(f"get module limits failed: {module}", LOGGER_COMMAND, e=exc)
|
||||
cls.module_limit_cache[module] = (current_time, [], True)
|
||||
cls.module_limit_error_cache[module] = (current_time, [])
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -39,6 +39,7 @@ if TYPE_CHECKING:
|
||||
class EventContext:
|
||||
bot_id: str
|
||||
platform: str
|
||||
platform_scope: str
|
||||
event_type: str
|
||||
message_id: str | int | None
|
||||
entity: EntityIDs
|
||||
@@ -170,6 +171,7 @@ def event_cache_key(
|
||||
*,
|
||||
bot_id: str,
|
||||
platform: str,
|
||||
platform_scope: str | None = None,
|
||||
entity: EntityIDs,
|
||||
) -> str:
|
||||
msg_id = _event_message_id(event)
|
||||
@@ -177,7 +179,11 @@ def event_cache_key(
|
||||
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}"
|
||||
scope = platform_scope or platform
|
||||
return (
|
||||
f"{scope}:{platform}:{bot_id}:{entity.user_id}:"
|
||||
f"{group_id}:{channel_id}:{msg_id}"
|
||||
)
|
||||
|
||||
|
||||
def get_event_cache(
|
||||
@@ -185,11 +191,18 @@ def get_event_cache(
|
||||
*,
|
||||
bot_id: str,
|
||||
platform: str,
|
||||
platform_scope: str | None = None,
|
||||
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)
|
||||
key = event_cache_key(
|
||||
event,
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
platform_scope=platform_scope,
|
||||
entity=entity,
|
||||
)
|
||||
try:
|
||||
return EVENT_CACHE[key]
|
||||
except KeyError:
|
||||
@@ -255,6 +268,7 @@ def get_or_create_event_context(
|
||||
entity = resolve_entity_ids(event, session)
|
||||
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
platform_scope = PlatformUtils.get_platform_scope(session)
|
||||
bot_id = str(bot.self_id)
|
||||
event_cache = state.get(STATE_EVENT_CACHE)
|
||||
if not isinstance(event_cache, dict):
|
||||
@@ -262,6 +276,7 @@ def get_or_create_event_context(
|
||||
event,
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
platform_scope=platform_scope,
|
||||
entity=entity,
|
||||
)
|
||||
|
||||
@@ -292,6 +307,7 @@ def get_or_create_event_context(
|
||||
context = EventContext(
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
platform_scope=platform_scope,
|
||||
event_type=event.get_type(),
|
||||
message_id=_event_message_id(event),
|
||||
entity=entity,
|
||||
|
||||
@@ -0,0 +1,151 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
BanMemoryCache,
|
||||
BotMemoryCache,
|
||||
BotSnapshot,
|
||||
GroupMemoryCache,
|
||||
GroupSnapshot,
|
||||
LevelUserMemoryCache,
|
||||
LevelUserSnapshot,
|
||||
PluginLimitMemoryCache,
|
||||
PluginLimitSnapshot,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
|
||||
|
||||
AdminLevels = tuple[LevelUserSnapshot | None, LevelUserSnapshot | None]
|
||||
|
||||
|
||||
class PermissionDataProvider:
|
||||
"""Auth data facade over runtime caches.
|
||||
|
||||
Permission checks should read stable runtime snapshots through this provider
|
||||
instead of reaching into individual cache classes from multiple auth modules.
|
||||
The provider does not own policy semantics and does not query the database
|
||||
directly.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def plugin_cache_loaded() -> bool:
|
||||
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
|
||||
|
||||
return PluginInfoMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_plugin_if_ready(module: str) -> "PluginInfo | None":
|
||||
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
|
||||
|
||||
return PluginInfoMemoryCache.get_by_module_if_ready(module)
|
||||
|
||||
@staticmethod
|
||||
async def get_plugin(module: str) -> "PluginInfo | None":
|
||||
from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache
|
||||
|
||||
return await PluginInfoMemoryCache.get_by_module(module)
|
||||
|
||||
@staticmethod
|
||||
def module_limit_cache_loaded() -> bool:
|
||||
return PluginLimitMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
async def ensure_module_limits_loaded() -> None:
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_module_limits_if_ready(
|
||||
module: str,
|
||||
) -> list[PluginLimitSnapshot] | None:
|
||||
return PluginLimitMemoryCache.get_limits_if_ready(module)
|
||||
|
||||
@staticmethod
|
||||
async def get_module_limits(module: str) -> list[PluginLimitSnapshot]:
|
||||
return await PluginLimitMemoryCache.get_limits(module)
|
||||
|
||||
@staticmethod
|
||||
async def get_all_module_limits() -> list[PluginLimitSnapshot]:
|
||||
if not PluginLimitMemoryCache.is_loaded():
|
||||
await PluginLimitMemoryCache.ensure_loaded()
|
||||
return PluginLimitMemoryCache.get_all_limits()
|
||||
|
||||
@staticmethod
|
||||
def bot_cache_loaded() -> bool:
|
||||
return BotMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_bot_if_ready(bot_id: str | None) -> BotSnapshot | None:
|
||||
return BotMemoryCache.get_if_ready(bot_id)
|
||||
|
||||
@staticmethod
|
||||
async def get_bot(bot_id: str | None) -> BotSnapshot | None:
|
||||
return await BotMemoryCache.get(bot_id)
|
||||
|
||||
@staticmethod
|
||||
def group_cache_loaded() -> bool:
|
||||
return GroupMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_group_if_ready(
|
||||
group_id: str | None,
|
||||
channel_id: str | None = None,
|
||||
) -> GroupSnapshot | None:
|
||||
return GroupMemoryCache.get_if_ready(group_id, channel_id)
|
||||
|
||||
@staticmethod
|
||||
async def get_group(
|
||||
group_id: str | None,
|
||||
channel_id: str | None = None,
|
||||
) -> GroupSnapshot | None:
|
||||
return await GroupMemoryCache.get(group_id, channel_id)
|
||||
|
||||
@staticmethod
|
||||
def admin_cache_loaded() -> bool:
|
||||
return LevelUserMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
def get_admin_levels_if_ready(
|
||||
user_id: str | None,
|
||||
group_id: str | None,
|
||||
) -> AdminLevels | None:
|
||||
return LevelUserMemoryCache.get_levels_if_ready(user_id, group_id)
|
||||
|
||||
@staticmethod
|
||||
async def get_admin_levels(
|
||||
user_id: str | None,
|
||||
group_id: str | None,
|
||||
) -> AdminLevels:
|
||||
return await LevelUserMemoryCache.get_levels(user_id, group_id)
|
||||
|
||||
@staticmethod
|
||||
def ban_cache_loaded() -> bool:
|
||||
return BanMemoryCache.is_loaded()
|
||||
|
||||
@staticmethod
|
||||
async def ensure_ban_loaded() -> None:
|
||||
await BanMemoryCache.ensure_loaded()
|
||||
|
||||
@staticmethod
|
||||
def is_banned(user_id: str | None, group_id: str | None) -> bool:
|
||||
return BanMemoryCache.is_banned(user_id, group_id)
|
||||
|
||||
@staticmethod
|
||||
def get_ban_remaining_time(user_id: str | None, group_id: str | None) -> int:
|
||||
return BanMemoryCache.remaining_time(user_id, group_id)
|
||||
|
||||
|
||||
DEFAULT_PERMISSION_DATA_PROVIDER = PermissionDataProvider()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_PERMISSION_DATA_PROVIDER",
|
||||
"AdminLevels",
|
||||
"BotSnapshot",
|
||||
"GroupSnapshot",
|
||||
"LevelUserSnapshot",
|
||||
"PermissionDataProvider",
|
||||
"PluginLimitSnapshot",
|
||||
]
|
||||
@@ -8,6 +8,7 @@ import re
|
||||
from typing import Any, Literal
|
||||
import weakref
|
||||
|
||||
from loguru import logger
|
||||
from nonebot.matcher import Matcher
|
||||
|
||||
ActivationDecision = Literal["match", "miss", "unknown"]
|
||||
@@ -23,6 +24,20 @@ ActivationLane = Literal[
|
||||
"passive_render",
|
||||
]
|
||||
|
||||
KNOWN_SAFE_RULE_NAMES = frozenset(
|
||||
{
|
||||
"CommandRule",
|
||||
"ShellCommandRule",
|
||||
"RegexRule",
|
||||
"StartswithRule",
|
||||
"EndswithRule",
|
||||
"FullmatchRule",
|
||||
"KeywordsRule",
|
||||
"IsTypeRule",
|
||||
"ToMeRule",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ActivationRuleDescriptor:
|
||||
@@ -48,9 +63,96 @@ class HandlerDescriptor:
|
||||
has_custom_rule: bool = False
|
||||
commands: tuple[str, ...] = ()
|
||||
shortcuts: tuple[str, ...] | None = None
|
||||
alconna: tuple[AlconnaDescriptor, ...] = ()
|
||||
rules: tuple[ActivationRuleDescriptor, ...] = ()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AlconnaShortcutDescriptor:
|
||||
pattern: str
|
||||
fuzzy: bool = False
|
||||
flags: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AlconnaDescriptor:
|
||||
command: str = ""
|
||||
aliases: tuple[str, ...] = ()
|
||||
prefixes: tuple[str, ...] = ()
|
||||
shortcuts: tuple[AlconnaShortcutDescriptor, ...] = ()
|
||||
compact: bool = False
|
||||
skip_for_unmatch: bool = True
|
||||
before_rule_count: int = 0
|
||||
after_rule_count: int = 0
|
||||
before_rule_known_safe: bool = True
|
||||
after_rule_known_safe: bool = True
|
||||
input_rewrite_extensions: tuple[str, ...] = ()
|
||||
|
||||
@property
|
||||
def has_reply_merge_extension(self) -> bool:
|
||||
return "ReplyMergeExtension" in self.input_rewrite_extensions
|
||||
|
||||
@property
|
||||
def regex_command(self) -> bool:
|
||||
return self.command.startswith("re:")
|
||||
|
||||
|
||||
class AlconnaActivationIndex:
|
||||
"""Safe, metadata-only Alconna prefilter.
|
||||
|
||||
This index never executes Alconna.parse() or Matcher.check_rule(). It only
|
||||
skips matchers when the command head/shortcut is known to miss. Anything
|
||||
custom or ambiguous stays fail-open to preserve NoneBot compatibility.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._safe = 0
|
||||
self._unknown = 0
|
||||
|
||||
@property
|
||||
def safe_count(self) -> int:
|
||||
return self._safe
|
||||
|
||||
@property
|
||||
def unknown_count(self) -> int:
|
||||
return self._unknown
|
||||
|
||||
def rebuild(self, descriptors: Iterable[HandlerDescriptor]) -> None:
|
||||
safe = 0
|
||||
unknown = 0
|
||||
for descriptor in descriptors:
|
||||
if descriptor.alconna:
|
||||
if any(_alconna_can_prefilter(item) for item in descriptor.alconna):
|
||||
safe += 1
|
||||
else:
|
||||
unknown += 1
|
||||
self._safe = safe
|
||||
self._unknown = unknown
|
||||
if logger.level("DEBUG"):
|
||||
logger.debug(
|
||||
"alconna activation index rebuilt: safe={}, unknown={}",
|
||||
safe,
|
||||
unknown,
|
||||
)
|
||||
|
||||
def select(
|
||||
self,
|
||||
descriptor: HandlerDescriptor,
|
||||
context: ActivationContext,
|
||||
texts: tuple[str, ...],
|
||||
) -> ActivationDecision:
|
||||
if not descriptor.alconna:
|
||||
return "unknown"
|
||||
saw_unknown = False
|
||||
for alconna in descriptor.alconna:
|
||||
decision = matcher_alconna_head_matches(alconna, texts, context)
|
||||
if decision == "match":
|
||||
return "match"
|
||||
if decision == "unknown":
|
||||
saw_unknown = True
|
||||
return "unknown" if saw_unknown else "miss"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ActivationContext:
|
||||
event_type: str
|
||||
@@ -69,37 +171,25 @@ class ActivationContext:
|
||||
@dataclass(slots=True)
|
||||
class ActivationResult:
|
||||
selected: list[type[Matcher]]
|
||||
fallback_required: bool = False
|
||||
selected_by_lane: dict[str, int] = field(default_factory=dict)
|
||||
skipped_by_lane: dict[str, int] = field(default_factory=dict)
|
||||
selected_by_reason: dict[str, int] = field(default_factory=dict)
|
||||
skipped_by_reason: dict[str, int] = field(default_factory=dict)
|
||||
deterministic_selected: set[type[Matcher]] = field(default_factory=set)
|
||||
total_descriptors: int = 0
|
||||
candidate_count: int = 0
|
||||
|
||||
def mark_selected(self, lane: str, reason: str = "selected") -> None:
|
||||
self.selected_by_lane[lane] = self.selected_by_lane.get(lane, 0) + 1
|
||||
self.selected_by_reason[reason] = self.selected_by_reason.get(reason, 0) + 1
|
||||
|
||||
def mark_skipped(self, lane: str, reason: str = "skipped") -> None:
|
||||
self.skipped_by_lane[lane] = self.skipped_by_lane.get(lane, 0) + 1
|
||||
self.skipped_by_reason[reason] = self.skipped_by_reason.get(reason, 0) + 1
|
||||
|
||||
|
||||
class HandlerActivationIndex:
|
||||
"""In-memory matcher activation index.
|
||||
|
||||
The index is intentionally fail-open: only proven misses are rejected before
|
||||
matcher task creation. Unknown custom rules, incomplete command metadata,
|
||||
and Alconna shortcut misses stay selected so plugin compatibility wins over
|
||||
dispatch aggressiveness.
|
||||
The index is intentionally fail-open: only known NoneBot rule misses are
|
||||
rejected before matcher task creation. Custom rules, incomplete command
|
||||
metadata, and Alconna shortcut misses stay selected so plugin compatibility
|
||||
wins over dispatch aggressiveness.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._by_priority: dict[int, list[HandlerDescriptor]] = {}
|
||||
self._matcher_map: dict[type[Matcher], HandlerDescriptor] = {}
|
||||
self._source_keys: set[tuple[int, tuple[int, ...]]] = set()
|
||||
self._alconna_index = AlconnaActivationIndex()
|
||||
self._compiled = False
|
||||
|
||||
@property
|
||||
@@ -122,6 +212,7 @@ class HandlerActivationIndex:
|
||||
(int(priority), tuple(id(matcher) for matcher in priority_matchers))
|
||||
)
|
||||
self._source_keys = source_keys
|
||||
self._alconna_index.rebuild(self._matcher_map.values())
|
||||
self._compiled = True
|
||||
|
||||
def ensure_fresh(self, matchers: dict[int, list[type[Matcher]]]) -> None:
|
||||
@@ -164,32 +255,16 @@ class HandlerActivationIndex:
|
||||
]
|
||||
for descriptor in descriptors:
|
||||
decision = self._select_descriptor(descriptor, context)
|
||||
lane = descriptor.lane
|
||||
if decision == "fallback":
|
||||
result.fallback_required = True
|
||||
if not _consume_uncertain_budget(descriptor, budget):
|
||||
result.mark_skipped(lane, "fallback_budget_exhausted")
|
||||
continue
|
||||
result.selected.append(descriptor.matcher)
|
||||
result.mark_selected(lane, "fallback_budgeted")
|
||||
continue
|
||||
if decision == "miss":
|
||||
result.mark_skipped(lane, _miss_reason(descriptor, context))
|
||||
continue
|
||||
if decision == "deterministic":
|
||||
result.selected.append(descriptor.matcher)
|
||||
result.deterministic_selected.add(descriptor.matcher)
|
||||
result.mark_selected(lane, "deterministic")
|
||||
continue
|
||||
if not _selection_is_guaranteed(descriptor, context):
|
||||
if not _consume_uncertain_budget(descriptor, budget):
|
||||
result.mark_skipped(lane, "unknown_budget_exhausted")
|
||||
if _is_throttleable_broad_passive(descriptor, context):
|
||||
if not _consume_broad_passive_budget(descriptor, budget):
|
||||
continue
|
||||
selected_reason = "unknown_budgeted"
|
||||
else:
|
||||
selected_reason = "guaranteed"
|
||||
result.selected.append(descriptor.matcher)
|
||||
result.mark_selected(lane, selected_reason)
|
||||
result.candidate_count = len(result.selected)
|
||||
return result
|
||||
|
||||
@@ -197,7 +272,7 @@ class HandlerActivationIndex:
|
||||
self,
|
||||
descriptor: HandlerDescriptor,
|
||||
context: ActivationContext,
|
||||
) -> Literal["select", "miss", "fallback", "deterministic"]:
|
||||
) -> Literal["select", "miss", "deterministic"]:
|
||||
if descriptor.temp:
|
||||
return "select"
|
||||
matcher_type = descriptor.matcher_type
|
||||
@@ -224,7 +299,7 @@ class HandlerActivationIndex:
|
||||
self,
|
||||
descriptor: HandlerDescriptor,
|
||||
context: ActivationContext,
|
||||
) -> Literal["select", "miss", "fallback", "deterministic"]:
|
||||
) -> Literal["select", "miss", "deterministic"]:
|
||||
texts = text_match_candidates(
|
||||
context.plain_text,
|
||||
context.raw_text,
|
||||
@@ -232,8 +307,25 @@ class HandlerActivationIndex:
|
||||
)
|
||||
if not texts:
|
||||
return "select"
|
||||
rule_match = matcher_rule_matches_text(
|
||||
descriptor.rules,
|
||||
context.raw_text,
|
||||
context.plain_text,
|
||||
event=context.event,
|
||||
to_me=context.to_me,
|
||||
)
|
||||
if rule_match == "miss":
|
||||
return "miss"
|
||||
command_matched = False
|
||||
if descriptor.commands:
|
||||
if descriptor.alconna:
|
||||
alconna_match = self._alconna_index.select(descriptor, context, texts)
|
||||
if alconna_match == "match":
|
||||
command_matched = True
|
||||
elif alconna_match == "miss":
|
||||
return "miss"
|
||||
else:
|
||||
return "select"
|
||||
elif descriptor.commands:
|
||||
if any(
|
||||
matcher_command_matches(text, command)
|
||||
for text in texts
|
||||
@@ -248,9 +340,7 @@ class HandlerActivationIndex:
|
||||
if shortcut_match == "match":
|
||||
command_matched = True
|
||||
else:
|
||||
# Command extraction for Alconna/custom matchers is incomplete by
|
||||
# design; a miss here is not proof that NoneBot will miss.
|
||||
return "select"
|
||||
return "miss" if not descriptor.has_custom_rule else "select"
|
||||
else:
|
||||
shortcut_match = matcher_alconna_shortcut_matches_any(
|
||||
descriptor.shortcuts,
|
||||
@@ -259,22 +349,19 @@ class HandlerActivationIndex:
|
||||
if shortcut_match == "match":
|
||||
command_matched = True
|
||||
|
||||
rule_match = matcher_rule_matches_text(
|
||||
descriptor.rules,
|
||||
context.raw_text,
|
||||
context.plain_text,
|
||||
event=context.event,
|
||||
to_me=context.to_me,
|
||||
)
|
||||
if rule_match == "match":
|
||||
if (
|
||||
rule_match == "match"
|
||||
and not descriptor.has_custom_rule
|
||||
and descriptor.shortcuts is None
|
||||
and not descriptor.alconna
|
||||
):
|
||||
command_matched = True
|
||||
elif rule_match == "miss":
|
||||
return "miss"
|
||||
elif (
|
||||
if (
|
||||
not (descriptor.commands or descriptor.shortcuts is not None)
|
||||
and not command_matched
|
||||
and not descriptor.alconna
|
||||
):
|
||||
return "fallback"
|
||||
return "select"
|
||||
|
||||
if (
|
||||
context.ai_route_modules
|
||||
@@ -284,7 +371,7 @@ class HandlerActivationIndex:
|
||||
return "select"
|
||||
|
||||
if command_matched:
|
||||
return "deterministic"
|
||||
return "select" if descriptor.has_custom_rule else "deterministic"
|
||||
return "select"
|
||||
|
||||
def _build_descriptor(
|
||||
@@ -298,7 +385,10 @@ class HandlerActivationIndex:
|
||||
if hasattr(matcher, "command"):
|
||||
command_like = True
|
||||
commands = extract_matcher_command_literals(matcher) or ()
|
||||
alconna_descriptors = extract_matcher_alconna_descriptors(matcher)
|
||||
shortcuts = extract_matcher_alconna_shortcuts(matcher)
|
||||
if alconna_descriptors:
|
||||
command_like = True
|
||||
if shortcuts is not None:
|
||||
command_like = True
|
||||
module = matcher_module_name(matcher)
|
||||
@@ -323,6 +413,7 @@ class HandlerActivationIndex:
|
||||
has_custom_rule=matcher_has_custom_rule(matcher),
|
||||
commands=commands,
|
||||
shortcuts=shortcuts,
|
||||
alconna=alconna_descriptors,
|
||||
rules=rules,
|
||||
)
|
||||
|
||||
@@ -393,6 +484,8 @@ def matcher_is_command_like(matcher_cls: type[Matcher]) -> bool:
|
||||
return True
|
||||
if hasattr(matcher_cls, "command"):
|
||||
return True
|
||||
if extract_matcher_alconna_descriptors(matcher_cls):
|
||||
return True
|
||||
return extract_matcher_alconna_shortcuts(matcher_cls) is not None
|
||||
|
||||
|
||||
@@ -420,7 +513,10 @@ def classify_matcher_lane(
|
||||
if hasattr(matcher_cls, "command"):
|
||||
command_like = True
|
||||
commands = extract_matcher_command_literals(matcher_cls) or ()
|
||||
alconna_descriptors = extract_matcher_alconna_descriptors(matcher_cls)
|
||||
shortcuts = extract_matcher_alconna_shortcuts(matcher_cls)
|
||||
if alconna_descriptors:
|
||||
command_like = True
|
||||
if shortcuts is not None:
|
||||
command_like = True
|
||||
return classify_lane(
|
||||
@@ -437,11 +533,6 @@ def extract_matcher_rule_descriptors(
|
||||
matcher_cls: type[Matcher],
|
||||
) -> tuple[ActivationRuleDescriptor, ...]:
|
||||
descriptors: list[ActivationRuleDescriptor] = []
|
||||
if hasattr(matcher_cls, "command"):
|
||||
descriptors.append(
|
||||
ActivationRuleDescriptor("matcher_command", command_like=True)
|
||||
)
|
||||
|
||||
rule = getattr(matcher_cls, "rule", None)
|
||||
checkers = getattr(rule, "checkers", ()) or ()
|
||||
for checker in checkers:
|
||||
@@ -530,48 +621,10 @@ def extract_matcher_rule_descriptors(
|
||||
):
|
||||
descriptors.append(ActivationRuleDescriptor("alconna", command_like=True))
|
||||
else:
|
||||
descriptors.append(_custom_rule_descriptor(call))
|
||||
descriptors.append(ActivationRuleDescriptor("custom"))
|
||||
return tuple(descriptors)
|
||||
|
||||
|
||||
def _custom_rule_descriptor(call: object) -> ActivationRuleDescriptor:
|
||||
keyword_regex = _extract_keyword_regex_pairs(call)
|
||||
if keyword_regex:
|
||||
return ActivationRuleDescriptor(
|
||||
"keyword_regex",
|
||||
keyword_regex,
|
||||
deterministic_text=True,
|
||||
)
|
||||
return ActivationRuleDescriptor("custom")
|
||||
|
||||
|
||||
def _extract_keyword_regex_pairs(
|
||||
call: object,
|
||||
) -> tuple[tuple[str, str, int], ...]:
|
||||
"""Recognize generic keyword + regex custom rules without plugin coupling."""
|
||||
|
||||
source = getattr(call, "key_pattern_list", None)
|
||||
if source is None:
|
||||
source = getattr(call, "keyword_patterns", None)
|
||||
if source is None:
|
||||
source = getattr(call, "patterns", None)
|
||||
if not isinstance(source, Iterable) or isinstance(source, str):
|
||||
return ()
|
||||
|
||||
pairs: list[tuple[str, str, int]] = []
|
||||
for item in source:
|
||||
if not isinstance(item, tuple | list) or len(item) < 2:
|
||||
continue
|
||||
keyword = str(item[0] or "").strip()
|
||||
pattern_obj = item[1]
|
||||
pattern = getattr(pattern_obj, "pattern", pattern_obj)
|
||||
if not keyword or not isinstance(pattern, str) or not pattern:
|
||||
continue
|
||||
flags = int(getattr(pattern_obj, "flags", 0) or 0)
|
||||
pairs.append((keyword, pattern, flags))
|
||||
return tuple(pairs)
|
||||
|
||||
|
||||
def normalize_rule_string_tuple(value: object) -> tuple[str, ...]:
|
||||
if isinstance(value, str):
|
||||
return (value,)
|
||||
@@ -633,13 +686,38 @@ def matcher_rule_matches_text(
|
||||
|
||||
for descriptor in descriptors:
|
||||
kind = descriptor.kind
|
||||
if kind == "regex":
|
||||
if kind in {"custom", "alconna"}:
|
||||
saw_unknown = True
|
||||
continue
|
||||
if kind in {"command", "shell_command"}:
|
||||
saw_deterministic = True
|
||||
commands: set[str] = set()
|
||||
collect_command_literals(descriptor.value, commands)
|
||||
normalized_commands = {
|
||||
normalized
|
||||
for item in commands
|
||||
if (normalized := normalize_command(item))
|
||||
}
|
||||
if any(
|
||||
matcher_command_matches(text, command)
|
||||
for text in plain_candidates
|
||||
for command in normalized_commands
|
||||
):
|
||||
matched_any = True
|
||||
else:
|
||||
return "miss"
|
||||
elif kind in {"regex", "regex_fullmatch"}:
|
||||
saw_deterministic = True
|
||||
pattern = str(descriptor.value or "")
|
||||
if not pattern:
|
||||
continue
|
||||
try:
|
||||
if re.search(pattern, message_text, descriptor.flags):
|
||||
matched = (
|
||||
re.fullmatch(pattern, message_text, descriptor.flags)
|
||||
if kind == "regex_fullmatch"
|
||||
else re.search(pattern, message_text, descriptor.flags)
|
||||
)
|
||||
if matched:
|
||||
matched_any = True
|
||||
else:
|
||||
return "miss"
|
||||
@@ -709,13 +787,6 @@ def matcher_rule_matches_text(
|
||||
matched_any = True
|
||||
else:
|
||||
return "miss"
|
||||
elif kind == "keyword_regex":
|
||||
saw_deterministic = True
|
||||
values = descriptor.value if isinstance(descriptor.value, tuple) else ()
|
||||
if _keyword_regex_matches(values, plain_candidates):
|
||||
matched_any = True
|
||||
else:
|
||||
return "miss"
|
||||
elif kind == "to_me":
|
||||
if not to_me:
|
||||
return "miss"
|
||||
@@ -729,43 +800,15 @@ def matcher_rule_matches_text(
|
||||
elif isinstance(types, tuple) and types:
|
||||
if not isinstance(event, types):
|
||||
return "miss"
|
||||
elif kind in {"custom", "alconna", "matcher_command"}:
|
||||
saw_unknown = True
|
||||
|
||||
if matched_any:
|
||||
return "match"
|
||||
if saw_deterministic:
|
||||
return "miss"
|
||||
return "unknown" if saw_unknown else "match"
|
||||
if saw_unknown:
|
||||
return "unknown"
|
||||
if saw_deterministic:
|
||||
return "miss"
|
||||
return "unknown"
|
||||
|
||||
|
||||
def _keyword_regex_matches(
|
||||
values: object,
|
||||
candidates: tuple[str, ...],
|
||||
) -> bool:
|
||||
if not isinstance(values, tuple):
|
||||
return False
|
||||
for item in values:
|
||||
if not isinstance(item, tuple | list) or len(item) < 3:
|
||||
continue
|
||||
keyword, pattern, flags = item[:3]
|
||||
keyword_text = str(keyword or "")
|
||||
pattern_text = str(pattern or "")
|
||||
if not keyword_text or not pattern_text:
|
||||
continue
|
||||
for text in candidates:
|
||||
if keyword_text not in text:
|
||||
continue
|
||||
try:
|
||||
if re.search(pattern_text, text, int(flags or 0)):
|
||||
return True
|
||||
except re.error:
|
||||
return False
|
||||
return False
|
||||
|
||||
|
||||
def extract_matcher_command_literals(
|
||||
matcher_cls: type[Matcher],
|
||||
) -> tuple[str, ...] | None:
|
||||
@@ -912,6 +955,293 @@ def collect_alconna_shortcuts(value: Any, target: set[str], depth: int = 0) -> N
|
||||
collect_alconna_shortcuts(nested, target, depth + 1)
|
||||
|
||||
|
||||
def extract_matcher_alconna_descriptors(
|
||||
matcher_cls: type[Matcher],
|
||||
) -> tuple[AlconnaDescriptor, ...]:
|
||||
descriptors: list[AlconnaDescriptor] = []
|
||||
rule = getattr(matcher_cls, "rule", None)
|
||||
checkers = getattr(rule, "checkers", ()) or ()
|
||||
for checker in checkers:
|
||||
call = getattr(checker, "call", None)
|
||||
if call is None:
|
||||
continue
|
||||
if call.__class__.__name__ != "AlconnaRule":
|
||||
continue
|
||||
if not call.__class__.__module__.startswith("nonebot_plugin_alconna.rule"):
|
||||
continue
|
||||
command = resolve_maybe_weakref(
|
||||
getattr(call, "command", None) or getattr(call, "alconna", None)
|
||||
)
|
||||
descriptor = _build_alconna_descriptor(call, command)
|
||||
if descriptor is not None:
|
||||
descriptors.append(descriptor)
|
||||
return tuple(descriptors)
|
||||
|
||||
|
||||
def _build_alconna_descriptor(call: Any, command: Any) -> AlconnaDescriptor | None:
|
||||
if command is None:
|
||||
return None
|
||||
command_text = str(getattr(command, "command", "") or "").strip()
|
||||
aliases = tuple(
|
||||
str(item).strip()
|
||||
for item in getattr(command, "aliases", ()) or ()
|
||||
if str(item).strip() and str(item).strip() != command_text
|
||||
)
|
||||
prefixes = tuple(
|
||||
str(item)
|
||||
for item in getattr(command, "prefixes", ()) or ()
|
||||
if isinstance(item, str)
|
||||
)
|
||||
meta = getattr(command, "meta", None)
|
||||
shortcuts = _extract_alconna_shortcut_descriptors(command)
|
||||
return AlconnaDescriptor(
|
||||
command=command_text,
|
||||
aliases=aliases,
|
||||
prefixes=prefixes,
|
||||
shortcuts=shortcuts,
|
||||
compact=bool(getattr(meta, "compact", False)),
|
||||
skip_for_unmatch=bool(getattr(call, "skip", True)),
|
||||
before_rule_count=_rule_checker_count(getattr(call, "before_rules", None)),
|
||||
after_rule_count=_rule_checker_count(getattr(call, "after_rules", None)),
|
||||
before_rule_known_safe=_alconna_rule_is_known_safe(
|
||||
getattr(call, "before_rules", None)
|
||||
),
|
||||
after_rule_known_safe=_alconna_rule_is_known_safe(
|
||||
getattr(call, "after_rules", None)
|
||||
),
|
||||
input_rewrite_extensions=_extract_alconna_input_rewrite_extensions(call),
|
||||
)
|
||||
|
||||
|
||||
def _rule_checker_count(rule: Any) -> int:
|
||||
checkers = getattr(rule, "checkers", ()) or ()
|
||||
with contextlib.suppress(TypeError):
|
||||
return len(checkers)
|
||||
return 1
|
||||
|
||||
|
||||
def _alconna_rule_is_known_safe(rule: Any) -> bool:
|
||||
"""Whether Alconna before/after Rule can be reasoned about statically.
|
||||
|
||||
We still do not execute these rules here. A rule is considered safe only
|
||||
when every checker is an official NoneBot rule whose negative result can be
|
||||
reproduced by the selector. Custom rules stay fail-open.
|
||||
"""
|
||||
|
||||
checkers = getattr(rule, "checkers", ()) or ()
|
||||
for checker in checkers:
|
||||
call = getattr(checker, "call", None)
|
||||
if call is None:
|
||||
return False
|
||||
call_module = call.__class__.__module__
|
||||
call_name = call.__class__.__name__
|
||||
if not call_module.startswith("nonebot.rule"):
|
||||
return False
|
||||
if call_name not in KNOWN_SAFE_RULE_NAMES:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _extract_alconna_shortcut_descriptors(
|
||||
command: Any,
|
||||
) -> tuple[AlconnaShortcutDescriptor, ...]:
|
||||
shortcuts: list[AlconnaShortcutDescriptor] = []
|
||||
with contextlib.suppress(Exception):
|
||||
from arclet.alconna import command_manager
|
||||
|
||||
raw_shortcuts = command_manager.get_shortcut(command) # type: ignore[arg-type]
|
||||
if isinstance(raw_shortcuts, dict):
|
||||
for key, args in raw_shortcuts.items():
|
||||
pattern = str(key or "").strip()
|
||||
if not pattern:
|
||||
continue
|
||||
shortcuts.append(
|
||||
AlconnaShortcutDescriptor(
|
||||
pattern=pattern,
|
||||
fuzzy=bool(getattr(args, "fuzzy", False)),
|
||||
flags=int(getattr(args, "flags", 0) or 0),
|
||||
)
|
||||
)
|
||||
if shortcuts:
|
||||
return tuple(shortcuts)
|
||||
fallback: set[str] = set()
|
||||
collect_alconna_shortcuts(command, fallback)
|
||||
return tuple(
|
||||
AlconnaShortcutDescriptor(pattern=item)
|
||||
for item in sorted(fallback)
|
||||
if item.strip()
|
||||
)
|
||||
|
||||
|
||||
def _extract_alconna_input_rewrite_extensions(call: Any) -> tuple[str, ...]:
|
||||
executor = getattr(call, "executor", None)
|
||||
if executor is None:
|
||||
return ()
|
||||
result: list[str] = []
|
||||
for attr in ("extensions", "_extensions", "exts", "context"):
|
||||
extensions = getattr(executor, attr, None)
|
||||
if not isinstance(extensions, list | tuple | set | frozenset):
|
||||
continue
|
||||
for extension in extensions:
|
||||
overrides = getattr(extension.__class__, "_overrides", None)
|
||||
if not isinstance(overrides, dict):
|
||||
overrides = getattr(extension, "_overrides", None)
|
||||
if not isinstance(overrides, dict):
|
||||
continue
|
||||
if not (
|
||||
bool(overrides.get("message_provider"))
|
||||
or bool(overrides.get("receive_wrapper"))
|
||||
):
|
||||
continue
|
||||
name = extension.__class__.__name__
|
||||
if name not in result:
|
||||
result.append(name)
|
||||
return tuple(result)
|
||||
|
||||
|
||||
def _alconna_can_prefilter(alconna: AlconnaDescriptor) -> bool:
|
||||
if not (alconna.command or alconna.aliases or alconna.shortcuts):
|
||||
return False
|
||||
if not (alconna.before_rule_known_safe and alconna.after_rule_known_safe):
|
||||
return False
|
||||
if any(name != "ReplyMergeExtension" for name in alconna.input_rewrite_extensions):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def matcher_alconna_head_matches(
|
||||
alconna: AlconnaDescriptor,
|
||||
texts: Iterable[str],
|
||||
context: ActivationContext,
|
||||
) -> ActivationDecision:
|
||||
if not _alconna_can_prefilter(alconna):
|
||||
return "unknown"
|
||||
if alconna.has_reply_merge_extension and _event_has_reply(context.event):
|
||||
return "unknown"
|
||||
|
||||
candidates = tuple(text.strip() for text in texts if text and text.strip())
|
||||
if not candidates:
|
||||
return "unknown"
|
||||
|
||||
saw_unknown = False
|
||||
command_heads = (alconna.command, *alconna.aliases)
|
||||
for text in candidates:
|
||||
for command in command_heads:
|
||||
if not command:
|
||||
continue
|
||||
decision = _alconna_command_head_matches(
|
||||
text,
|
||||
command,
|
||||
alconna.prefixes,
|
||||
compact=alconna.compact,
|
||||
)
|
||||
if decision == "match":
|
||||
return "match"
|
||||
if decision == "unknown":
|
||||
saw_unknown = True
|
||||
for shortcut in alconna.shortcuts:
|
||||
decision = _alconna_shortcut_matches(text, shortcut)
|
||||
if decision == "match":
|
||||
return "match"
|
||||
if decision == "unknown":
|
||||
saw_unknown = True
|
||||
return "unknown" if saw_unknown else "miss"
|
||||
|
||||
|
||||
def _alconna_command_head_matches(
|
||||
text: str,
|
||||
command: str,
|
||||
prefixes: tuple[str, ...],
|
||||
*,
|
||||
compact: bool,
|
||||
) -> ActivationDecision:
|
||||
normalized = command.strip()
|
||||
if not normalized:
|
||||
return "unknown"
|
||||
if normalized.startswith("re:"):
|
||||
pattern = normalized.removeprefix("re:").strip()
|
||||
if not pattern:
|
||||
return "unknown"
|
||||
for prefix in prefixes or ("",):
|
||||
try:
|
||||
if re.match(rf"^{re.escape(prefix)}(?:{pattern})", text):
|
||||
return "match"
|
||||
except re.error:
|
||||
return "unknown"
|
||||
return "miss"
|
||||
for prefix in prefixes or ("",):
|
||||
head = f"{prefix}{normalized}"
|
||||
if _alconna_literal_head_matches(text, head, compact=compact):
|
||||
return "match"
|
||||
return "miss"
|
||||
|
||||
|
||||
def _alconna_literal_head_matches(text: str, head: str, *, compact: bool) -> bool:
|
||||
if not text or not head:
|
||||
return False
|
||||
if text == head:
|
||||
return True
|
||||
if not text.startswith(head):
|
||||
return False
|
||||
if len(text) == len(head):
|
||||
return True
|
||||
rest = text[len(head) :]
|
||||
if rest and rest[0].isspace():
|
||||
return True
|
||||
return bool(compact or not head[-1].isascii())
|
||||
|
||||
|
||||
def _alconna_shortcut_matches(
|
||||
text: str,
|
||||
shortcut: AlconnaShortcutDescriptor,
|
||||
) -> ActivationDecision:
|
||||
pattern = shortcut.pattern.strip()
|
||||
if not pattern:
|
||||
return "unknown"
|
||||
normalized = normalize_shortcut_pattern(pattern)
|
||||
placeholder_match = _placeholder_shortcut_decision(text, normalized)
|
||||
if placeholder_match == "match":
|
||||
return "match"
|
||||
if placeholder_match == "unknown":
|
||||
return "unknown"
|
||||
if not is_regex_like_shortcut(normalized) and matcher_command_matches(
|
||||
text,
|
||||
normalized,
|
||||
):
|
||||
return "match"
|
||||
try:
|
||||
if shortcut.fuzzy:
|
||||
return "match" if re.match(f"^{pattern}", text, shortcut.flags) else "miss"
|
||||
return "match" if re.fullmatch(pattern, text, shortcut.flags) else "miss"
|
||||
except re.error:
|
||||
return "unknown"
|
||||
|
||||
|
||||
def _event_has_reply(event: object | None) -> bool:
|
||||
if event is None:
|
||||
return False
|
||||
with contextlib.suppress(Exception):
|
||||
message = event.get_message() # type: ignore[attr-defined]
|
||||
for segment in message:
|
||||
segment_type = getattr(segment, "type", None)
|
||||
if segment_type == "reply":
|
||||
return True
|
||||
if isinstance(segment, dict) and segment.get("type") == "reply":
|
||||
return True
|
||||
for attr in ("reply", "reply_message", "source"):
|
||||
with contextlib.suppress(Exception):
|
||||
if getattr(event, attr, None) is not None:
|
||||
return True
|
||||
raw_text = ""
|
||||
with contextlib.suppress(Exception):
|
||||
raw_text = str(event.get_message()) # type: ignore[attr-defined]
|
||||
lowered = raw_text.casefold()
|
||||
return any(
|
||||
marker in lowered
|
||||
for marker in ("[cq:reply", "type=reply", '"reply"', "'reply'")
|
||||
)
|
||||
|
||||
|
||||
def resolve_maybe_weakref(value: Any) -> Any:
|
||||
if isinstance(value, weakref.ReferenceType):
|
||||
resolved = value()
|
||||
@@ -1026,8 +1356,12 @@ def shortcut_matches_text(text: str, shortcut: str) -> bool:
|
||||
|
||||
|
||||
def placeholder_shortcut_matches(text: str, pattern: str) -> bool:
|
||||
return _placeholder_shortcut_decision(text, pattern) == "match"
|
||||
|
||||
|
||||
def _placeholder_shortcut_decision(text: str, pattern: str) -> ActivationDecision:
|
||||
if "{" not in pattern or "}" not in pattern:
|
||||
return False
|
||||
return "miss"
|
||||
pieces: list[str] = []
|
||||
last = 0
|
||||
for match in re.finditer(r"\{[^{}]+\}", pattern):
|
||||
@@ -1035,12 +1369,16 @@ def placeholder_shortcut_matches(text: str, pattern: str) -> bool:
|
||||
pieces.append(r"\S+")
|
||||
last = match.end()
|
||||
if not pieces:
|
||||
return False
|
||||
return "miss"
|
||||
pieces.append(re.escape(pattern[last:]))
|
||||
try:
|
||||
return re.match(rf"^{''.join(pieces)}(?:\s|$)", text) is not None
|
||||
return (
|
||||
"match"
|
||||
if re.match(rf"^{''.join(pieces)}(?:\s|$)", text) is not None
|
||||
else "miss"
|
||||
)
|
||||
except re.error:
|
||||
return False
|
||||
return "unknown"
|
||||
|
||||
|
||||
def is_regex_like_shortcut(pattern: str) -> bool:
|
||||
@@ -1072,21 +1410,27 @@ def matcher_matches_ai_route_heads(
|
||||
return False
|
||||
|
||||
|
||||
def _selection_is_guaranteed(
|
||||
def _is_throttleable_broad_passive(
|
||||
descriptor: HandlerDescriptor,
|
||||
context: ActivationContext,
|
||||
) -> bool:
|
||||
"""Return True for candidates that must not be budget-throttled."""
|
||||
"""Only broad, no-rule passive message matchers may be budget-throttled."""
|
||||
|
||||
if descriptor.temp or descriptor.lane == "system":
|
||||
return True
|
||||
if context.event_type != "message":
|
||||
return not descriptor.command_like
|
||||
return False
|
||||
if descriptor.temp or descriptor.lane == "system":
|
||||
return False
|
||||
if not descriptor.lane.startswith("passive_"):
|
||||
return False
|
||||
if descriptor.command_like or descriptor.deterministic_text:
|
||||
return False
|
||||
if descriptor.has_custom_rule or descriptor.rules:
|
||||
return False
|
||||
if descriptor.lane == "passive_http" and (
|
||||
context.has_url or _looks_like_rich_message(context.raw_text)
|
||||
):
|
||||
return True
|
||||
return False
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _looks_like_rich_message(text: str) -> bool:
|
||||
@@ -1106,34 +1450,11 @@ def _looks_like_rich_message(text: str) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _miss_reason(descriptor: HandlerDescriptor, context: ActivationContext) -> str:
|
||||
matcher_type = descriptor.matcher_type
|
||||
if matcher_type and matcher_type != context.event_type:
|
||||
return "type_miss"
|
||||
if context.event_type != "message" and descriptor.command_like:
|
||||
return "non_message_command"
|
||||
if descriptor.command_like:
|
||||
return "command_or_rule_miss"
|
||||
return "rule_miss"
|
||||
|
||||
|
||||
def _uncertain_budget_lane(descriptor: HandlerDescriptor) -> str:
|
||||
lane = descriptor.lane
|
||||
if lane.startswith("passive_"):
|
||||
return lane
|
||||
# Unknown command-like matchers are fail-open for compatibility, but they
|
||||
# should not fan out as unbounded command tasks when no deterministic signal
|
||||
# matched. Put them into the cheapest passive bucket.
|
||||
if lane.startswith("command_"):
|
||||
return "passive_light"
|
||||
return lane
|
||||
|
||||
|
||||
def _consume_uncertain_budget(
|
||||
def _consume_broad_passive_budget(
|
||||
descriptor: HandlerDescriptor,
|
||||
budget: dict[str, int],
|
||||
) -> bool:
|
||||
lane = _uncertain_budget_lane(descriptor)
|
||||
lane = descriptor.lane
|
||||
if lane not in budget:
|
||||
return True
|
||||
if budget[lane] <= 0:
|
||||
|
||||
@@ -15,19 +15,9 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services.auth_observability import (
|
||||
append_auth_decision_log,
|
||||
append_runtime_backpressure_log,
|
||||
build_auth_observability_report,
|
||||
)
|
||||
from zhenxun.services.cache.cache_containers import CacheDict
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
PluginInfoMemoryCache,
|
||||
PluginLimitMemoryCache,
|
||||
)
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import is_overloaded, signal_overload
|
||||
from zhenxun.services.message_load import signal_overload
|
||||
from zhenxun.utils.enum import GoldHandle, PluginType
|
||||
from zhenxun.utils.exception import InsufficientGold
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
@@ -48,6 +38,7 @@ from .auth.context import (
|
||||
set_route_modules,
|
||||
store_permission_context,
|
||||
)
|
||||
from .auth.data_provider import DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
from .auth.exception import (
|
||||
IsSuperuserException,
|
||||
PermissionExemption,
|
||||
@@ -55,7 +46,6 @@ from .auth.exception import (
|
||||
)
|
||||
from .auth_activation import (
|
||||
ActivationContext,
|
||||
ActivationResult,
|
||||
HandlerActivationIndex,
|
||||
classify_matcher_lane,
|
||||
extract_matcher_alconna_shortcuts,
|
||||
@@ -147,9 +137,7 @@ _ROUTE_COMMAND_MAP: dict[str, set[str]] = {}
|
||||
_ROUTE_PREFIX_MAP: dict[str, set[str]] = {}
|
||||
_ROUTE_MODULES_WITH_COMMANDS: set[str] = set()
|
||||
MATCHER_ROUTE_PREFILTER_TTL = AUTH_DISPATCH_RUNTIME_CONFIG.matcher_route_prefilter_ttl
|
||||
PREFILTER_STATS_LOG_INTERVAL = AUTH_DISPATCH_RUNTIME_CONFIG.prefilter_stats_log_interval
|
||||
CACHE_SWEEP_INTERVAL = AUTH_DISPATCH_RUNTIME_CONFIG.cache_sweep_interval
|
||||
DISPATCH_STATS_LOG_INTERVAL = AUTH_DISPATCH_RUNTIME_CONFIG.dispatch_stats_log_interval
|
||||
|
||||
# 全局信号量与计数器
|
||||
HOOKS_ACTIVE_COUNT = 0
|
||||
@@ -182,56 +170,6 @@ _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,
|
||||
"empty_text": 0,
|
||||
}
|
||||
_PREFILTER_LAST_LOG = 0.0
|
||||
_DISPATCH_SELECTED = 0
|
||||
_DISPATCH_SKIPPED = 0
|
||||
_DISPATCH_SELECTED_BY_LANE: dict[str, int] = {
|
||||
"command_exact": 0,
|
||||
"command_shortcut": 0,
|
||||
"command_regex": 0,
|
||||
"system": 0,
|
||||
"passive_light": 0,
|
||||
"passive_db": 0,
|
||||
"passive_http": 0,
|
||||
"passive_ai": 0,
|
||||
"passive_render": 0,
|
||||
}
|
||||
_DISPATCH_SKIPPED_BY_LANE: dict[str, int] = {
|
||||
"command_exact": 0,
|
||||
"command_shortcut": 0,
|
||||
"command_regex": 0,
|
||||
"system": 0,
|
||||
"passive_light": 0,
|
||||
"passive_db": 0,
|
||||
"passive_http": 0,
|
||||
"passive_ai": 0,
|
||||
"passive_render": 0,
|
||||
}
|
||||
_DISPATCH_LANE_WAIT_MS: dict[str, float] = {
|
||||
"command_exact": 0.0,
|
||||
"command_shortcut": 0.0,
|
||||
"command_regex": 0.0,
|
||||
"system": 0.0,
|
||||
"passive_light": 0.0,
|
||||
"passive_db": 0.0,
|
||||
"passive_http": 0.0,
|
||||
"passive_ai": 0.0,
|
||||
"passive_render": 0.0,
|
||||
}
|
||||
_DISPATCH_LAST_LOG = 0.0
|
||||
_DISPATCH_SHADOW_LAST_LOG = 0.0
|
||||
_CACHE_SWEEP_TASK: asyncio.Task | None = None
|
||||
_BOT_WAKE_COMMAND_PATTERN = re.compile(r"^bot醒来(?:\s+\S+)?$", re.IGNORECASE)
|
||||
_BOT_WAKE_CANONICAL_PATTERN = re.compile(
|
||||
@@ -367,12 +305,6 @@ def _is_bot_wake_command(module: str, text: str | None) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _debug_log(message: str, *args, **kwargs) -> None:
|
||||
if is_overloaded():
|
||||
return
|
||||
logger.debug(message, *args, **kwargs)
|
||||
|
||||
|
||||
def _is_command_matcher_class(matcher_cls: type[Matcher]) -> bool:
|
||||
if matcher_cls in _MATCHER_COMMAND_TYPE_CACHE:
|
||||
return _MATCHER_COMMAND_TYPE_CACHE[matcher_cls]
|
||||
@@ -728,41 +660,6 @@ def _activation_context_from_dispatch(
|
||||
)
|
||||
|
||||
|
||||
def _record_dispatch_selection(lane: str, selected: bool, wait_ms: float = 0.0) -> None:
|
||||
global _DISPATCH_LAST_LOG, _DISPATCH_SELECTED, _DISPATCH_SKIPPED
|
||||
if lane == "command":
|
||||
lane = "command_exact"
|
||||
lane = lane if lane in _DISPATCH_SELECTED_BY_LANE else "passive_light"
|
||||
if selected:
|
||||
_DISPATCH_SELECTED += 1
|
||||
_DISPATCH_SELECTED_BY_LANE[lane] += 1
|
||||
_DISPATCH_LANE_WAIT_MS[lane] += wait_ms
|
||||
else:
|
||||
_DISPATCH_SKIPPED += 1
|
||||
_DISPATCH_SKIPPED_BY_LANE[lane] += 1
|
||||
|
||||
now = time.monotonic()
|
||||
if now - _DISPATCH_LAST_LOG < DISPATCH_STATS_LOG_INTERVAL or is_overloaded():
|
||||
return
|
||||
_DISPATCH_LAST_LOG = now
|
||||
wait_snapshot = {
|
||||
lane: round(wait, 2) for lane, wait in _DISPATCH_LANE_WAIT_MS.items()
|
||||
}
|
||||
lane_snapshot = " ".join(
|
||||
f"{lane}={count}" for lane, count in _DISPATCH_SELECTED_BY_LANE.items()
|
||||
)
|
||||
_debug_log(
|
||||
(
|
||||
"dispatch stats: "
|
||||
f"selected={_DISPATCH_SELECTED} "
|
||||
f"skipped={_DISPATCH_SKIPPED} "
|
||||
f"{lane_snapshot} "
|
||||
f"wait_ms={wait_snapshot}"
|
||||
),
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
def _new_dispatch_budget() -> dict[str, int]:
|
||||
return dict(_DISPATCH_LANE_LIMITS)
|
||||
|
||||
@@ -775,52 +672,6 @@ def _merge_dispatch_budget(
|
||||
target[lane] = source.get(lane, target.get(lane, 0))
|
||||
|
||||
|
||||
def _record_activation_result(activation_result: ActivationResult) -> None:
|
||||
for lane, count in activation_result.skipped_by_lane.items():
|
||||
for _ in range(count):
|
||||
_record_dispatch_selection(lane, False)
|
||||
|
||||
|
||||
def _compact_counter(counter: dict[str, int], limit: int = 6) -> str:
|
||||
if not counter:
|
||||
return "-"
|
||||
items = sorted(counter.items(), key=lambda item: item[1], reverse=True)[:limit]
|
||||
return ",".join(f"{key}={value}" for key, value in items)
|
||||
|
||||
|
||||
def _debug_activation_shadow(
|
||||
*,
|
||||
priority: int,
|
||||
activation_result: ActivationResult,
|
||||
context: EventDispatchContext,
|
||||
) -> None:
|
||||
global _DISPATCH_SHADOW_LAST_LOG
|
||||
if is_overloaded():
|
||||
return
|
||||
now = time.monotonic()
|
||||
if now - _DISPATCH_SHADOW_LAST_LOG < DISPATCH_STATS_LOG_INTERVAL:
|
||||
return
|
||||
_DISPATCH_SHADOW_LAST_LOG = now
|
||||
text_hint = (context.plain_text or context.trie_raw_command or "")[:48]
|
||||
_debug_log(
|
||||
(
|
||||
"dispatch shadow: "
|
||||
f"priority={priority} "
|
||||
f"event={context.event_type} "
|
||||
f"selected={activation_result.candidate_count}/"
|
||||
f"{activation_result.total_descriptors} "
|
||||
f"selected_lane={_compact_counter(activation_result.selected_by_lane)} "
|
||||
f"skipped_lane={_compact_counter(activation_result.skipped_by_lane)} "
|
||||
f"selected_reason="
|
||||
f"{_compact_counter(activation_result.selected_by_reason)} "
|
||||
f"skipped_reason="
|
||||
f"{_compact_counter(activation_result.skipped_by_reason)} "
|
||||
f"text={text_hint!r}"
|
||||
),
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
def _auth_scope_key(context: EventContext) -> str:
|
||||
group_id = context.group_id or ""
|
||||
channel_id = context.channel_id or ""
|
||||
@@ -876,7 +727,6 @@ async def _dispatch_lane_section(lane: str):
|
||||
wait_ms = (time.perf_counter() - started) * 1000
|
||||
if wait_ms >= AUTH_OVERLOAD_LANE_WAIT_MS:
|
||||
signal_overload(2.0)
|
||||
_record_dispatch_selection(lane, True, wait_ms=wait_ms)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
@@ -891,11 +741,6 @@ def get_dispatch_snapshot() -> dict[str, object]:
|
||||
value = getattr(semaphore, "_value", limit)
|
||||
lane_active[lane] = max(limit - int(value), 0)
|
||||
return {
|
||||
"selected": _DISPATCH_SELECTED,
|
||||
"skipped": _DISPATCH_SKIPPED,
|
||||
"selected_by_lane": dict(_DISPATCH_SELECTED_BY_LANE),
|
||||
"skipped_by_lane": dict(_DISPATCH_SKIPPED_BY_LANE),
|
||||
"lane_wait_ms": dict(_DISPATCH_LANE_WAIT_MS),
|
||||
"lane_active": lane_active,
|
||||
"lane_limits": dict(_DISPATCH_LANE_LIMITS),
|
||||
}
|
||||
@@ -975,58 +820,6 @@ async def _run_selected_matcher(
|
||||
)
|
||||
|
||||
|
||||
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":
|
||||
_PREFILTER_STATS["route_miss"] += 1
|
||||
elif reason == "command_miss":
|
||||
_PREFILTER_STATS["command_miss"] += 1
|
||||
elif reason == "empty_text":
|
||||
_PREFILTER_STATS["empty_text"] += 1
|
||||
|
||||
if _PREFILTER_STATS["checked"] % 1024 == 0:
|
||||
with contextlib.suppress(Exception):
|
||||
_ = len(_CHECK_MATCHER_ROUTE_CACHE)
|
||||
|
||||
now = time.monotonic()
|
||||
if now - _PREFILTER_LAST_LOG < PREFILTER_STATS_LOG_INTERVAL or is_overloaded():
|
||||
return
|
||||
_PREFILTER_LAST_LOG = now
|
||||
_debug_log(
|
||||
(
|
||||
"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']} "
|
||||
f"empty_text={_PREFILTER_STATS['empty_text']}"
|
||||
),
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
_MAX_MATCHER_CACHE = 512
|
||||
|
||||
|
||||
@@ -1038,8 +831,6 @@ _SELECTOR_DEPS = HandleEventSelectorDependencies(
|
||||
activation_context_from_dispatch=_activation_context_from_dispatch,
|
||||
new_dispatch_budget=_new_dispatch_budget,
|
||||
dispatch_lane_for_matcher=_dispatch_lane_for_matcher,
|
||||
record_activation_result=_record_activation_result,
|
||||
debug_activation_shadow=_debug_activation_shadow,
|
||||
merge_dispatch_budget=_merge_dispatch_budget,
|
||||
build_matcher_state=_build_matcher_state,
|
||||
run_selected_matcher=_run_selected_matcher,
|
||||
@@ -1068,15 +859,6 @@ async def _get_route_context(text: str, event_cache: dict | None) -> set[str]:
|
||||
return matched
|
||||
|
||||
|
||||
def _get_auth_route_precheck_deps() -> dict:
|
||||
return {
|
||||
"route_modules_with_commands": _ROUTE_MODULES_WITH_COMMANDS,
|
||||
"get_route_context": _get_route_context,
|
||||
"is_command_matcher_class": _is_command_matcher_class,
|
||||
"matcher_has_alconna_shortcuts": _matcher_has_alconna_shortcuts,
|
||||
}
|
||||
|
||||
|
||||
async def _cache_sweep_loop() -> None:
|
||||
while True:
|
||||
await asyncio.sleep(CACHE_SWEEP_INTERVAL)
|
||||
@@ -1145,7 +927,8 @@ async def _has_limits_cached(
|
||||
if event_cache is not None:
|
||||
event_cache.setdefault("module_limits_ready", {})[module] = True
|
||||
return has_limits
|
||||
limit_entries = PluginLimitMemoryCache.get_limits_if_ready(module)
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
limit_entries = provider.get_module_limits_if_ready(module)
|
||||
if limit_entries is not None:
|
||||
has_limits = bool(limit_entries)
|
||||
module_limit_cache[module] = has_limits
|
||||
@@ -1278,16 +1061,17 @@ async def _get_plugin_cache_first(
|
||||
*,
|
||||
allow_cache_load: bool,
|
||||
) -> tuple[PluginInfo | None, bool]:
|
||||
provider = DEFAULT_PERMISSION_DATA_PROVIDER
|
||||
plugin = None
|
||||
if event_cache is not None:
|
||||
plugin_cache = event_cache.setdefault("plugin_cache", {})
|
||||
if module in plugin_cache:
|
||||
return cast(PluginInfo | None, plugin_cache[module]), False
|
||||
|
||||
plugin = PluginInfoMemoryCache.get_by_module_if_ready(module)
|
||||
cache_miss = plugin is None and not PluginInfoMemoryCache.is_loaded()
|
||||
plugin = provider.get_plugin_if_ready(module)
|
||||
cache_miss = plugin is None and not provider.plugin_cache_loaded()
|
||||
if plugin is None and allow_cache_load:
|
||||
plugin = await PluginInfoMemoryCache.get_by_module(module)
|
||||
plugin = await provider.get_plugin(module)
|
||||
cache_miss = False
|
||||
if event_cache is not None:
|
||||
event_cache.setdefault("plugin_cache", {})[module] = plugin
|
||||
@@ -1351,7 +1135,6 @@ async def reserve_gold(
|
||||
)
|
||||
except InsufficientGold:
|
||||
raise
|
||||
await DataAccess(UserConsole).clear_cache(user_id=user_id)
|
||||
logger.debug(f"预扣功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session)
|
||||
return reservation
|
||||
|
||||
@@ -1384,14 +1167,12 @@ async def _record_backpressure(
|
||||
action: str,
|
||||
duration_ms: float = 0.0,
|
||||
) -> None:
|
||||
await append_runtime_backpressure_log(
|
||||
scope_key=lane_context.scope_key,
|
||||
reason=reason,
|
||||
lane=lane_context.lane,
|
||||
action=action,
|
||||
queue_size=lane_context.queue_size,
|
||||
active_count=HOOKS_ACTIVE_COUNT,
|
||||
duration_ms=duration_ms,
|
||||
logger.debug(
|
||||
"auth backpressure: "
|
||||
f"scope={lane_context.scope_key}, lane={lane_context.lane}, "
|
||||
f"reason={reason}, action={action}, queue={lane_context.queue_size}, "
|
||||
f"active={HOOKS_ACTIVE_COUNT}, duration_ms={duration_ms:.1f}",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
@@ -1437,7 +1218,6 @@ async def _prepare_auth_state(
|
||||
context: EventContext,
|
||||
bot: Bot,
|
||||
event_cache: dict | None,
|
||||
route_skip_checks: bool,
|
||||
skip_ban: bool,
|
||||
hook_recorder: HookTraceRecorder,
|
||||
state: dict | None,
|
||||
@@ -1502,7 +1282,6 @@ async def _prepare_auth_state(
|
||||
|
||||
policy_context = PolicyContext(
|
||||
snapshot=snapshot,
|
||||
route_skip_checks=route_skip_checks,
|
||||
allow_sleep_bypass=_is_bot_wake_command(module, context.plain_text),
|
||||
allow_group_sleep_bypass=_is_group_wake_command(plugin, context.plain_text),
|
||||
)
|
||||
@@ -1522,7 +1301,6 @@ async def _prepare_auth_state_with_fallback(
|
||||
context: EventContext,
|
||||
bot: Bot,
|
||||
event_cache: dict | None,
|
||||
route_skip_checks: bool,
|
||||
skip_ban: bool,
|
||||
hook_recorder: HookTraceRecorder,
|
||||
state: dict | None,
|
||||
@@ -1533,7 +1311,6 @@ async def _prepare_auth_state_with_fallback(
|
||||
context=context,
|
||||
bot=bot,
|
||||
event_cache=event_cache,
|
||||
route_skip_checks=route_skip_checks,
|
||||
skip_ban=skip_ban,
|
||||
hook_recorder=hook_recorder,
|
||||
state=state,
|
||||
@@ -1548,7 +1325,6 @@ async def _prepare_auth_state_with_fallback(
|
||||
context=context,
|
||||
bot=bot,
|
||||
event_cache=event_cache,
|
||||
route_skip_checks=route_skip_checks,
|
||||
skip_ban=skip_ban,
|
||||
hook_recorder=hook_recorder,
|
||||
state=state,
|
||||
@@ -1614,16 +1390,26 @@ async def _reserve_limit_side_effect(
|
||||
async def _resolve_cost_gold(
|
||||
*,
|
||||
prep: AuthPreparation,
|
||||
route_skip_checks: bool,
|
||||
hook_recorder: HookTraceRecorder,
|
||||
session: Uninfo,
|
||||
) -> int:
|
||||
plugin = prep.plugin
|
||||
if route_skip_checks or prep.profile.cost_gold <= 0:
|
||||
if prep.profile.cost_gold <= 0:
|
||||
hook_recorder.set("cost_gold", "skipped")
|
||||
return 0
|
||||
cost_start = time.time()
|
||||
try:
|
||||
if prep.user is None:
|
||||
user_start = time.time()
|
||||
prep.user = await with_timeout(
|
||||
UserConsole.get_user(
|
||||
prep.permission_context.user_id,
|
||||
PlatformUtils.get_platform(session),
|
||||
),
|
||||
name="get_cost_user",
|
||||
)
|
||||
prep.permission_context.user = prep.user
|
||||
hook_recorder.set("get_cost_user", f"{time.time() - user_start:.3f}s")
|
||||
cost_gold = await with_timeout(
|
||||
get_plugin_cost(
|
||||
prep.user,
|
||||
@@ -1649,7 +1435,6 @@ async def _run_auth_hooks(
|
||||
prep: AuthPreparation,
|
||||
session: Uninfo,
|
||||
event_cache: dict | None,
|
||||
route_skip_checks: bool,
|
||||
lane_context: AuthLaneContext,
|
||||
hook_recorder: HookTraceRecorder,
|
||||
side_effect_commit: SideEffectCommit,
|
||||
@@ -1660,26 +1445,23 @@ async def _run_auth_hooks(
|
||||
await _enter_hooks_section(lane_context)
|
||||
hook_tasks = []
|
||||
try:
|
||||
if not route_skip_checks:
|
||||
has_limits = await _has_limits_cached(
|
||||
profile.module,
|
||||
event_cache,
|
||||
known=profile.has_limit,
|
||||
)
|
||||
if has_limits:
|
||||
hook_tasks.append(
|
||||
time_hook(
|
||||
_reserve_limit_side_effect(
|
||||
prep=prep,
|
||||
session=session,
|
||||
side_effect_commit=side_effect_commit,
|
||||
),
|
||||
"auth_limit",
|
||||
hook_recorder,
|
||||
)
|
||||
has_limits = await _has_limits_cached(
|
||||
profile.module,
|
||||
event_cache,
|
||||
known=profile.has_limit,
|
||||
)
|
||||
if has_limits:
|
||||
hook_tasks.append(
|
||||
time_hook(
|
||||
_reserve_limit_side_effect(
|
||||
prep=prep,
|
||||
session=session,
|
||||
side_effect_commit=side_effect_commit,
|
||||
),
|
||||
"auth_limit",
|
||||
hook_recorder,
|
||||
)
|
||||
else:
|
||||
hook_recorder.set("auth_limit", "skipped")
|
||||
)
|
||||
else:
|
||||
hook_recorder.set("auth_limit", "skipped")
|
||||
|
||||
@@ -1703,13 +1485,6 @@ async def _run_auth_hooks(
|
||||
return time.time() - hooks_start
|
||||
|
||||
|
||||
async def build_auth_decision_backpressure_report(
|
||||
*,
|
||||
hours: float = 24.0,
|
||||
) -> dict:
|
||||
return await build_auth_observability_report(hours=hours)
|
||||
|
||||
|
||||
_AUTH_PIPELINE_DEPS = AuthPipelineDependencies(
|
||||
route_modules_with_commands=_ROUTE_MODULES_WITH_COMMANDS,
|
||||
get_route_context=_get_route_context,
|
||||
@@ -1726,7 +1501,6 @@ _AUTH_PIPELINE_DEPS = AuthPipelineDependencies(
|
||||
run_auth_hooks=_run_auth_hooks,
|
||||
bot_filter=bot_filter,
|
||||
reserve_gold=reserve_gold,
|
||||
append_auth_decision_log=append_auth_decision_log,
|
||||
insufficient_gold_error=InsufficientGold,
|
||||
logger=logger,
|
||||
log_command=LOGGER_COMMAND,
|
||||
|
||||
@@ -31,8 +31,6 @@ class HandleEventSelectorDependencies:
|
||||
activation_context_from_dispatch: Callable[[EventDispatchContext, Event], Any]
|
||||
new_dispatch_budget: Callable[[], dict[str, int]]
|
||||
dispatch_lane_for_matcher: Callable[[type[Matcher], EventDispatchContext], str]
|
||||
record_activation_result: Callable[[Any], None]
|
||||
debug_activation_shadow: Callable[..., None]
|
||||
merge_dispatch_budget: Callable[[dict[str, int], dict[str, int]], None]
|
||||
build_matcher_state: Callable[[dict], dict]
|
||||
run_selected_matcher: Callable[..., Awaitable[None]]
|
||||
@@ -43,11 +41,81 @@ _ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None
|
||||
_ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {}
|
||||
|
||||
|
||||
def _trim_leading_text(message: Any) -> None:
|
||||
if not message:
|
||||
return
|
||||
segment = message[0]
|
||||
if getattr(segment, "type", None) != "text":
|
||||
return
|
||||
data = getattr(segment, "data", None)
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
data["text"] = str(data.get("text", "")).lstrip("\xa0").lstrip()
|
||||
if not data["text"]:
|
||||
del message[0]
|
||||
|
||||
|
||||
def _is_self_mention_segment(segment: Any, bot: Bot) -> bool:
|
||||
segment_type = getattr(segment, "type", None)
|
||||
if segment_type not in {"mention_user", "group_mention_user"}:
|
||||
return False
|
||||
data = getattr(segment, "data", None)
|
||||
if not isinstance(data, dict):
|
||||
return False
|
||||
if data.get("is_you") or data.get("is_bot"):
|
||||
return True
|
||||
user_id = data.get("user_id")
|
||||
return user_id is not None and str(user_id) == str(bot.self_id)
|
||||
|
||||
|
||||
def _ensure_nonempty_qq_message(message: Any) -> None:
|
||||
if message:
|
||||
return
|
||||
with contextlib.suppress(Exception):
|
||||
message_module = importlib.import_module("nonebot.adapters.qq.message")
|
||||
MessageSegment = getattr(message_module, "MessageSegment")
|
||||
message.append(MessageSegment.text(""))
|
||||
|
||||
|
||||
def _normalize_qq_self_at_message(bot: Bot, event: Event) -> None:
|
||||
"""Remove the leading bot mention left by QQ official @ events.
|
||||
|
||||
nonebot-adapter-qq's @ event branches can mark ``to_me`` but keep the
|
||||
synthetic leading mention segment. Alconna command heads then see
|
||||
``<@bot>命令`` and fail to match, while regular ``event.get_plaintext()``
|
||||
still looks correct. Normalizing here keeps the runtime behavior aligned
|
||||
with OneBot/standard to_me preprocessing without changing plugin code or
|
||||
database state.
|
||||
"""
|
||||
if event.__class__.__name__ not in {
|
||||
"AtMessageCreateEvent",
|
||||
"GroupAtMessageCreateEvent",
|
||||
}:
|
||||
return
|
||||
adapter = getattr(bot, "adapter", None)
|
||||
adapter_name = ""
|
||||
get_name = getattr(adapter, "get_name", None)
|
||||
if callable(get_name):
|
||||
with contextlib.suppress(Exception):
|
||||
adapter_name = str(get_name()).lower()
|
||||
if adapter_name != "qq":
|
||||
return
|
||||
with contextlib.suppress(Exception):
|
||||
message = event.get_message()
|
||||
if not message or not _is_self_mention_segment(message[0], bot):
|
||||
return
|
||||
message.pop(0)
|
||||
setattr(event, "to_me", True)
|
||||
_trim_leading_text(message)
|
||||
_ensure_nonempty_qq_message(message)
|
||||
|
||||
|
||||
async def patched_handle_event(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
deps: HandleEventSelectorDependencies,
|
||||
) -> None:
|
||||
_normalize_qq_self_at_message(bot, event)
|
||||
show_log = True
|
||||
escape_tag = getattr(nb_message, "escape_tag")
|
||||
logger_ = getattr(nb_message, "logger")
|
||||
@@ -153,12 +221,6 @@ async def patched_handle_event(
|
||||
|
||||
if activation_result is not None:
|
||||
selected_matchers = activation_result.selected
|
||||
deps.record_activation_result(activation_result)
|
||||
deps.debug_activation_shadow(
|
||||
priority=priority,
|
||||
activation_result=activation_result,
|
||||
context=dispatch_context,
|
||||
)
|
||||
if (
|
||||
activation_result.candidate_count
|
||||
> deps.overload_selected_threshold
|
||||
@@ -186,12 +248,6 @@ async def patched_handle_event(
|
||||
except Exception:
|
||||
single_result = None
|
||||
if single_result is not None:
|
||||
deps.record_activation_result(single_result)
|
||||
deps.debug_activation_shadow(
|
||||
priority=priority,
|
||||
activation_result=single_result,
|
||||
context=dispatch_context,
|
||||
)
|
||||
deps.merge_dispatch_budget(
|
||||
priority_budget,
|
||||
single_budget,
|
||||
@@ -238,6 +294,7 @@ def install_handle_event_selector(deps: HandleEventSelectorDependencies) -> None
|
||||
for module_name in (
|
||||
"nonebot.adapters.onebot.v11.bot",
|
||||
"nonebot.adapters.onebot.v12.bot",
|
||||
"nonebot.adapters.qq.bot",
|
||||
"onebug.mixin.process",
|
||||
):
|
||||
with contextlib.suppress(Exception):
|
||||
|
||||
@@ -26,13 +26,11 @@ from .auth.context import (
|
||||
)
|
||||
from .auth_checker import (
|
||||
LimitManager,
|
||||
_get_auth_route_precheck_deps,
|
||||
_get_route_context,
|
||||
auth,
|
||||
start_auth_runtime_tasks,
|
||||
stop_auth_runtime_tasks,
|
||||
)
|
||||
from .auth_route import route_precheck
|
||||
|
||||
_SKIP_AUTH_PLUGINS = {"chat_history", "chat_message"}
|
||||
_BOT_CONNECT_TS: float | None = None
|
||||
@@ -113,9 +111,6 @@ async def _auth_preprocessor(
|
||||
)
|
||||
set_route_modules(state, event_context, route_modules)
|
||||
|
||||
if await route_precheck(matcher, event_context, **_get_auth_route_precheck_deps()):
|
||||
return
|
||||
|
||||
try:
|
||||
await auth(
|
||||
matcher,
|
||||
|
||||
@@ -10,7 +10,6 @@ from nonebot.adapters import Bot, Event
|
||||
from nonebot.matcher import Matcher
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.services.message_load import is_overloaded
|
||||
from zhenxun.utils.utils import EntityIDs
|
||||
|
||||
from .auth.context import (
|
||||
@@ -86,7 +85,6 @@ class AuthPipelineContext:
|
||||
event_cache: dict | None = None
|
||||
text: str = ""
|
||||
route_modules: set[str] | None = None
|
||||
route_skip_checks: bool = False
|
||||
is_command_matcher: bool = False
|
||||
lane_context: AuthLaneContext | None = None
|
||||
side_effect_cache: PermissionSideEffectCache | None = None
|
||||
@@ -149,7 +147,6 @@ class AuthPipelineDependencies:
|
||||
run_auth_hooks: Callable[..., Awaitable[float]]
|
||||
bot_filter: Callable[..., None]
|
||||
reserve_gold: Callable[..., Awaitable[Any]]
|
||||
append_auth_decision_log: Callable[..., Awaitable[None]]
|
||||
insufficient_gold_error: type[Exception]
|
||||
logger: Any
|
||||
log_command: str
|
||||
@@ -173,10 +170,7 @@ def apply_policy_precheck(
|
||||
hook_recorder.set("auth_core", f"policy:{decision.reason}")
|
||||
if decision.denied:
|
||||
raise_for_policy(decision, deps.policy_skip_message(decision.reason))
|
||||
if decision.allowed and decision.reason in {
|
||||
"hidden_plugin_skip_auth",
|
||||
"route_miss_skip_checks",
|
||||
}:
|
||||
if decision.allowed and decision.reason in {"hidden_plugin_skip_auth"}:
|
||||
flags.should_return_allowed = True
|
||||
return flags
|
||||
|
||||
@@ -260,17 +254,16 @@ async def route_gate_stage(
|
||||
if ctx.route_modules is None:
|
||||
ctx.route_modules = await deps.get_route_context(ctx.text, ctx.event_cache)
|
||||
set_route_modules(ctx.state, ctx.event_context, ctx.route_modules)
|
||||
ctx.route_skip_checks = (
|
||||
route_missed = (
|
||||
ctx.is_command_matcher
|
||||
and ctx.module in deps.route_modules_with_commands
|
||||
and ctx.module not in ctx.route_modules
|
||||
and not deps.matcher_has_alconna_shortcuts(type(ctx.matcher))
|
||||
)
|
||||
if ctx.route_skip_checks:
|
||||
if route_missed:
|
||||
if ctx.event_cache is not None:
|
||||
ctx.event_cache["route_skip"] = True
|
||||
ctx.event_cache["route_miss_after_native_match"] = True
|
||||
_recorder(ctx).set("route", "miss")
|
||||
ctx.stop(allowed=True, effect="allow", reason="route_miss_skip_checks")
|
||||
|
||||
|
||||
async def prepare_snapshot_stage(
|
||||
@@ -282,7 +275,6 @@ async def prepare_snapshot_stage(
|
||||
context=ctx.event_context,
|
||||
bot=ctx.bot,
|
||||
event_cache=ctx.event_cache,
|
||||
route_skip_checks=ctx.route_skip_checks,
|
||||
skip_ban=ctx.skip_ban,
|
||||
hook_recorder=ctx.hook_recorder,
|
||||
state=ctx.state,
|
||||
@@ -305,7 +297,6 @@ async def policy_precheck_stage(
|
||||
context=ctx.event_context,
|
||||
bot=ctx.bot,
|
||||
event_cache=ctx.event_cache,
|
||||
route_skip_checks=ctx.route_skip_checks,
|
||||
skip_ban=ctx.skip_ban,
|
||||
hook_recorder=ctx.hook_recorder,
|
||||
state=ctx.state,
|
||||
@@ -340,7 +331,6 @@ async def policy_precheck_stage(
|
||||
)
|
||||
ctx.cost_gold = await deps.resolve_cost_gold(
|
||||
prep=ctx.prep,
|
||||
route_skip_checks=ctx.route_skip_checks,
|
||||
hook_recorder=ctx.hook_recorder,
|
||||
session=ctx.session,
|
||||
)
|
||||
@@ -356,7 +346,6 @@ async def legacy_hook_adapter_stage(
|
||||
prep=prep,
|
||||
session=ctx.session,
|
||||
event_cache=ctx.event_cache,
|
||||
route_skip_checks=ctx.route_skip_checks,
|
||||
lane_context=_lane_context(ctx),
|
||||
hook_recorder=_recorder(ctx),
|
||||
side_effect_commit=_side_effect_commit(ctx),
|
||||
@@ -425,35 +414,12 @@ async def decision_log_stage(
|
||||
ctx.auth_allowed,
|
||||
None if ctx.auth_allowed else ctx.decision_reason,
|
||||
)
|
||||
side_effect_state = commit.snapshot() if commit is not None else None
|
||||
shadow_effect = None
|
||||
shadow_reason = None
|
||||
if has_deferred_commit:
|
||||
shadow_effect = "defer"
|
||||
shadow_reason = "side_effect_pending:" + ",".join(
|
||||
commit.pending_kinds if commit is not None else ()
|
||||
)
|
||||
if ctx.entered_side_effect_lock and ctx.side_effect_lock is not None:
|
||||
try:
|
||||
ctx.side_effect_lock.release()
|
||||
except Exception:
|
||||
pass
|
||||
ctx.entered_side_effect_lock = False
|
||||
latency_ms = (time.time() - ctx.start_time) * 1000
|
||||
await deps.append_auth_decision_log(
|
||||
bot_id=ctx.event_context.bot_id,
|
||||
platform=ctx.event_context.platform,
|
||||
group_id=_entity(ctx).group_id,
|
||||
user_id=_entity(ctx).user_id,
|
||||
module=ctx.module,
|
||||
effect=ctx.decision_effect or "error",
|
||||
reason=ctx.decision_reason,
|
||||
shadow_effect=shadow_effect,
|
||||
shadow_reason=shadow_reason,
|
||||
side_effect_state=side_effect_state,
|
||||
latency_ms=latency_ms,
|
||||
overloaded=is_overloaded(),
|
||||
)
|
||||
|
||||
|
||||
def build_auth_pipeline(deps: AuthPipelineDependencies) -> AuthPipeline:
|
||||
|
||||
@@ -60,7 +60,6 @@ class PolicyResource:
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PolicyContext:
|
||||
snapshot: AuthSnapshot
|
||||
route_skip_checks: bool = False
|
||||
allow_sleep_bypass: bool = False
|
||||
allow_group_sleep_bypass: bool = False
|
||||
|
||||
@@ -101,8 +100,6 @@ class PolicyDecisionPoint:
|
||||
profile = resource.profile
|
||||
if profile.hidden:
|
||||
return PolicyDecision("allow", "hidden_plugin_skip_auth")
|
||||
if context.route_skip_checks:
|
||||
return PolicyDecision("allow", "route_miss_skip_checks")
|
||||
if snapshot.ban_state is True and not principal.is_superuser:
|
||||
return PolicyDecision("deny", "user_or_group_banned")
|
||||
if profile.superuser_only and not principal.is_superuser:
|
||||
|
||||
@@ -2,11 +2,13 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
PluginLimitMemoryCache,
|
||||
from zhenxun.utils.enum import BlockType, PluginType
|
||||
|
||||
from .auth.data_provider import (
|
||||
DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
PermissionDataProvider,
|
||||
PluginLimitSnapshot,
|
||||
)
|
||||
from zhenxun.utils.enum import BlockType, PluginType
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
@@ -81,6 +83,7 @@ async def get_plugin_auth_profile(
|
||||
*,
|
||||
event_cache: dict | None = None,
|
||||
allow_cache_load: bool = True,
|
||||
provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
) -> PluginAuthProfile:
|
||||
module = str(getattr(plugin, "module", "") or "")
|
||||
profile_cache: dict[str, PluginAuthProfile] = {}
|
||||
@@ -98,10 +101,10 @@ async def get_plugin_auth_profile(
|
||||
limits = limit_cache[module]
|
||||
limits_ready = True
|
||||
if limits is None:
|
||||
limits = PluginLimitMemoryCache.get_limits_if_ready(module)
|
||||
limits = provider.get_module_limits_if_ready(module)
|
||||
limits_ready = limits is not None
|
||||
if limits is None and allow_cache_load:
|
||||
limits = await PluginLimitMemoryCache.get_limits(module)
|
||||
limits = await provider.get_module_limits(module)
|
||||
limits_ready = True
|
||||
if limits is None:
|
||||
limits = []
|
||||
|
||||
@@ -1,60 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from nonebot.matcher import Matcher
|
||||
|
||||
from zhenxun.utils.enum import PluginType
|
||||
|
||||
from .auth.context import EventContext, set_route_modules
|
||||
|
||||
RouteContextGetter = Callable[[str, dict | None], Awaitable[set[str]]]
|
||||
CommandMatcherChecker = Callable[[type[Matcher]], bool]
|
||||
AlconnaShortcutChecker = Callable[[type[Matcher]], bool]
|
||||
|
||||
|
||||
async def route_precheck(
|
||||
matcher: Matcher,
|
||||
context: EventContext,
|
||||
*,
|
||||
route_modules_with_commands: set[str],
|
||||
get_route_context: RouteContextGetter,
|
||||
is_command_matcher_class: CommandMatcherChecker,
|
||||
matcher_has_alconna_shortcuts: AlconnaShortcutChecker,
|
||||
) -> bool:
|
||||
"""Skip expensive auth checks for command matchers proven to be off-route."""
|
||||
|
||||
module = matcher.plugin_name or ""
|
||||
if not module:
|
||||
return False
|
||||
if _is_hidden_plugin(matcher):
|
||||
return False
|
||||
if not is_command_matcher_class(type(matcher)):
|
||||
return False
|
||||
|
||||
route_modules = context.route_modules if context.route_modules_loaded else None
|
||||
if route_modules is None:
|
||||
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 context.event_cache is not None:
|
||||
context.event_cache["route_skip"] = True
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _is_hidden_plugin(matcher: Matcher) -> bool:
|
||||
plugin = matcher.plugin
|
||||
if not plugin or not plugin.metadata:
|
||||
return False
|
||||
extra = plugin.metadata.extra or {}
|
||||
return extra.get("plugin_type") == PluginType.HIDDEN
|
||||
|
||||
|
||||
__all__ = ["route_precheck"]
|
||||
@@ -1,24 +1,120 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from dataclasses import dataclass, field
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from zhenxun.services.cache.runtime_cache import (
|
||||
BanMemoryCache,
|
||||
BotMemoryCache,
|
||||
BotSnapshot,
|
||||
GroupMemoryCache,
|
||||
GroupSnapshot,
|
||||
LevelUserMemoryCache,
|
||||
LevelUserSnapshot,
|
||||
)
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .auth.config import LOGGER_COMMAND
|
||||
from .auth.context import EventContext
|
||||
from .auth.data_provider import (
|
||||
DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
PermissionDataProvider,
|
||||
)
|
||||
from .auth_profile import PluginAuthProfile
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nonebot.adapters import Bot
|
||||
|
||||
QQ_CLIENT_GROUP_REPAIR_TTL = 60
|
||||
_QQ_CLIENT_GROUP_REPAIR_FAILURES: dict[tuple[str, str], float] = {}
|
||||
_QQ_CLIENT_GROUP_REPAIR_LOCKS: dict[tuple[str, str], asyncio.Lock] = {}
|
||||
|
||||
|
||||
def _build_runtime_group_snapshot(context: EventContext) -> GroupSnapshot | None:
|
||||
"""Provide a non-persistent default group for QQ official runtime auth."""
|
||||
if context.platform_scope != "qq_api" or not context.group_id:
|
||||
return None
|
||||
return GroupSnapshot(
|
||||
group_id=context.group_id,
|
||||
channel_id=context.channel_id,
|
||||
group_name="",
|
||||
max_member_count=0,
|
||||
member_count=0,
|
||||
status=True,
|
||||
level=5,
|
||||
is_super=False,
|
||||
group_flag=0,
|
||||
block_plugin="",
|
||||
superuser_block_plugin="",
|
||||
block_task="",
|
||||
superuser_block_task="",
|
||||
platform=context.platform,
|
||||
)
|
||||
|
||||
|
||||
def _qq_client_group_repair_key(context: EventContext) -> tuple[str, str] | None:
|
||||
if context.platform_scope != "qq_client" or not context.group_id:
|
||||
return None
|
||||
return (context.group_id, context.channel_id or "")
|
||||
|
||||
|
||||
def _qq_client_group_repair_on_cooldown(key: tuple[str, str]) -> bool:
|
||||
expire_at = _QQ_CLIENT_GROUP_REPAIR_FAILURES.get(key)
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
_QQ_CLIENT_GROUP_REPAIR_FAILURES.pop(key, None)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
async def _repair_missing_qq_client_group(
|
||||
context: EventContext,
|
||||
*,
|
||||
provider: PermissionDataProvider,
|
||||
) -> GroupSnapshot | None:
|
||||
"""Persist a minimal OneBot group when startup group sync returned empty."""
|
||||
key = _qq_client_group_repair_key(context)
|
||||
if key is None or not provider.group_cache_loaded():
|
||||
return None
|
||||
group_id, _ = key
|
||||
if _qq_client_group_repair_on_cooldown(key):
|
||||
return None
|
||||
|
||||
try:
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
|
||||
lock = _QQ_CLIENT_GROUP_REPAIR_LOCKS.setdefault(key, asyncio.Lock())
|
||||
async with lock:
|
||||
existing = provider.get_group_if_ready(
|
||||
group_id,
|
||||
context.channel_id,
|
||||
)
|
||||
if existing is not None:
|
||||
return existing
|
||||
defaults = {
|
||||
"group_name": "",
|
||||
"max_member_count": 0,
|
||||
"member_count": 0,
|
||||
"group_flag": 1,
|
||||
"platform": context.platform,
|
||||
}
|
||||
group, _ = await GroupConsole.get_or_create_root_group(
|
||||
group_id=group_id,
|
||||
defaults=defaults,
|
||||
)
|
||||
from zhenxun.services.cache.runtime_cache import GroupMemoryCache
|
||||
|
||||
await GroupMemoryCache.upsert_from_model(group)
|
||||
return GroupSnapshot.from_model(group)
|
||||
except Exception as exc:
|
||||
_QQ_CLIENT_GROUP_REPAIR_FAILURES[key] = time.time() + QQ_CLIENT_GROUP_REPAIR_TTL
|
||||
logger.warning(
|
||||
"协议端群记录缺失自愈失败,已短期跳过重复修复",
|
||||
LOGGER_COMMAND,
|
||||
group_id=context.group_id,
|
||||
e=exc,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class AuthSnapshot:
|
||||
@@ -72,6 +168,7 @@ async def build_auth_snapshot(
|
||||
bot: "Bot",
|
||||
skip_ban: bool = False,
|
||||
allow_cache_load: bool = False,
|
||||
provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
) -> AuthSnapshot:
|
||||
event_cache = context.event_cache
|
||||
entity = context.entity
|
||||
@@ -85,15 +182,15 @@ async def build_auth_snapshot(
|
||||
):
|
||||
bot_data = event_cache.get("bot_data")
|
||||
else:
|
||||
bot_data = BotMemoryCache.get_if_ready(bot.self_id)
|
||||
bot_data = provider.get_bot_if_ready(bot.self_id)
|
||||
if bot_data is None:
|
||||
if allow_cache_load:
|
||||
bot_data = await BotMemoryCache.get(bot.self_id)
|
||||
elif not BotMemoryCache.is_loaded():
|
||||
bot_data = await provider.get_bot(bot.self_id)
|
||||
elif not provider.bot_cache_loaded():
|
||||
cache_misses.add("bot")
|
||||
if event_cache is not None:
|
||||
event_cache["bot_data"] = bot_data
|
||||
event_cache["bot_cache_ready"] = BotMemoryCache.is_loaded()
|
||||
event_cache["bot_cache_ready"] = provider.bot_cache_loaded()
|
||||
|
||||
group = None
|
||||
if entity.group_id:
|
||||
@@ -104,14 +201,31 @@ async def build_auth_snapshot(
|
||||
):
|
||||
group = event_cache.get("group")
|
||||
else:
|
||||
group = GroupMemoryCache.get_if_ready(entity.group_id, entity.channel_id)
|
||||
if group is None and not GroupMemoryCache.is_loaded():
|
||||
group = provider.get_group_if_ready(entity.group_id, entity.channel_id)
|
||||
if group is None and not provider.group_cache_loaded():
|
||||
cache_misses.add("group")
|
||||
elif group is None and allow_cache_load:
|
||||
group = await GroupMemoryCache.get(entity.group_id, entity.channel_id)
|
||||
group = await provider.get_group(entity.group_id, entity.channel_id)
|
||||
if event_cache is not None:
|
||||
event_cache["group"] = group
|
||||
event_cache["group_cache_ready"] = GroupMemoryCache.is_loaded()
|
||||
event_cache["group_cache_ready"] = provider.group_cache_loaded()
|
||||
if group is None:
|
||||
group = await _repair_missing_qq_client_group(
|
||||
context,
|
||||
provider=provider,
|
||||
)
|
||||
if group is None and (runtime_group := _build_runtime_group_snapshot(context)):
|
||||
group = runtime_group
|
||||
cache_misses.discard("group")
|
||||
if event_cache is not None:
|
||||
event_cache["group"] = group
|
||||
event_cache["group_cache_ready"] = True
|
||||
event_cache["group_runtime_virtual"] = True
|
||||
elif group is not None:
|
||||
cache_misses.discard("group")
|
||||
if event_cache is not None:
|
||||
event_cache["group"] = group
|
||||
event_cache["group_cache_ready"] = True
|
||||
|
||||
admin_levels = None
|
||||
if profile.need_admin:
|
||||
@@ -122,13 +236,13 @@ async def build_auth_snapshot(
|
||||
):
|
||||
admin_levels = event_cache.get("admin_levels")
|
||||
else:
|
||||
admin_levels = LevelUserMemoryCache.get_levels_if_ready(
|
||||
admin_levels = provider.get_admin_levels_if_ready(
|
||||
entity.user_id,
|
||||
entity.group_id,
|
||||
)
|
||||
if admin_levels is None:
|
||||
if allow_cache_load:
|
||||
admin_levels = await LevelUserMemoryCache.get_levels(
|
||||
admin_levels = await provider.get_admin_levels(
|
||||
entity.user_id,
|
||||
entity.group_id,
|
||||
)
|
||||
@@ -136,19 +250,19 @@ async def build_auth_snapshot(
|
||||
cache_misses.add("admin_levels")
|
||||
if event_cache is not None:
|
||||
event_cache["admin_levels"] = admin_levels
|
||||
event_cache["admin_cache_ready"] = LevelUserMemoryCache.is_loaded()
|
||||
event_cache["admin_cache_ready"] = provider.admin_cache_loaded()
|
||||
|
||||
ban_state = None
|
||||
if not skip_ban:
|
||||
if event_cache is not None and "ban_state" in event_cache:
|
||||
ban_state = event_cache.get("ban_state")
|
||||
elif BanMemoryCache.is_loaded():
|
||||
ban_state = BanMemoryCache.is_banned(entity.user_id, entity.group_id)
|
||||
elif provider.ban_cache_loaded():
|
||||
ban_state = provider.is_banned(entity.user_id, entity.group_id)
|
||||
if event_cache is not None:
|
||||
event_cache["ban_state"] = ban_state
|
||||
elif allow_cache_load:
|
||||
await BanMemoryCache.ensure_loaded()
|
||||
ban_state = BanMemoryCache.is_banned(entity.user_id, entity.group_id)
|
||||
await provider.ensure_ban_loaded()
|
||||
ban_state = provider.is_banned(entity.user_id, entity.group_id)
|
||||
if event_cache is not None:
|
||||
event_cache["ban_state"] = ban_state
|
||||
else:
|
||||
@@ -174,6 +288,7 @@ async def get_or_build_auth_snapshot(
|
||||
bot: "Bot",
|
||||
skip_ban: bool = False,
|
||||
allow_cache_load: bool = False,
|
||||
provider: PermissionDataProvider = DEFAULT_PERMISSION_DATA_PROVIDER,
|
||||
) -> AuthSnapshot:
|
||||
event_cache = context.event_cache
|
||||
module = profile.module
|
||||
@@ -190,6 +305,7 @@ async def get_or_build_auth_snapshot(
|
||||
bot=bot,
|
||||
skip_ban=skip_ban,
|
||||
allow_cache_load=allow_cache_load,
|
||||
provider=provider,
|
||||
)
|
||||
if event_cache is not None:
|
||||
event_cache.setdefault("auth_snapshots", {})[module] = snapshot
|
||||
|
||||
@@ -1,11 +1,9 @@
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from nonebot.adapters import Bot, Message
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.bot_message_store import BotMessageStore
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import BotSentType
|
||||
from zhenxun.utils.log_sanitizer import sanitize_for_logging
|
||||
from zhenxun.utils.manager.message_manager import MessageManager
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
@@ -45,13 +43,15 @@ def replace_message(message: Message) -> str:
|
||||
async def handle_api_result(
|
||||
bot: Bot, exception: Exception | None, api: str, data: dict[str, Any], result: Any
|
||||
):
|
||||
if exception or api != "send_msg":
|
||||
if (
|
||||
exception
|
||||
or api != "send_msg"
|
||||
or PlatformUtils.get_platform_scope(bot) != "qq_client"
|
||||
):
|
||||
return
|
||||
user_id = data.get("user_id")
|
||||
group_id = data.get("group_id")
|
||||
message_id = result.get("message_id")
|
||||
message_id = result.get("message_id") if isinstance(result, Mapping) else None
|
||||
message: Message = data.get("message", "")
|
||||
message_type = data.get("message_type")
|
||||
try:
|
||||
if user_id and message_id:
|
||||
MessageManager.add(str(user_id), str(message_id))
|
||||
@@ -62,27 +62,5 @@ async def handle_api_result(
|
||||
logger.warning(
|
||||
f"收集消息id发生错误...data: {data}, result: {result}", LOG_COMMAND, e=e
|
||||
)
|
||||
if not Config.get_config("hook", "RECORD_BOT_SENT_MESSAGES"):
|
||||
return
|
||||
try:
|
||||
await BotMessageStore.append_buffered(
|
||||
bot_id=bot.self_id,
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
sent_type=BotSentType.GROUP
|
||||
if message_type == "group"
|
||||
else BotSentType.PRIVATE,
|
||||
text=replace_message(message),
|
||||
plain_text=message.extract_plain_text()
|
||||
if isinstance(message, Message)
|
||||
else replace_message(message),
|
||||
platform=PlatformUtils.get_platform(bot),
|
||||
)
|
||||
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
|
||||
logger.debug(f"消息发送记录,message: {sanitized_message}")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"消息发送记录发生错误...data: {data}, result: {result}",
|
||||
LOG_COMMAND,
|
||||
e=e,
|
||||
)
|
||||
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
|
||||
logger.debug(f"消息发送记录,message: {sanitized_message}")
|
||||
|
||||
@@ -30,7 +30,7 @@ async def _(bot: Bot):
|
||||
参数:
|
||||
bot: Bot
|
||||
"""
|
||||
if PlatformUtils.get_platform(bot) != "qq":
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
return
|
||||
|
||||
logger.debug(f"更新Bot: {bot.self_id} 的群认证...", "群认证同步")
|
||||
@@ -44,6 +44,13 @@ async def _(bot: Bot):
|
||||
)
|
||||
return
|
||||
|
||||
if not current_group_list:
|
||||
logger.warning(
|
||||
f"Bot: {bot.self_id} 未获取到任何群组,"
|
||||
"本次不会创建群认证;后续群消息将尝试按事件自愈。",
|
||||
"群认证同步",
|
||||
)
|
||||
|
||||
db_group_list: list[str] = await GroupConsole.all().values_list(
|
||||
"group_id", flat=True
|
||||
) # pyright: ignore[reportAssignmentType]
|
||||
|
||||
@@ -8,7 +8,7 @@ from nonebot_plugin_apscheduler import scheduler
|
||||
from zhenxun.configs.utils import PluginExtraData, Task
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.task_info import TaskInfo
|
||||
from zhenxun.services.cache.runtime_cache import TaskInfoMemoryCache
|
||||
from zhenxun.services.cache.runtime_cache import GroupMemoryCache, TaskInfoMemoryCache
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
@@ -63,6 +63,8 @@ async def update_to_group(create_list: list[tuple[bool, TaskInfo]]):
|
||||
)
|
||||
group.block_task = CommonUtils.convert_module_format(block_tasks)
|
||||
await GroupConsole.bulk_update(group_list, ["block_task"], 10)
|
||||
for group in group_list:
|
||||
await GroupMemoryCache.upsert_from_model(group)
|
||||
|
||||
|
||||
async def to_db(
|
||||
@@ -146,8 +148,10 @@ async def _():
|
||||
for plugin in get_loaded_plugins():
|
||||
await _handle_setting(plugin, task_info_list, task_list)
|
||||
if not task_info_list:
|
||||
await TaskInfo.all().update(load_status=False)
|
||||
await TaskInfoMemoryCache.refresh()
|
||||
logger.warning(
|
||||
"未扫描到任何被动技能,跳过 TaskInfo.load_status 全量关闭,"
|
||||
"避免插件加载异常时误关闭全部被动技能。",
|
||||
)
|
||||
return
|
||||
module_dict = {t[1]: t[0] for t in await TaskInfo.all().values_list("id", "module")}
|
||||
load_task = []
|
||||
|
||||
@@ -2,6 +2,7 @@ from pathlib import Path
|
||||
|
||||
import nonebot
|
||||
|
||||
from zhenxun.configs.config import BotConfig
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
path = Path(__file__).parent
|
||||
@@ -15,11 +16,12 @@ except ImportError:
|
||||
logger.warning("未安装 onebot-adapter,无法加载QQ平台专用插件...")
|
||||
|
||||
|
||||
try:
|
||||
from nonebot.adapters.qq import ( # noqa: F401 # pyright: ignore [reportMissingImports]
|
||||
Bot,
|
||||
)
|
||||
if BotConfig.qq_adapter_load:
|
||||
try:
|
||||
from nonebot.adapters.qq import ( # noqa: F401 # pyright: ignore [reportMissingImports]
|
||||
Bot,
|
||||
)
|
||||
|
||||
nonebot.load_plugins(str((path / "qq_api").resolve()))
|
||||
except ImportError:
|
||||
logger.warning("未安装 qq-adapter,无法加载QQ官平台专用插件...")
|
||||
nonebot.load_plugins(str((path / "qq_api").resolve()))
|
||||
except ImportError:
|
||||
logger.warning("未安装 qq-adapter,无法加载QQ官平台专用插件...")
|
||||
|
||||
@@ -108,9 +108,7 @@ async def _(
|
||||
):
|
||||
if session.user.id == bot.self_id:
|
||||
"""新成员为bot本身"""
|
||||
group, _ = await GroupConsole.get_or_create(
|
||||
group_id=str(event.group_id), channel_id__isnull=True
|
||||
)
|
||||
group, _ = await GroupConsole.get_or_create_root_group(str(event.group_id))
|
||||
try:
|
||||
await GroupManager.add_bot(
|
||||
bot, str(event.operator_id), str(event.group_id), group
|
||||
|
||||
@@ -124,7 +124,7 @@ class GroupManager:
|
||||
group_id=group_id,
|
||||
)
|
||||
return
|
||||
await GroupConsole.update_or_create(
|
||||
await GroupConsole.get_or_create_root_group(
|
||||
group_id=group_info["group_id"],
|
||||
defaults={
|
||||
"group_name": group_info["group_name"],
|
||||
@@ -134,6 +134,7 @@ class GroupManager:
|
||||
"block_plugin": block_plugin,
|
||||
"platform": "qq",
|
||||
},
|
||||
update_defaults=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -1,14 +1,13 @@
|
||||
"""QQ official platform observer.
|
||||
|
||||
Official QQ identifiers are not in the same namespace as OneBot QQ numbers.
|
||||
This observer intentionally avoids writing legacy identity tables; runtime auth
|
||||
uses a non-persistent group snapshot when needed.
|
||||
"""
|
||||
|
||||
from nonebot import on_message
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.models.friend_user import FriendUser
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.group_member_info import GroupInfoUser
|
||||
from zhenxun.services.hot_query_cache import (
|
||||
invalidate_group_members,
|
||||
invalidate_member_names,
|
||||
)
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
|
||||
@@ -20,21 +19,5 @@ _matcher = on_message(priority=999, block=False, rule=rule)
|
||||
|
||||
|
||||
@_matcher.handle()
|
||||
async def _(session: Uninfo):
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
if session.group:
|
||||
if not await GroupConsole.exists(group_id=session.group.id):
|
||||
await GroupConsole.create(group_id=session.group.id)
|
||||
logger.info("添加当前群组ID信息", session=session)
|
||||
await GroupInfoUser.update_or_create(
|
||||
user_id=session.user.id,
|
||||
group_id=session.group.id,
|
||||
platform=PlatformUtils.get_platform(session),
|
||||
)
|
||||
await invalidate_group_members(session.group.id, [session.user.id])
|
||||
await invalidate_member_names([session.user.id])
|
||||
elif not await FriendUser.exists(user_id=session.user.id, platform=platform):
|
||||
await FriendUser.create(
|
||||
user_id=session.user.id, platform=PlatformUtils.get_platform(session)
|
||||
)
|
||||
logger.info("添加当前好友用户信息", "", session=session)
|
||||
async def _():
|
||||
return
|
||||
|
||||
@@ -203,7 +203,7 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
|
||||
session=event.user_id,
|
||||
target=event.group_id,
|
||||
)
|
||||
group, _ = await GroupConsole.update_or_create(
|
||||
group, _ = await GroupConsole.get_or_create_root_group(
|
||||
group_id=str(event.group_id),
|
||||
defaults={
|
||||
"group_name": "",
|
||||
@@ -211,6 +211,7 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
|
||||
"member_count": 0,
|
||||
"group_flag": 1,
|
||||
},
|
||||
update_defaults=True,
|
||||
)
|
||||
await bot.set_group_add_request(
|
||||
flag=event.flag, sub_type="invite", approve=True
|
||||
|
||||
@@ -18,6 +18,8 @@ async def _():
|
||||
return
|
||||
bots = nonebot.get_bots()
|
||||
for bot in bots.values():
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
continue
|
||||
try:
|
||||
await PlatformUtils.update_group(bot)
|
||||
except Exception as e:
|
||||
@@ -36,6 +38,8 @@ async def _():
|
||||
return
|
||||
bots = nonebot.get_bots()
|
||||
for bot in bots.values():
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
continue
|
||||
try:
|
||||
await PlatformUtils.update_friend(bot)
|
||||
except Exception as e:
|
||||
|
||||
@@ -52,8 +52,8 @@ async def _():
|
||||
if last_message:
|
||||
now = datetime.now(pytz.timezone("Asia/Shanghai"))
|
||||
if now - timedelta(days=2) > last_message.create_time:
|
||||
_group, _ = await GroupConsole.get_or_create(
|
||||
group_id=group.group_id, channel_id__isnull=True
|
||||
_group, _ = await GroupConsole.get_or_create_root_group(
|
||||
group.group_id
|
||||
)
|
||||
modules = [f"<{module}" for module in modules]
|
||||
_group.block_task = ",".join(modules) + "," # type: ignore
|
||||
|
||||
@@ -147,7 +147,7 @@ def CheckGroupId():
|
||||
@_matcher.assign("modify-level", parameterless=[CheckGroupId()])
|
||||
async def _(session: EventSession, arparma: Arparma, state: T_State, level: int):
|
||||
gid = state["group_id"]
|
||||
group, _ = await GroupConsole.get_or_create(group_id=gid)
|
||||
group, _ = await GroupConsole.get_or_create_root_group(gid)
|
||||
old_level = group.level
|
||||
group.level = level
|
||||
await group.save(update_fields=["level"])
|
||||
@@ -176,10 +176,10 @@ async def _(session: EventSession, arparma: Arparma, state: T_State):
|
||||
@_matcher.assign("auth-handle", parameterless=[CheckGroupId()])
|
||||
async def _(session: EventSession, arparma: Arparma, state: T_State):
|
||||
gid = state["group_id"]
|
||||
await GroupConsole.update_or_create(
|
||||
await GroupConsole.get_or_create_root_group(
|
||||
group_id=gid,
|
||||
channel_id__isnull=True,
|
||||
defaults={"group_flag": 0 if arparma.find("delete") else 1},
|
||||
update_defaults=True,
|
||||
)
|
||||
s = "删除" if arparma.find("delete") else "添加"
|
||||
await MessageUtils.build_message(f"{s}群认证成功!").send(reply_to=True)
|
||||
|
||||
@@ -54,6 +54,11 @@ async def _(
|
||||
arparma: Arparma,
|
||||
):
|
||||
try:
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
await MessageUtils.build_message(
|
||||
"当前平台不支持旧群组信息同步,仅 OneBot 协议端可用。"
|
||||
).send()
|
||||
return
|
||||
num = await PlatformUtils.update_group(bot)
|
||||
logger.info(
|
||||
f"更新群聊信息完成,共更新了 {num} 个群组的信息!",
|
||||
@@ -75,6 +80,11 @@ async def _(
|
||||
arparma: Arparma,
|
||||
):
|
||||
try:
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
await MessageUtils.build_message(
|
||||
"当前平台不支持旧好友信息同步,仅 OneBot 协议端可用。"
|
||||
).send()
|
||||
return
|
||||
num = await PlatformUtils.update_friend(bot)
|
||||
logger.info(
|
||||
f"更新好友信息完成,共更新了 {num} 个好友的信息!",
|
||||
|
||||
@@ -213,9 +213,10 @@ async def _(param: HandleRequest) -> Result:
|
||||
group.group_flag = 1
|
||||
await group.save(update_fields=["group_flag"])
|
||||
else:
|
||||
await GroupConsole.update_or_create(
|
||||
await GroupConsole.get_or_create_root_group(
|
||||
group_id=req.group_id,
|
||||
defaults={"group_flag": 1},
|
||||
update_defaults=True,
|
||||
)
|
||||
try:
|
||||
await FgRequest.approve(bot, param.id)
|
||||
@@ -240,8 +241,7 @@ async def _(param: HandleRequest) -> Result:
|
||||
async def _(param: LeaveGroup) -> Result:
|
||||
try:
|
||||
bot = nonebot.get_bot(param.bot_id)
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
if platform != "qq":
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
return Result.warning_("该平台不支持退群操作...")
|
||||
group_list, _ = await PlatformUtils.get_group_list(bot)
|
||||
if param.group_id not in [g.group_id for g in group_list]:
|
||||
@@ -265,8 +265,7 @@ async def _(param: LeaveGroup) -> Result:
|
||||
async def _(param: DeleteFriend) -> Result:
|
||||
try:
|
||||
bot = nonebot.get_bot(param.bot_id)
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
if platform != "qq":
|
||||
if PlatformUtils.get_platform_scope(bot) != "qq_client":
|
||||
return Result.warning_("该平台不支持删除好友操作...")
|
||||
friend_list, _ = await PlatformUtils.get_friend_list(bot)
|
||||
if param.user_id not in [f.user_id for f in friend_list]:
|
||||
|
||||
@@ -40,7 +40,7 @@ def reply_check() -> Rule:
|
||||
if event.get_type() == "message":
|
||||
return (
|
||||
bool(await reply_fetch(event, bot))
|
||||
and PlatformUtils.get_platform(session) == "qq"
|
||||
and PlatformUtils.get_platform_scope(session) == "qq_client"
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
@@ -1,113 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections import deque
|
||||
import contextlib
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .bot_message_store import BotMessageStore
|
||||
|
||||
LOG_COMMAND = "BotMessageStore"
|
||||
|
||||
_BUFFER_MAX_RETAIN = 20_000
|
||||
_FLUSH_TRIGGER_SIZE = 64
|
||||
_FLUSH_BATCH_SIZE = 500
|
||||
_FLUSH_INTERVAL_SECONDS = 5.0
|
||||
_DROP_LOG_INTERVAL_SECONDS = 10.0
|
||||
|
||||
_buffer: deque[BotMessageStore] = deque()
|
||||
_buffer_lock = asyncio.Lock()
|
||||
_flush_lock = asyncio.Lock()
|
||||
_flush_task: asyncio.Task[None] | None = None
|
||||
_dropped = 0
|
||||
_last_drop_log_at = 0.0
|
||||
|
||||
|
||||
def _ensure_flush_task() -> None:
|
||||
global _flush_task
|
||||
if _flush_task is not None and not _flush_task.done():
|
||||
return
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return
|
||||
_flush_task = loop.create_task(_flush_loop())
|
||||
|
||||
|
||||
def _record_drop() -> None:
|
||||
global _dropped, _last_drop_log_at
|
||||
_dropped += 1
|
||||
now = time.monotonic()
|
||||
if now - _last_drop_log_at < _DROP_LOG_INTERVAL_SECONDS:
|
||||
return
|
||||
_last_drop_log_at = now
|
||||
logger.warning(
|
||||
f"bot_message_store buffer full, dropped {_dropped} records, "
|
||||
f"backlog={len(_buffer)}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
async def _flush_loop() -> None:
|
||||
while True:
|
||||
await asyncio.sleep(_FLUSH_INTERVAL_SECONDS)
|
||||
try:
|
||||
await flush_bot_message_store_buffer("定时")
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.warning("定时批量写入 Bot 发送记录失败", LOG_COMMAND, e=exc)
|
||||
|
||||
|
||||
async def append_bot_message_store_record(record: BotMessageStore) -> None:
|
||||
_ensure_flush_task()
|
||||
async with _buffer_lock:
|
||||
if len(_buffer) >= _BUFFER_MAX_RETAIN:
|
||||
_buffer.popleft()
|
||||
_record_drop()
|
||||
_buffer.append(record)
|
||||
should_flush = len(_buffer) >= _FLUSH_TRIGGER_SIZE and not _flush_lock.locked()
|
||||
if should_flush:
|
||||
await flush_bot_message_store_buffer("缓冲区触发")
|
||||
|
||||
|
||||
async def flush_bot_message_store_buffer(reason: str) -> int:
|
||||
from .bot_message_store import BotMessageStore
|
||||
|
||||
async with _flush_lock:
|
||||
written = 0
|
||||
while True:
|
||||
batch: list[BotMessageStore] = []
|
||||
async with _buffer_lock:
|
||||
while _buffer and len(batch) < _FLUSH_BATCH_SIZE:
|
||||
batch.append(_buffer.popleft())
|
||||
if not batch:
|
||||
break
|
||||
try:
|
||||
await BotMessageStore.bulk_create(batch, batch_size=_FLUSH_BATCH_SIZE)
|
||||
except Exception as exc:
|
||||
async with _buffer_lock:
|
||||
retain_count = max(_BUFFER_MAX_RETAIN - len(_buffer), 0)
|
||||
for record in reversed(batch[-retain_count:]):
|
||||
_buffer.appendleft(record)
|
||||
logger.error(f"{reason}批量写入 Bot 发送记录失败", LOG_COMMAND, e=exc)
|
||||
return written
|
||||
written += len(batch)
|
||||
if written:
|
||||
logger.debug(f"{reason}批量写入 Bot 发送记录 {written} 条", LOG_COMMAND)
|
||||
return written
|
||||
|
||||
|
||||
async def stop_bot_message_store_buffer() -> int:
|
||||
global _flush_task
|
||||
task = _flush_task
|
||||
_flush_task = None
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
with contextlib.suppress(BaseException):
|
||||
await task
|
||||
return await flush_bot_message_store_buffer("关闭")
|
||||
@@ -1,49 +0,0 @@
|
||||
from typing import ClassVar
|
||||
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.db_context import Model
|
||||
|
||||
|
||||
class AuthDecisionLog(Model):
|
||||
id = fields.IntField(pk=True, generated=True, auto_increment=True)
|
||||
"""自增id"""
|
||||
bot_id = fields.CharField(255, null=True, description="Bot ID")
|
||||
"""Bot ID"""
|
||||
platform = fields.CharField(64, null=True, description="平台")
|
||||
"""平台"""
|
||||
group_id = fields.CharField(255, null=True, description="群组id")
|
||||
"""群组id"""
|
||||
user_id = fields.CharField(255, null=True, description="用户id")
|
||||
"""用户id"""
|
||||
module = fields.CharField(255, null=True, description="插件模块")
|
||||
"""插件模块"""
|
||||
effect = fields.CharField(32, description="决策结果")
|
||||
"""决策结果 allow/deny/skip/defer/error"""
|
||||
reason = fields.CharField(255, null=True, description="原因")
|
||||
"""原因"""
|
||||
shadow_effect = fields.CharField(32, null=True, description="影子决策结果")
|
||||
"""影子决策结果"""
|
||||
shadow_reason = fields.CharField(255, null=True, description="影子决策原因")
|
||||
"""影子决策原因"""
|
||||
side_effect_state = fields.TextField(null=True, description="副作用状态")
|
||||
"""副作用状态 JSON 摘要"""
|
||||
latency_ms = fields.FloatField(default=0, description="耗时毫秒")
|
||||
"""耗时毫秒"""
|
||||
overloaded = fields.BooleanField(default=False, description="是否过载")
|
||||
"""是否过载"""
|
||||
create_time = fields.DatetimeField(auto_now_add=True, description="创建时间")
|
||||
"""创建时间"""
|
||||
|
||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||
table = "auth_decision_log"
|
||||
table_description = "权限决策追加审计日志"
|
||||
indexes: ClassVar = [
|
||||
("create_time",),
|
||||
("module", "create_time"),
|
||||
("effect", "create_time"),
|
||||
]
|
||||
|
||||
@classmethod
|
||||
async def _run_script(cls):
|
||||
return []
|
||||
@@ -3,6 +3,7 @@ from typing import ClassVar
|
||||
from typing_extensions import Self
|
||||
|
||||
from tortoise import fields
|
||||
from tortoise.expressions import Q
|
||||
|
||||
from zhenxun.services.cache import CacheException, CacheRegistry, CacheRoot
|
||||
from zhenxun.services.cache.runtime_cache import BanMemoryCache
|
||||
@@ -89,15 +90,18 @@ class BanConsole(Model):
|
||||
cls._ensure_cache_registered()
|
||||
if not user_id and not group_id:
|
||||
raise UserAndGroupIsNone()
|
||||
dao = DataAccess(cls)
|
||||
if user_id:
|
||||
dao = DataAccess(cls)
|
||||
return (
|
||||
await dao.safe_get_or_none(user_id=user_id, group_id=group_id)
|
||||
if group_id
|
||||
else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
|
||||
)
|
||||
else:
|
||||
return await dao.safe_get_or_none(user_id="", group_id=group_id)
|
||||
return await cls.safe_get_or_none(
|
||||
Q(user_id__isnull=True) | Q(user_id=""),
|
||||
group_id=group_id,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def check_ban_level(
|
||||
|
||||
@@ -26,7 +26,7 @@ class BotConsole(Model):
|
||||
available_plugins = fields.TextField(default="", description="可用插件")
|
||||
"""可用插件"""
|
||||
available_tasks = fields.TextField(default="", description="可用被动技能")
|
||||
"""可用被动技能"""
|
||||
"""可用被动技能管理镜像,不作为运行白名单。"""
|
||||
|
||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||
table = "bot_console"
|
||||
@@ -87,7 +87,11 @@ class BotConsole(Model):
|
||||
@classmethod
|
||||
async def get_tasks(cls, bot_id: str | None = None, status: bool | None = True):
|
||||
"""
|
||||
获取bot被动技能
|
||||
获取bot被动技能管理镜像。
|
||||
|
||||
available_tasks 只服务管理命令和展示,不参与运行白名单判断。
|
||||
被动运行真源是 TaskInfo.status/load_status、BotConsole.block_tasks、
|
||||
GroupConsole.block_task/superuser_block_task。
|
||||
|
||||
参数:
|
||||
bot_id (str | None, optional): bot_id. Defaults to None.
|
||||
|
||||
@@ -1,55 +0,0 @@
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.utils.enum import BotSentType
|
||||
|
||||
from ._bot_message_buffer import append_bot_message_store_record
|
||||
|
||||
|
||||
class BotMessageStore(Model):
|
||||
id = fields.IntField(pk=True, generated=True, auto_increment=True)
|
||||
"""自增id"""
|
||||
bot_id = fields.CharField(255, null=True)
|
||||
"""bot id"""
|
||||
user_id = fields.CharField(255, null=True)
|
||||
"""目标id"""
|
||||
group_id = fields.CharField(255, null=True)
|
||||
"""群组id"""
|
||||
sent_type = fields.CharEnumField(BotSentType)
|
||||
"""类型"""
|
||||
text = fields.TextField(null=True)
|
||||
"""文本内容"""
|
||||
plain_text = fields.TextField(null=True)
|
||||
"""纯文本"""
|
||||
platform = fields.CharField(255, null=True)
|
||||
"""平台"""
|
||||
create_time = fields.DatetimeField(auto_now_add=True)
|
||||
"""创建时间"""
|
||||
|
||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||
table = "bot_message_store"
|
||||
table_description = "Bot发送消息列表"
|
||||
|
||||
@classmethod
|
||||
async def append_buffered(
|
||||
cls,
|
||||
*,
|
||||
bot_id: str | None = None,
|
||||
user_id: str | None = None,
|
||||
group_id: str | None = None,
|
||||
sent_type: BotSentType,
|
||||
text: str | None = None,
|
||||
plain_text: str | None = None,
|
||||
platform: str | None = None,
|
||||
) -> None:
|
||||
await append_bot_message_store_record(
|
||||
cls(
|
||||
bot_id=bot_id,
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
sent_type=sent_type,
|
||||
text=text,
|
||||
plain_text=plain_text,
|
||||
platform=platform,
|
||||
)
|
||||
)
|
||||
@@ -151,8 +151,10 @@ class FgRequest(Model):
|
||||
"添加好友自动发送BOT自我介绍图片", session=req.user_id
|
||||
)
|
||||
else:
|
||||
await GroupConsole.update_or_create(
|
||||
group_id=req.group_id, defaults={"group_flag": 1}
|
||||
await GroupConsole.get_or_create_root_group(
|
||||
group_id=req.group_id,
|
||||
defaults={"group_flag": 1},
|
||||
update_defaults=True,
|
||||
)
|
||||
if req.flag == "0":
|
||||
# 用户手动申请入群,创建群认证后提醒用户拉群
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, cast, overload
|
||||
from typing_extensions import Self
|
||||
|
||||
@@ -103,6 +104,8 @@ class GroupConsole(Model):
|
||||
"""缓存键字段"""
|
||||
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
|
||||
"""开启锁"""
|
||||
_root_group_locks: ClassVar[dict[str, asyncio.Lock]] = {}
|
||||
"""普通群记录应用层锁,规避 channel_id=NULL 唯一键语义差异。"""
|
||||
|
||||
@classmethod
|
||||
async def _get_task_modules(cls, *, default_status: bool) -> list[str]:
|
||||
@@ -250,6 +253,143 @@ class GroupConsole(Model):
|
||||
|
||||
return group, is_create
|
||||
|
||||
@classmethod
|
||||
def _clean_root_group_defaults(cls, defaults: dict | None) -> dict[str, Any]:
|
||||
cleaned = {}
|
||||
for field, value in (defaults or {}).items():
|
||||
if field in {"id", "group_id", "channel_id", "channel_id__isnull"}:
|
||||
continue
|
||||
if value is None:
|
||||
continue
|
||||
cleaned[field] = value
|
||||
return cleaned
|
||||
|
||||
@classmethod
|
||||
def _root_group_score(cls, group: Self) -> tuple[int, int]:
|
||||
score = 0
|
||||
score += 8 if group.group_name else 0
|
||||
score += 4 if group.max_member_count else 0
|
||||
score += 4 if group.member_count else 0
|
||||
score += 4 if group.group_flag else 0
|
||||
score += 4 if group.is_super else 0
|
||||
score += 3 if not group.status else 0
|
||||
score += 3 if group.level != 5 else 0
|
||||
score += len(convert_module_format(group.block_plugin))
|
||||
score += len(convert_module_format(group.superuser_block_plugin))
|
||||
score += len(convert_module_format(group.block_task))
|
||||
score += len(convert_module_format(group.superuser_block_task))
|
||||
return score, int(group.id or 0)
|
||||
|
||||
@classmethod
|
||||
def _merge_module_field(cls, groups: list[Self], field: str) -> str:
|
||||
modules: list[str] = []
|
||||
seen = set()
|
||||
for group in groups:
|
||||
value = getattr(group, field, "") or ""
|
||||
for module in cast(list[str], convert_module_format(value)):
|
||||
if module not in seen:
|
||||
seen.add(module)
|
||||
modules.append(module)
|
||||
return cast(str, convert_module_format(modules))
|
||||
|
||||
@classmethod
|
||||
async def _deduplicate_root_group_records(cls, groups: list[Self]) -> Self:
|
||||
if len(groups) == 1:
|
||||
return groups[0]
|
||||
|
||||
keep = max(groups, key=cls._root_group_score)
|
||||
newest_first = sorted(
|
||||
groups, key=lambda group: int(group.id or 0), reverse=True
|
||||
)
|
||||
|
||||
merged = {
|
||||
"group_name": next(
|
||||
(g.group_name for g in newest_first if g.group_name), ""
|
||||
),
|
||||
"max_member_count": max(g.max_member_count for g in groups),
|
||||
"member_count": max(g.member_count for g in groups),
|
||||
"status": all(g.status for g in groups),
|
||||
"level": min(g.level for g in groups),
|
||||
"is_super": any(g.is_super for g in groups),
|
||||
"group_flag": max(g.group_flag for g in groups),
|
||||
"block_plugin": cls._merge_module_field(groups, "block_plugin"),
|
||||
"superuser_block_plugin": cls._merge_module_field(
|
||||
groups, "superuser_block_plugin"
|
||||
),
|
||||
"block_task": cls._merge_module_field(groups, "block_task"),
|
||||
"superuser_block_task": cls._merge_module_field(
|
||||
groups, "superuser_block_task"
|
||||
),
|
||||
"platform": next((g.platform for g in newest_first if g.platform), "qq"),
|
||||
}
|
||||
|
||||
update_fields = []
|
||||
for field, value in merged.items():
|
||||
if getattr(keep, field) != value:
|
||||
setattr(keep, field, value)
|
||||
update_fields.append(field)
|
||||
if update_fields:
|
||||
await keep.save(update_fields=update_fields)
|
||||
|
||||
for group in groups:
|
||||
if group.id != keep.id:
|
||||
await group.delete()
|
||||
await GroupMemoryCache.upsert_from_model(keep)
|
||||
return keep
|
||||
|
||||
@classmethod
|
||||
async def get_or_create_root_group(
|
||||
cls,
|
||||
group_id: str | int,
|
||||
defaults: dict | None = None,
|
||||
*,
|
||||
update_defaults: bool = False,
|
||||
) -> tuple[Self, bool]:
|
||||
"""获取或创建普通群记录,并收敛 channel_id=NULL 重复数据。
|
||||
|
||||
普通群固定使用 ``channel_id IS NULL``;频道记录必须继续显式传
|
||||
``channel_id`` 走原有 get_or_create/update_or_create。
|
||||
"""
|
||||
gid = str(group_id).strip()
|
||||
if not gid:
|
||||
raise ValueError("group_id cannot be empty")
|
||||
|
||||
lock = cls._root_group_locks.setdefault(gid, asyncio.Lock())
|
||||
async with lock:
|
||||
defaults = cls._clean_root_group_defaults(defaults)
|
||||
records = await cls.filter(group_id=gid, channel_id__isnull=True).all()
|
||||
if records:
|
||||
group = await cls._deduplicate_root_group_records(records)
|
||||
if update_defaults:
|
||||
update_fields = []
|
||||
for field, value in defaults.items():
|
||||
if hasattr(group, field) and getattr(group, field) != value:
|
||||
setattr(group, field, value)
|
||||
update_fields.append(field)
|
||||
if update_fields:
|
||||
await group.save(update_fields=update_fields)
|
||||
await GroupMemoryCache.upsert_from_model(group)
|
||||
return group, False
|
||||
|
||||
group = await cls.create(group_id=gid, channel_id=None, **defaults)
|
||||
return group, True
|
||||
|
||||
@classmethod
|
||||
async def _get_or_create_group_for_write(
|
||||
cls,
|
||||
group_id: str,
|
||||
channel_id: str | None,
|
||||
defaults: dict | None = None,
|
||||
) -> tuple[Self, bool]:
|
||||
defaults = cls._clean_root_group_defaults(defaults)
|
||||
if channel_id:
|
||||
return await cls.get_or_create(
|
||||
group_id=group_id,
|
||||
channel_id=channel_id,
|
||||
defaults=defaults,
|
||||
)
|
||||
return await cls.get_or_create_root_group(group_id, defaults=defaults)
|
||||
|
||||
async def save(self, *args, **kwargs):
|
||||
await super().save(*args, **kwargs)
|
||||
await GroupMemoryCache.upsert_from_model(self)
|
||||
@@ -326,6 +466,7 @@ class GroupConsole(Model):
|
||||
module: str,
|
||||
is_superuser: bool = False,
|
||||
platform: str | None = None,
|
||||
channel_id: str | None = None,
|
||||
):
|
||||
"""禁用群组插件
|
||||
|
||||
@@ -335,8 +476,10 @@ class GroupConsole(Model):
|
||||
is_superuser: 是否为超级用户
|
||||
platform: 平台
|
||||
"""
|
||||
group, _ = await cls.get_or_create(
|
||||
group_id=group_id, defaults={"platform": platform}
|
||||
group, _ = await cls._get_or_create_group_for_write(
|
||||
group_id=group_id,
|
||||
channel_id=channel_id,
|
||||
defaults={"platform": platform},
|
||||
)
|
||||
update_fields = []
|
||||
if is_superuser:
|
||||
@@ -365,6 +508,7 @@ class GroupConsole(Model):
|
||||
module: str,
|
||||
is_superuser: bool = False,
|
||||
platform: str | None = None,
|
||||
channel_id: str | None = None,
|
||||
):
|
||||
"""禁用群组插件
|
||||
|
||||
@@ -374,8 +518,10 @@ class GroupConsole(Model):
|
||||
is_superuser: 是否为超级用户
|
||||
platform: 平台
|
||||
"""
|
||||
group, _ = await cls.get_or_create(
|
||||
group_id=group_id, defaults={"platform": platform}
|
||||
group, _ = await cls._get_or_create_group_for_write(
|
||||
group_id=group_id,
|
||||
channel_id=channel_id,
|
||||
defaults={"platform": platform},
|
||||
)
|
||||
update_fields = []
|
||||
if is_superuser:
|
||||
@@ -447,6 +593,7 @@ class GroupConsole(Model):
|
||||
task: str,
|
||||
is_superuser: bool = False,
|
||||
platform: str | None = None,
|
||||
channel_id: str | None = None,
|
||||
):
|
||||
"""禁用群组插件
|
||||
|
||||
@@ -456,13 +603,15 @@ class GroupConsole(Model):
|
||||
is_superuser: 是否为超级用户
|
||||
platform: 平台
|
||||
"""
|
||||
group, _ = await cls.get_or_create(
|
||||
group_id=group_id, defaults={"platform": platform}
|
||||
group, _ = await cls._get_or_create_group_for_write(
|
||||
group_id=group_id,
|
||||
channel_id=channel_id,
|
||||
defaults={"platform": platform},
|
||||
)
|
||||
update_fields = []
|
||||
if is_superuser:
|
||||
superuser_block_task = convert_module_format(group.superuser_block_task)
|
||||
if task not in group.superuser_block_task:
|
||||
if task not in superuser_block_task:
|
||||
superuser_block_task.append(task)
|
||||
group.superuser_block_task = convert_module_format(superuser_block_task)
|
||||
update_fields.append("superuser_block_task")
|
||||
@@ -484,6 +633,7 @@ class GroupConsole(Model):
|
||||
task: str,
|
||||
is_superuser: bool = False,
|
||||
platform: str | None = None,
|
||||
channel_id: str | None = None,
|
||||
):
|
||||
"""禁用群组插件
|
||||
|
||||
@@ -493,8 +643,10 @@ class GroupConsole(Model):
|
||||
is_superuser: 是否为超级用户
|
||||
platform: 平台
|
||||
"""
|
||||
group, _ = await cls.get_or_create(
|
||||
group_id=group_id, defaults={"platform": platform}
|
||||
group, _ = await cls._get_or_create_group_for_write(
|
||||
group_id=group_id,
|
||||
channel_id=channel_id,
|
||||
defaults={"platform": platform},
|
||||
)
|
||||
update_fields = []
|
||||
if is_superuser:
|
||||
|
||||
@@ -1,39 +0,0 @@
|
||||
from typing import ClassVar
|
||||
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.db_context import Model
|
||||
|
||||
|
||||
class RuntimeBackpressureLog(Model):
|
||||
id = fields.IntField(pk=True, generated=True, auto_increment=True)
|
||||
"""自增id"""
|
||||
scope_key = fields.CharField(255, null=True, description="作用域")
|
||||
"""作用域"""
|
||||
reason = fields.CharField(255, null=True, description="原因")
|
||||
"""原因"""
|
||||
lane = fields.CharField(64, null=True, description="调度通道")
|
||||
"""调度通道"""
|
||||
action = fields.CharField(64, description="处理动作")
|
||||
"""处理动作 execute/skip/defer/signal"""
|
||||
queue_size = fields.IntField(default=0, description="队列长度")
|
||||
"""队列长度"""
|
||||
active_count = fields.IntField(default=0, description="活跃数量")
|
||||
"""活跃数量"""
|
||||
duration_ms = fields.FloatField(default=0, description="持续耗时毫秒")
|
||||
"""持续耗时毫秒"""
|
||||
create_time = fields.DatetimeField(auto_now_add=True, description="创建时间")
|
||||
"""创建时间"""
|
||||
|
||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||
table = "runtime_backpressure_log"
|
||||
table_description = "运行时背压追加审计日志"
|
||||
indexes: ClassVar = [
|
||||
("create_time",),
|
||||
("scope_key", "create_time"),
|
||||
("lane", "create_time"),
|
||||
]
|
||||
|
||||
@classmethod
|
||||
async def _run_script(cls):
|
||||
return []
|
||||
@@ -219,6 +219,7 @@ class UserConsole(Model):
|
||||
)
|
||||
if not updated:
|
||||
raise InsufficientGold()
|
||||
await cls.invalidate_user_cache(user_id)
|
||||
return GoldReservation(
|
||||
user_id=user_id,
|
||||
gold=gold,
|
||||
|
||||
@@ -1,547 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections import deque
|
||||
import contextlib
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta
|
||||
import json
|
||||
import random
|
||||
import time
|
||||
from typing import Any, TypeVar
|
||||
|
||||
from tortoise import Tortoise
|
||||
|
||||
from zhenxun.builtin_plugins.hooks.auth_runtime_config import (
|
||||
AUTH_OBSERVABILITY_RUNTIME_CONFIG,
|
||||
)
|
||||
from zhenxun.models.auth_decision_log import AuthDecisionLog
|
||||
from zhenxun.models.runtime_backpressure_log import RuntimeBackpressureLog
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
|
||||
LOG_COMMAND = "AuthObservability"
|
||||
|
||||
_BUFFER_MAX_RETAIN = AUTH_OBSERVABILITY_RUNTIME_CONFIG.buffer_max_retain
|
||||
_FLUSH_TRIGGER_SIZE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_trigger_size
|
||||
_FLUSH_BATCH_SIZE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_batch_size
|
||||
_FLUSH_INTERVAL_SECONDS = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_interval_seconds
|
||||
_DROP_LOG_INTERVAL_SECONDS = AUTH_OBSERVABILITY_RUNTIME_CONFIG.drop_log_interval_seconds
|
||||
_ALLOW_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.allow_sample_rate
|
||||
_OVERLOADED_ALLOW_SAMPLE_RATE = (
|
||||
AUTH_OBSERVABILITY_RUNTIME_CONFIG.overloaded_allow_sample_rate
|
||||
)
|
||||
_NON_ALLOW_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.non_allow_sample_rate
|
||||
_BACKPRESSURE_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.backpressure_sample_rate
|
||||
_BACKPRESSURE_SEVERE_ACTIVE_THRESHOLD = (
|
||||
AUTH_OBSERVABILITY_RUNTIME_CONFIG.backpressure_severe_active_threshold
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class AuthDecisionLogRecord:
|
||||
bot_id: str | None
|
||||
platform: str | None
|
||||
group_id: str | None
|
||||
user_id: str | None
|
||||
module: str | None
|
||||
effect: str
|
||||
reason: str | None = None
|
||||
shadow_effect: str | None = None
|
||||
shadow_reason: str | None = None
|
||||
side_effect_state: dict[str, Any] | None = None
|
||||
latency_ms: float = 0.0
|
||||
overloaded: bool = False
|
||||
|
||||
def to_model(self) -> AuthDecisionLog:
|
||||
return AuthDecisionLog(
|
||||
bot_id=self.bot_id,
|
||||
platform=self.platform,
|
||||
group_id=self.group_id,
|
||||
user_id=self.user_id,
|
||||
module=self.module,
|
||||
effect=self.effect,
|
||||
reason=self.reason,
|
||||
shadow_effect=self.shadow_effect,
|
||||
shadow_reason=self.shadow_reason,
|
||||
side_effect_state=json.dumps(
|
||||
self.side_effect_state,
|
||||
ensure_ascii=False,
|
||||
separators=(",", ":"),
|
||||
)[:4000]
|
||||
if self.side_effect_state
|
||||
else None,
|
||||
latency_ms=self.latency_ms,
|
||||
overloaded=self.overloaded,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RuntimeBackpressureLogRecord:
|
||||
scope_key: str | None
|
||||
reason: str | None
|
||||
lane: str | None
|
||||
action: str
|
||||
queue_size: int = 0
|
||||
active_count: int = 0
|
||||
duration_ms: float = 0.0
|
||||
|
||||
def to_model(self) -> RuntimeBackpressureLog:
|
||||
return RuntimeBackpressureLog(
|
||||
scope_key=self.scope_key,
|
||||
reason=self.reason,
|
||||
lane=self.lane,
|
||||
action=self.action,
|
||||
queue_size=self.queue_size,
|
||||
active_count=self.active_count,
|
||||
duration_ms=self.duration_ms,
|
||||
)
|
||||
|
||||
|
||||
_auth_decision_buffer: deque[AuthDecisionLogRecord] = deque()
|
||||
_backpressure_buffer: deque[RuntimeBackpressureLogRecord] = deque()
|
||||
_buffer_lock = asyncio.Lock()
|
||||
_flush_lock = asyncio.Lock()
|
||||
_flush_task: asyncio.Task[None] | None = None
|
||||
_dropped = 0
|
||||
_last_drop_log_at = 0.0
|
||||
_last_schema_repair_at = 0.0
|
||||
_SCHEMA_REPAIR_INTERVAL_SECONDS = 300.0
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def _ensure_flush_task() -> None:
|
||||
global _flush_task
|
||||
if _flush_task is not None and not _flush_task.done():
|
||||
return
|
||||
_flush_task = asyncio.create_task(_flush_loop())
|
||||
|
||||
|
||||
def _record_drop() -> None:
|
||||
global _dropped, _last_drop_log_at
|
||||
_dropped += 1
|
||||
now = time.monotonic()
|
||||
if now - _last_drop_log_at < _DROP_LOG_INTERVAL_SECONDS:
|
||||
return
|
||||
_last_drop_log_at = now
|
||||
logger.warning(
|
||||
"auth observability buffer full, dropped "
|
||||
f"{_dropped} records, auth_backlog={len(_auth_decision_buffer)}, "
|
||||
f"backpressure_backlog={len(_backpressure_buffer)}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
def _sample(rate: float) -> bool:
|
||||
if rate >= 1:
|
||||
return True
|
||||
if rate <= 0:
|
||||
return False
|
||||
return random.random() < rate
|
||||
|
||||
|
||||
def _auth_decision_sample_rate(effect: str, overloaded: bool) -> float:
|
||||
if effect != "allow":
|
||||
return _NON_ALLOW_SAMPLE_RATE
|
||||
if overloaded:
|
||||
return _OVERLOADED_ALLOW_SAMPLE_RATE
|
||||
return _ALLOW_SAMPLE_RATE
|
||||
|
||||
|
||||
def _backpressure_sample_rate(record: RuntimeBackpressureLogRecord) -> float:
|
||||
if record.reason and record.reason.startswith("hooks_"):
|
||||
return 1.0
|
||||
if record.active_count >= _BACKPRESSURE_SEVERE_ACTIVE_THRESHOLD:
|
||||
return 1.0
|
||||
if record.action in {"skip", "defer"}:
|
||||
return _BACKPRESSURE_SAMPLE_RATE
|
||||
return min(_BACKPRESSURE_SAMPLE_RATE, 0.02)
|
||||
|
||||
|
||||
async def _append_auth_decision_record(record: AuthDecisionLogRecord) -> None:
|
||||
_ensure_flush_task()
|
||||
async with _buffer_lock:
|
||||
total = len(_auth_decision_buffer) + len(_backpressure_buffer)
|
||||
if total >= _BUFFER_MAX_RETAIN:
|
||||
if len(_auth_decision_buffer) >= len(_backpressure_buffer):
|
||||
with contextlib.suppress(IndexError):
|
||||
_auth_decision_buffer.popleft()
|
||||
else:
|
||||
with contextlib.suppress(IndexError):
|
||||
_backpressure_buffer.popleft()
|
||||
_record_drop()
|
||||
_auth_decision_buffer.append(record)
|
||||
should_flush = (
|
||||
len(_auth_decision_buffer) + len(_backpressure_buffer)
|
||||
>= _FLUSH_TRIGGER_SIZE
|
||||
and not _flush_lock.locked()
|
||||
)
|
||||
if should_flush:
|
||||
# Fire-and-forget keeps auth hot path independent of database stalls.
|
||||
asyncio.create_task(flush_auth_observability_buffer("缓冲区触发")) # noqa: RUF006
|
||||
|
||||
|
||||
async def _append_backpressure_record(record: RuntimeBackpressureLogRecord) -> None:
|
||||
_ensure_flush_task()
|
||||
async with _buffer_lock:
|
||||
total = len(_auth_decision_buffer) + len(_backpressure_buffer)
|
||||
if total >= _BUFFER_MAX_RETAIN:
|
||||
if len(_auth_decision_buffer) >= len(_backpressure_buffer):
|
||||
with contextlib.suppress(IndexError):
|
||||
_auth_decision_buffer.popleft()
|
||||
else:
|
||||
with contextlib.suppress(IndexError):
|
||||
_backpressure_buffer.popleft()
|
||||
_record_drop()
|
||||
_backpressure_buffer.append(record)
|
||||
should_flush = (
|
||||
len(_auth_decision_buffer) + len(_backpressure_buffer)
|
||||
>= _FLUSH_TRIGGER_SIZE
|
||||
and not _flush_lock.locked()
|
||||
)
|
||||
if should_flush:
|
||||
# Fire-and-forget keeps auth hot path independent of database stalls.
|
||||
asyncio.create_task(flush_auth_observability_buffer("缓冲区触发")) # noqa: RUF006
|
||||
|
||||
|
||||
async def append_auth_decision_log(
|
||||
*,
|
||||
bot_id: str | None,
|
||||
platform: str | None,
|
||||
group_id: str | None,
|
||||
user_id: str | None,
|
||||
module: str | None,
|
||||
effect: str,
|
||||
reason: str | None = None,
|
||||
shadow_effect: str | None = None,
|
||||
shadow_reason: str | None = None,
|
||||
side_effect_state: dict[str, Any] | None = None,
|
||||
latency_ms: float = 0.0,
|
||||
overloaded: bool = False,
|
||||
) -> None:
|
||||
if shadow_effect is None and not _sample(
|
||||
_auth_decision_sample_rate(effect, overloaded)
|
||||
):
|
||||
return
|
||||
record = AuthDecisionLogRecord(
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
group_id=group_id,
|
||||
user_id=user_id,
|
||||
module=module,
|
||||
effect=effect,
|
||||
reason=(reason or "")[:255] or None,
|
||||
shadow_effect=(shadow_effect or "")[:32] or None,
|
||||
shadow_reason=(shadow_reason or "")[:255] or None,
|
||||
side_effect_state=side_effect_state,
|
||||
latency_ms=latency_ms,
|
||||
overloaded=overloaded,
|
||||
)
|
||||
await _append_auth_decision_record(record)
|
||||
|
||||
|
||||
async def append_runtime_backpressure_log(
|
||||
*,
|
||||
scope_key: str | None,
|
||||
reason: str | None,
|
||||
lane: str | None,
|
||||
action: str,
|
||||
queue_size: int = 0,
|
||||
active_count: int = 0,
|
||||
duration_ms: float = 0.0,
|
||||
) -> None:
|
||||
record = RuntimeBackpressureLogRecord(
|
||||
scope_key=(scope_key or "")[:255] or None,
|
||||
reason=(reason or "")[:255] or None,
|
||||
lane=(lane or "")[:64] or None,
|
||||
action=action,
|
||||
queue_size=queue_size,
|
||||
active_count=active_count,
|
||||
duration_ms=duration_ms,
|
||||
)
|
||||
if not _sample(_backpressure_sample_rate(record)):
|
||||
return
|
||||
await _append_backpressure_record(record)
|
||||
|
||||
|
||||
async def _flush_loop() -> None:
|
||||
while True:
|
||||
await asyncio.sleep(_FLUSH_INTERVAL_SECONDS)
|
||||
try:
|
||||
await flush_auth_observability_buffer("定时")
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.warning("定时批量写入权限观测日志失败", LOG_COMMAND, e=exc)
|
||||
|
||||
|
||||
async def _drain_batch(buffer: deque[T]) -> list[T]:
|
||||
batch: list[T] = []
|
||||
async with _buffer_lock:
|
||||
while buffer and len(batch) < _FLUSH_BATCH_SIZE:
|
||||
batch.append(buffer.popleft())
|
||||
return batch
|
||||
|
||||
|
||||
async def _restore_batch(buffer: deque[T], batch: list[T]) -> None:
|
||||
async with _buffer_lock:
|
||||
retain_count = max(_BUFFER_MAX_RETAIN - len(buffer), 0)
|
||||
for record in reversed(batch[-retain_count:]):
|
||||
buffer.appendleft(record)
|
||||
|
||||
|
||||
def _is_schema_mismatch_error(exc: Exception) -> bool:
|
||||
message = str(exc).lower()
|
||||
return any(
|
||||
marker in message
|
||||
for marker in (
|
||||
"no column named",
|
||||
"unknown column",
|
||||
"column does not exist",
|
||||
"no such column",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def _try_repair_auth_schema_once() -> bool:
|
||||
global _last_schema_repair_at
|
||||
now = time.monotonic()
|
||||
if now - _last_schema_repair_at < _SCHEMA_REPAIR_INTERVAL_SECONDS:
|
||||
return False
|
||||
_last_schema_repair_at = now
|
||||
try:
|
||||
from zhenxun.services.db_context.schema_guard import repair_table_schema
|
||||
|
||||
await repair_table_schema("auth_decision_log")
|
||||
await repair_table_schema("runtime_backpressure_log")
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.warning("权限观测日志表结构自修复失败", LOG_COMMAND, e=exc)
|
||||
return False
|
||||
|
||||
|
||||
async def flush_auth_observability_buffer(reason: str) -> int:
|
||||
async with _flush_lock:
|
||||
written = 0
|
||||
while True:
|
||||
auth_batch = await _drain_batch(_auth_decision_buffer)
|
||||
backpressure_batch = await _drain_batch(_backpressure_buffer)
|
||||
if not auth_batch and not backpressure_batch:
|
||||
break
|
||||
try:
|
||||
if auth_batch:
|
||||
await AuthDecisionLog.bulk_create(
|
||||
[record.to_model() for record in auth_batch],
|
||||
_FLUSH_BATCH_SIZE,
|
||||
)
|
||||
written += len(auth_batch)
|
||||
if backpressure_batch:
|
||||
await RuntimeBackpressureLog.bulk_create(
|
||||
[record.to_model() for record in backpressure_batch],
|
||||
_FLUSH_BATCH_SIZE,
|
||||
)
|
||||
written += len(backpressure_batch)
|
||||
except Exception as exc:
|
||||
if _is_schema_mismatch_error(exc):
|
||||
if await _try_repair_auth_schema_once():
|
||||
try:
|
||||
if auth_batch:
|
||||
await AuthDecisionLog.bulk_create(
|
||||
[record.to_model() for record in auth_batch],
|
||||
_FLUSH_BATCH_SIZE,
|
||||
)
|
||||
written += len(auth_batch)
|
||||
if backpressure_batch:
|
||||
await RuntimeBackpressureLog.bulk_create(
|
||||
[
|
||||
record.to_model()
|
||||
for record in backpressure_batch
|
||||
],
|
||||
_FLUSH_BATCH_SIZE,
|
||||
)
|
||||
written += len(backpressure_batch)
|
||||
continue
|
||||
except Exception as retry_exc:
|
||||
exc = retry_exc
|
||||
dropped = len(auth_batch) + len(backpressure_batch)
|
||||
logger.warning(
|
||||
f"{reason}批量写入权限观测日志遇到表结构不匹配,"
|
||||
f"已丢弃低优先级观测日志 {dropped} 条,等待下次启动修复",
|
||||
LOG_COMMAND,
|
||||
e=exc,
|
||||
)
|
||||
return written
|
||||
await _restore_batch(_auth_decision_buffer, auth_batch)
|
||||
await _restore_batch(_backpressure_buffer, backpressure_batch)
|
||||
logger.error(f"{reason}批量写入权限观测日志失败", LOG_COMMAND, e=exc)
|
||||
return written
|
||||
if written:
|
||||
logger.debug(f"{reason}批量写入权限观测日志 {written} 条", LOG_COMMAND)
|
||||
return written
|
||||
|
||||
|
||||
async def stop_auth_observability_buffer() -> int:
|
||||
global _flush_task
|
||||
task = _flush_task
|
||||
_flush_task = None
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
with contextlib.suppress(BaseException):
|
||||
await task
|
||||
return await flush_auth_observability_buffer("关闭")
|
||||
|
||||
|
||||
def _percentile(values: list[float], ratio: float) -> float:
|
||||
if not values:
|
||||
return 0.0
|
||||
ordered = sorted(values)
|
||||
index = min(max(round((len(ordered) - 1) * ratio), 0), len(ordered) - 1)
|
||||
return round(ordered[index], 3)
|
||||
|
||||
|
||||
def _bucket_counts(rows: list[dict[str, Any]], field: str) -> dict[str, int]:
|
||||
counts: dict[str, int] = {}
|
||||
for row in rows:
|
||||
key = str(row.get(field) or "<none>")
|
||||
counts[key] = counts.get(key, 0) + 1
|
||||
return counts
|
||||
|
||||
|
||||
def _lane_budget_advice(rows: list[dict[str, Any]]) -> dict[str, dict[str, Any]]:
|
||||
buckets: dict[str, list[dict[str, Any]]] = {}
|
||||
for row in rows:
|
||||
lane = str(row.get("lane") or "<unknown>")
|
||||
buckets.setdefault(lane, []).append(row)
|
||||
advice: dict[str, dict[str, Any]] = {}
|
||||
for lane, items in buckets.items():
|
||||
if lane == "<unknown>":
|
||||
continue
|
||||
durations = [float(item.get("duration_ms") or 0.0) for item in items]
|
||||
slow_waits = sum(1 for value in durations if value >= 200.0)
|
||||
active_max = max(
|
||||
(int(item.get("active_count") or 0) for item in items), default=0
|
||||
)
|
||||
total = len(items)
|
||||
if not total:
|
||||
continue
|
||||
pressure_ratio = slow_waits / total
|
||||
if pressure_ratio >= 0.2 or active_max >= 5:
|
||||
action = "increase_or_split"
|
||||
elif pressure_ratio == 0 and active_max <= 1 and total >= 20:
|
||||
action = "can_reduce"
|
||||
else:
|
||||
action = "keep"
|
||||
advice[lane] = {
|
||||
"samples": total,
|
||||
"slow_waits": slow_waits,
|
||||
"pressure_ratio": round(pressure_ratio, 3),
|
||||
"active_max": active_max,
|
||||
"p95_duration_ms": _percentile(durations, 0.95),
|
||||
"action": action,
|
||||
}
|
||||
return advice
|
||||
|
||||
|
||||
def _query_placeholder() -> str:
|
||||
try:
|
||||
connection = Tortoise.get_connection("default")
|
||||
if (
|
||||
getattr(connection, "capabilities", None)
|
||||
and getattr(
|
||||
connection.capabilities,
|
||||
"dialect",
|
||||
"",
|
||||
)
|
||||
== "postgres"
|
||||
):
|
||||
return "$1"
|
||||
except Exception:
|
||||
return "?"
|
||||
return "?"
|
||||
|
||||
|
||||
async def build_auth_observability_report(*, hours: float = 24.0) -> dict[str, Any]:
|
||||
since = datetime.now() - timedelta(hours=hours)
|
||||
db = Tortoise.get_connection("default")
|
||||
placeholder = _query_placeholder()
|
||||
auth_rows = await db.execute_query_dict(
|
||||
"SELECT module, effect, reason, shadow_effect, shadow_reason, latency_ms, "
|
||||
f"overloaded FROM auth_decision_log WHERE create_time >= {placeholder} "
|
||||
"ORDER BY create_time DESC LIMIT 100000",
|
||||
[since],
|
||||
)
|
||||
backpressure_rows = await db.execute_query_dict(
|
||||
"SELECT scope_key, lane, reason, action, queue_size, active_count, duration_ms "
|
||||
f"FROM runtime_backpressure_log WHERE create_time >= {placeholder} "
|
||||
"ORDER BY create_time DESC LIMIT 100000",
|
||||
[since],
|
||||
)
|
||||
|
||||
module_buckets: dict[str, list[dict[str, Any]]] = {}
|
||||
for row in auth_rows:
|
||||
module_buckets.setdefault(str(row.get("module") or "<unknown>"), []).append(row)
|
||||
module_stats: list[dict[str, Any]] = []
|
||||
for module, items in module_buckets.items():
|
||||
latencies = [float(item.get("latency_ms") or 0.0) for item in items]
|
||||
module_stats.append(
|
||||
{
|
||||
"module": module,
|
||||
"total": len(items),
|
||||
"effects": _bucket_counts(items, "effect"),
|
||||
"shadow_effects": _bucket_counts(items, "shadow_effect"),
|
||||
"avg_latency_ms": round(sum(latencies) / len(latencies), 3)
|
||||
if latencies
|
||||
else 0.0,
|
||||
"p95_latency_ms": _percentile(latencies, 0.95),
|
||||
"overloaded": sum(1 for item in items if bool(item.get("overloaded"))),
|
||||
}
|
||||
)
|
||||
|
||||
backpressure_buckets: dict[str, list[dict[str, Any]]] = {}
|
||||
for row in backpressure_rows:
|
||||
key = f"{row.get('lane') or '<unknown>'}:{row.get('reason') or '<none>'}"
|
||||
backpressure_buckets.setdefault(key, []).append(row)
|
||||
backpressure_stats: list[dict[str, Any]] = []
|
||||
for key, items in backpressure_buckets.items():
|
||||
durations = [float(item.get("duration_ms") or 0.0) for item in items]
|
||||
backpressure_stats.append(
|
||||
{
|
||||
"key": key,
|
||||
"total": len(items),
|
||||
"actions": _bucket_counts(items, "action"),
|
||||
"avg_duration_ms": round(sum(durations) / len(durations), 3)
|
||||
if durations
|
||||
else 0.0,
|
||||
"p95_duration_ms": _percentile(durations, 0.95),
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"created_at": datetime.now().isoformat(timespec="seconds"),
|
||||
"window_hours": hours,
|
||||
"auth_decisions": {
|
||||
"total": len(auth_rows),
|
||||
"effects": _bucket_counts(auth_rows, "effect"),
|
||||
"shadow_effects": _bucket_counts(auth_rows, "shadow_effect"),
|
||||
"top_modules_by_p95": sorted(
|
||||
module_stats,
|
||||
key=lambda item: (item["p95_latency_ms"], item["total"]),
|
||||
reverse=True,
|
||||
)[:30],
|
||||
},
|
||||
"backpressure": {
|
||||
"total": len(backpressure_rows),
|
||||
"lane_budget_advice": _lane_budget_advice(backpressure_rows),
|
||||
"top_reasons": sorted(
|
||||
backpressure_stats,
|
||||
key=lambda item: (item["total"], item["p95_duration_ms"]),
|
||||
reverse=True,
|
||||
)[:30],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@PriorityLifecycle.on_shutdown(priority=90)
|
||||
async def _flush_auth_observability_buffer_on_shutdown() -> None:
|
||||
await stop_auth_observability_buffer()
|
||||
+378
-72
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass, field
|
||||
import json
|
||||
import os
|
||||
@@ -35,6 +36,8 @@ def _coerce_int(value, default: int) -> int:
|
||||
|
||||
# RuntimeCache 以模型 save/delete 主动失效为主,周期 refresh 只做兜底。
|
||||
# 这些默认值避免低压力运行时频繁全量扫表。
|
||||
# 权限检查热路径以 RuntimeCache/AuthSnapshot 为唯一数据入口;普通业务的
|
||||
# DataAccess/CacheRoot 缓存不能替代这里的运行态快照。
|
||||
PLUGININFO_MEM_REFRESH_INTERVAL = 3600 # 60分钟
|
||||
BAN_MEM_REFRESH_INTERVAL = 300
|
||||
BAN_MEM_CLEAN_INTERVAL = 60
|
||||
@@ -52,10 +55,16 @@ LIMIT_MEM_REFRESH_INTERVAL = 900 # 15分钟
|
||||
LIMIT_MEM_NEGATIVE_TTL = 30
|
||||
RUNTIME_CACHE_SYNC_ENABLED = True
|
||||
RUNTIME_CACHE_SYNC_CHANNEL = "ZHENXUN_RUNTIME_CACHE_SYNC"
|
||||
RUNTIME_CACHE_LOAD_RETRY_SECONDS = 1.0
|
||||
RUNTIME_CACHE_STARTUP_REFRESH_SKIP_SECONDS = 5.0
|
||||
|
||||
|
||||
INSTANCE_ID = uuid.uuid4().hex
|
||||
_CACHE_READY_EVENT = asyncio.Event()
|
||||
_APPLYING_REMOTE_CACHE_EVENT: ContextVar[bool] = ContextVar(
|
||||
"APPLYING_REMOTE_RUNTIME_CACHE_EVENT",
|
||||
default=False,
|
||||
)
|
||||
|
||||
|
||||
def _env_get(name: str, default: str | None = None) -> str | None:
|
||||
@@ -180,6 +189,65 @@ class PluginInfoSnapshot:
|
||||
plugin._saved_in_db = True
|
||||
return plugin
|
||||
|
||||
def to_payload(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"module": self.module,
|
||||
"module_path": self.module_path,
|
||||
"name": self.name,
|
||||
"status": self.status,
|
||||
"block_type": self.block_type.value if self.block_type else None,
|
||||
"load_status": self.load_status,
|
||||
"author": self.author,
|
||||
"version": self.version,
|
||||
"level": self.level,
|
||||
"default_status": self.default_status,
|
||||
"limit_superuser": self.limit_superuser,
|
||||
"menu_type": self.menu_type,
|
||||
"plugin_type": self.plugin_type.value if self.plugin_type else None,
|
||||
"cost_gold": self.cost_gold,
|
||||
"admin_level": self.admin_level,
|
||||
"ignore_prompt": self.ignore_prompt,
|
||||
"is_delete": self.is_delete,
|
||||
"parent": self.parent,
|
||||
"is_show": self.is_show,
|
||||
"ignore_statistics": self.ignore_statistics,
|
||||
"impression": self.impression,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_payload(cls, payload: dict[str, Any]) -> "PluginInfoSnapshot":
|
||||
block_type = payload.get("block_type")
|
||||
if block_type is not None and not isinstance(block_type, BlockType):
|
||||
block_type = BlockType(block_type)
|
||||
plugin_type = payload.get("plugin_type")
|
||||
if plugin_type is not None and not isinstance(plugin_type, PluginType):
|
||||
plugin_type = PluginType(plugin_type)
|
||||
return cls(
|
||||
id=int(payload.get("id", 0) or 0),
|
||||
module=str(payload.get("module", "") or ""),
|
||||
module_path=str(payload.get("module_path", "") or ""),
|
||||
name=str(payload.get("name", "") or ""),
|
||||
status=bool(payload.get("status", True)),
|
||||
block_type=block_type,
|
||||
load_status=bool(payload.get("load_status", True)),
|
||||
author=payload.get("author"),
|
||||
version=payload.get("version"),
|
||||
level=int(payload.get("level", 0) or 0),
|
||||
default_status=bool(payload.get("default_status", True)),
|
||||
limit_superuser=bool(payload.get("limit_superuser", False)),
|
||||
menu_type=str(payload.get("menu_type", "") or ""),
|
||||
plugin_type=plugin_type,
|
||||
cost_gold=int(payload.get("cost_gold", 0) or 0),
|
||||
admin_level=payload.get("admin_level"),
|
||||
ignore_prompt=bool(payload.get("ignore_prompt", False)),
|
||||
is_delete=bool(payload.get("is_delete", False)),
|
||||
parent=payload.get("parent"),
|
||||
is_show=bool(payload.get("is_show", True)),
|
||||
ignore_statistics=bool(payload.get("ignore_statistics", False)),
|
||||
impression=float(payload.get("impression", 0) or 0),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BanEntry:
|
||||
@@ -648,18 +716,87 @@ class RuntimeCacheSync:
|
||||
cache_type = payload.get("type")
|
||||
action = payload.get("action")
|
||||
data = payload.get("data") or {}
|
||||
if cache_type == "bot":
|
||||
await BotMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "group":
|
||||
await GroupMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "ban":
|
||||
await BanMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "level":
|
||||
await LevelUserMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "task":
|
||||
await TaskInfoMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "plugin_limit":
|
||||
await PluginLimitMemoryCache.apply_sync_event(action, data)
|
||||
token = _APPLYING_REMOTE_CACHE_EVENT.set(True)
|
||||
try:
|
||||
if cache_type == "bot":
|
||||
await BotMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "group":
|
||||
await GroupMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "ban":
|
||||
await BanMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "level":
|
||||
await LevelUserMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "task":
|
||||
await TaskInfoMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "plugin_limit":
|
||||
await PluginLimitMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "plugin":
|
||||
await PluginInfoMemoryCache.apply_sync_event(action, data)
|
||||
finally:
|
||||
_APPLYING_REMOTE_CACHE_EVENT.reset(token)
|
||||
|
||||
|
||||
class RuntimeCacheMutation:
|
||||
"""Small helpers for runtime cache mutation bookkeeping.
|
||||
|
||||
Cache classes still own their storage layout. This helper centralizes the
|
||||
shared mutation side effects: health markers, negative-cache cleanup and
|
||||
cross-process publish.
|
||||
"""
|
||||
|
||||
_load_locks: ClassVar[dict[str, asyncio.Lock]] = {}
|
||||
_retry_after: ClassVar[dict[str, float]] = {}
|
||||
|
||||
@classmethod
|
||||
async def ensure_loaded(cls, cache_cls: type, label: str) -> None:
|
||||
if getattr(cache_cls, "_loaded", False):
|
||||
return
|
||||
now = time.monotonic()
|
||||
if cls._retry_after.get(label, 0.0) > now:
|
||||
return
|
||||
lock = cls._load_locks.setdefault(label, asyncio.Lock())
|
||||
async with lock:
|
||||
if getattr(cache_cls, "_loaded", False):
|
||||
return
|
||||
now = time.monotonic()
|
||||
if cls._retry_after.get(label, 0.0) > now:
|
||||
return
|
||||
try:
|
||||
await cache_cls.refresh()
|
||||
except Exception as exc:
|
||||
cls.mark_error(cache_cls, exc)
|
||||
cls._retry_after[label] = (
|
||||
time.monotonic() + RUNTIME_CACHE_LOAD_RETRY_SECONDS
|
||||
)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def mark_refreshed(cache_cls: type) -> None:
|
||||
setattr(cache_cls, "_loaded", True)
|
||||
setattr(cache_cls, "_last_refresh", time.time())
|
||||
setattr(cache_cls, "_last_error", None)
|
||||
|
||||
@staticmethod
|
||||
def mark_error(cache_cls: type, exc: Exception) -> None:
|
||||
setattr(cache_cls, "_last_error", f"{type(exc).__name__}: {exc}")
|
||||
|
||||
@staticmethod
|
||||
def clear_negative_key(cache_cls: type, key: object) -> None:
|
||||
negative = getattr(cache_cls, "_negative", None)
|
||||
if isinstance(negative, dict):
|
||||
negative.pop(key, None)
|
||||
|
||||
@staticmethod
|
||||
def clear_negative_all(cache_cls: type) -> None:
|
||||
negative = getattr(cache_cls, "_negative", None)
|
||||
if isinstance(negative, dict):
|
||||
negative.clear()
|
||||
|
||||
@staticmethod
|
||||
def publish(cache_type: str, action: str, data: dict[str, Any]) -> None:
|
||||
if _APPLYING_REMOTE_CACHE_EVENT.get():
|
||||
return
|
||||
RuntimeCacheSync.publish_event(cache_type, action, data)
|
||||
|
||||
|
||||
class PluginInfoMemoryCache:
|
||||
@@ -669,6 +806,7 @@ class PluginInfoMemoryCache:
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _to_model(cls, snapshot: PluginInfoSnapshot | None) -> "PluginInfo | None":
|
||||
@@ -687,6 +825,10 @@ class PluginInfoMemoryCache:
|
||||
cls._by_module.pop(old.module, None)
|
||||
cls._by_module_path[snapshot.module_path] = snapshot
|
||||
|
||||
@staticmethod
|
||||
def _module_snapshot_rank(snapshot: PluginInfoSnapshot) -> tuple[int, int]:
|
||||
return (1 if snapshot.load_status else 0, snapshot.id)
|
||||
|
||||
@classmethod
|
||||
async def refresh(cls) -> None:
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
@@ -698,13 +840,16 @@ class PluginInfoMemoryCache:
|
||||
for plugin in plugins:
|
||||
snapshot = PluginInfoSnapshot.from_model(plugin)
|
||||
if snapshot.module:
|
||||
by_module[snapshot.module] = snapshot
|
||||
current = by_module.get(snapshot.module)
|
||||
if current is None or cls._module_snapshot_rank(
|
||||
snapshot
|
||||
) >= cls._module_snapshot_rank(current):
|
||||
by_module[snapshot.module] = snapshot
|
||||
if snapshot.module_path:
|
||||
by_module_path[snapshot.module_path] = snapshot
|
||||
cls._by_module = by_module
|
||||
cls._by_module_path = by_module_path
|
||||
cls._loaded = True
|
||||
cls._last_refresh = time.time()
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(
|
||||
f"plugin cache refreshed: {len(by_module)} entries", LOG_COMMAND
|
||||
)
|
||||
@@ -713,7 +858,7 @@ class PluginInfoMemoryCache:
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "plugin")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -749,8 +894,6 @@ class PluginInfoMemoryCache:
|
||||
return
|
||||
snapshot = PluginInfoSnapshot.from_model(plugin)
|
||||
cls._store_snapshot(snapshot)
|
||||
cls._loaded = True
|
||||
cls._last_refresh = time.time()
|
||||
|
||||
@classmethod
|
||||
def remove_by_module(cls, module: str) -> None:
|
||||
@@ -765,8 +908,16 @@ class PluginInfoMemoryCache:
|
||||
async with cls._lock:
|
||||
snapshot = PluginInfoSnapshot.from_model(plugin)
|
||||
cls._store_snapshot(snapshot)
|
||||
cls._loaded = True
|
||||
cls._last_refresh = time.time()
|
||||
RuntimeCacheMutation.publish("plugin", "upsert", snapshot.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
try:
|
||||
snapshot = PluginInfoSnapshot.from_payload(payload)
|
||||
except Exception:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._store_snapshot(snapshot)
|
||||
|
||||
@classmethod
|
||||
async def remove(
|
||||
@@ -783,6 +934,20 @@ class PluginInfoMemoryCache:
|
||||
snapshot = cls._by_module_path.pop(module_path, None)
|
||||
if snapshot and snapshot.module:
|
||||
cls._by_module.pop(snapshot.module, None)
|
||||
RuntimeCacheMutation.publish(
|
||||
"plugin",
|
||||
"delete",
|
||||
{"module": module, "module_path": module_path},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
if action == "upsert":
|
||||
await cls.upsert_from_payload(data)
|
||||
elif action == "delete":
|
||||
await cls.remove(data.get("module"), data.get("module_path"))
|
||||
elif action == "refresh":
|
||||
await cls.refresh()
|
||||
|
||||
@classmethod
|
||||
async def _refresh_loop(cls, interval: int) -> None:
|
||||
@@ -815,6 +980,8 @@ class BotMemoryCache:
|
||||
_negative: ClassVar[dict[str, float]] = {}
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize(cls, bot_id: str | None) -> str | None:
|
||||
@@ -833,7 +1000,7 @@ class BotMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(bot_id, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, bot_id)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -851,15 +1018,15 @@ class BotMemoryCache:
|
||||
async with cls._lock:
|
||||
records = await BotConsole.all()
|
||||
cls._by_id = {str(r.bot_id): BotSnapshot.from_model(r) for r in records}
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(f"bot cache refreshed: {len(cls._by_id)} entries", LOG_COMMAND)
|
||||
|
||||
@classmethod
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "bot")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -918,15 +1085,16 @@ class BotMemoryCache:
|
||||
available_tasks=entry.available_tasks,
|
||||
)
|
||||
cls._by_id[bot_id] = updated
|
||||
RuntimeCacheSync.publish_event("bot", "upsert", updated.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, bot_id)
|
||||
RuntimeCacheMutation.publish("bot", "upsert", updated.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_model(cls, record) -> None:
|
||||
entry = BotSnapshot.from_model(record)
|
||||
async with cls._lock:
|
||||
cls._by_id[entry.bot_id] = entry
|
||||
cls._negative.pop(entry.bot_id, None)
|
||||
RuntimeCacheSync.publish_event("bot", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.bot_id)
|
||||
RuntimeCacheMutation.publish("bot", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -935,7 +1103,7 @@ class BotMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_id[entry.bot_id] = entry
|
||||
cls._negative.pop(entry.bot_id, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.bot_id)
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, bot_id: str | None) -> None:
|
||||
@@ -944,7 +1112,8 @@ class BotMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_id.pop(bot_id, None)
|
||||
RuntimeCacheSync.publish_event("bot", "delete", {"bot_id": bot_id})
|
||||
RuntimeCacheMutation.clear_negative_key(cls, bot_id)
|
||||
RuntimeCacheMutation.publish("bot", "delete", {"bot_id": bot_id})
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
@@ -986,6 +1155,8 @@ class GroupMemoryCache:
|
||||
_negative: ClassVar[dict[tuple[str, str], float]] = {}
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize(cls, value: str | None) -> str | None:
|
||||
@@ -1014,7 +1185,7 @@ class GroupMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1038,15 +1209,15 @@ class GroupMemoryCache:
|
||||
if key:
|
||||
by_key[key] = entry
|
||||
cls._by_key = by_key
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(f"group cache refreshed: {len(by_key)} entries", LOG_COMMAND)
|
||||
|
||||
@classmethod
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "group")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -1094,8 +1265,8 @@ class GroupMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_key[key] = entry
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheSync.publish_event("group", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
RuntimeCacheMutation.publish("group", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -1105,7 +1276,7 @@ class GroupMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_key[key] = entry
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, group_id: str | None, channel_id: str | None = None) -> None:
|
||||
@@ -1114,7 +1285,8 @@ class GroupMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_key.pop(key, None)
|
||||
RuntimeCacheSync.publish_event(
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
RuntimeCacheMutation.publish(
|
||||
"group", "delete", {"group_id": key[0], "channel_id": key[1] or None}
|
||||
)
|
||||
|
||||
@@ -1186,7 +1358,7 @@ class LevelUserMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1215,16 +1387,15 @@ class LevelUserMemoryCache:
|
||||
by_user_max[entry.user_id] = entry.user_level
|
||||
cls._by_key = by_key
|
||||
cls._by_user_max = by_user_max
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
cls._last_refresh = time.time()
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(f"level cache refreshed: {len(by_key)} entries", LOG_COMMAND)
|
||||
|
||||
@classmethod
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "level")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -1310,13 +1481,13 @@ class LevelUserMemoryCache:
|
||||
async with cls._lock:
|
||||
prev = cls._by_key.get(key)
|
||||
cls._by_key[key] = entry
|
||||
cls._negative.pop(key, None)
|
||||
current = cls._by_user_max.get(entry.user_id, 0)
|
||||
if entry.user_level >= current:
|
||||
cls._by_user_max[entry.user_id] = entry.user_level
|
||||
elif prev and prev.user_level == current and entry.user_level < current:
|
||||
cls._recalc_user_max(entry.user_id)
|
||||
RuntimeCacheSync.publish_event("level", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
RuntimeCacheMutation.publish("level", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -1327,12 +1498,12 @@ class LevelUserMemoryCache:
|
||||
async with cls._lock:
|
||||
prev = cls._by_key.get(key)
|
||||
cls._by_key[key] = entry
|
||||
cls._negative.pop(key, None)
|
||||
current = cls._by_user_max.get(entry.user_id, 0)
|
||||
if entry.user_level >= current:
|
||||
cls._by_user_max[entry.user_id] = entry.user_level
|
||||
elif prev and prev.user_level == current and entry.user_level < current:
|
||||
cls._recalc_user_max(entry.user_id)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, user_id: str | None, group_id: str | None) -> None:
|
||||
@@ -1343,7 +1514,8 @@ class LevelUserMemoryCache:
|
||||
removed = cls._by_key.pop(key, None)
|
||||
if removed and cls._by_user_max.get(removed.user_id) == removed.user_level:
|
||||
cls._recalc_user_max(removed.user_id)
|
||||
RuntimeCacheSync.publish_event(
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
RuntimeCacheMutation.publish(
|
||||
"level", "delete", {"user_id": key[0], "group_id": key[1] or None}
|
||||
)
|
||||
|
||||
@@ -1399,6 +1571,8 @@ class TaskInfoMemoryCache:
|
||||
_negative: ClassVar[dict[str, float]] = {}
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize(cls, module: str | None) -> str | None:
|
||||
@@ -1417,7 +1591,7 @@ class TaskInfoMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(module, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, module)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1443,8 +1617,8 @@ class TaskInfoMemoryCache:
|
||||
by_name[entry.name] = entry
|
||||
cls._by_module = by_module
|
||||
cls._by_name = by_name
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(
|
||||
f"task info cache refreshed: {len(cls._by_module)} entries",
|
||||
LOG_COMMAND,
|
||||
@@ -1454,7 +1628,7 @@ class TaskInfoMemoryCache:
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "task")
|
||||
|
||||
@classmethod
|
||||
async def get(cls, module: str | None) -> TaskInfoSnapshot | None:
|
||||
@@ -1488,10 +1662,21 @@ class TaskInfoMemoryCache:
|
||||
|
||||
@classmethod
|
||||
async def is_disabled(cls, module: str | None) -> bool:
|
||||
"""Backward-compatible runtime disabled check for passive tasks."""
|
||||
return await cls.is_runtime_disabled(module)
|
||||
|
||||
@classmethod
|
||||
async def is_runtime_disabled(cls, module: str | None) -> bool:
|
||||
"""Return whether a passive task is unavailable at runtime.
|
||||
|
||||
Runtime passive availability is defined by TaskInfo.status and
|
||||
TaskInfo.load_status. Bot/group scoped block lists are checked by
|
||||
CommonUtils.task_is_block().
|
||||
"""
|
||||
entry = await cls.get(module)
|
||||
if not entry:
|
||||
return False
|
||||
return not entry.status
|
||||
return not entry.status or not entry.load_status
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_model(cls, record) -> None:
|
||||
@@ -1500,8 +1685,8 @@ class TaskInfoMemoryCache:
|
||||
cls._by_module[entry.module] = entry
|
||||
if entry.name:
|
||||
cls._by_name[entry.name] = entry
|
||||
cls._negative.pop(entry.module, None)
|
||||
RuntimeCacheSync.publish_event("task", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
RuntimeCacheMutation.publish("task", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -1512,7 +1697,7 @@ class TaskInfoMemoryCache:
|
||||
cls._by_module[entry.module] = entry
|
||||
if entry.name:
|
||||
cls._by_name[entry.name] = entry
|
||||
cls._negative.pop(entry.module, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, module: str | None) -> None:
|
||||
@@ -1525,7 +1710,8 @@ class TaskInfoMemoryCache:
|
||||
current = cls._by_name.get(removed.name)
|
||||
if current and current.module == removed.module:
|
||||
cls._by_name.pop(removed.name, None)
|
||||
RuntimeCacheSync.publish_event("task", "delete", {"module": module})
|
||||
RuntimeCacheMutation.clear_negative_key(cls, module)
|
||||
RuntimeCacheMutation.publish("task", "delete", {"module": module})
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
@@ -1568,6 +1754,8 @@ class PluginLimitMemoryCache:
|
||||
_negative: ClassVar[dict[str, float]] = {}
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize(cls, value: str | None) -> str | None:
|
||||
@@ -1586,7 +1774,7 @@ class PluginLimitMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(module, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, module)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1611,8 +1799,8 @@ class PluginLimitMemoryCache:
|
||||
by_module.setdefault(entry.module, []).append(entry)
|
||||
cls._by_id = by_id
|
||||
cls._by_module = by_module
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(
|
||||
f"plugin limit cache refreshed: {len(by_id)} entries",
|
||||
LOG_COMMAND,
|
||||
@@ -1622,7 +1810,7 @@ class PluginLimitMemoryCache:
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "plugin_limit")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -1668,7 +1856,7 @@ class PluginLimitMemoryCache:
|
||||
async def upsert_from_model(cls, record) -> None:
|
||||
entry = PluginLimitSnapshot.from_model(record)
|
||||
await cls._upsert_entry(entry)
|
||||
RuntimeCacheSync.publish_event("plugin_limit", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.publish("plugin_limit", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -1695,6 +1883,7 @@ class PluginLimitMemoryCache:
|
||||
for item in cls._by_module.get(entry.module, [])
|
||||
if item.id != entry.id
|
||||
]
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
return
|
||||
cls._by_id[entry.id] = entry
|
||||
module_limits = [
|
||||
@@ -1704,7 +1893,7 @@ class PluginLimitMemoryCache:
|
||||
]
|
||||
module_limits.append(entry)
|
||||
cls._by_module[entry.module] = module_limits
|
||||
cls._negative.pop(entry.module, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
|
||||
@classmethod
|
||||
async def remove_by_id(cls, limit_id: int | None) -> None:
|
||||
@@ -1718,7 +1907,8 @@ class PluginLimitMemoryCache:
|
||||
for item in cls._by_module.get(entry.module, [])
|
||||
if item.id != entry.id
|
||||
]
|
||||
RuntimeCacheSync.publish_event("plugin_limit", "delete", {"id": int(limit_id)})
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
RuntimeCacheMutation.publish("plugin_limit", "delete", {"id": int(limit_id)})
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
@@ -1764,6 +1954,8 @@ class BanMemoryCache:
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_cleanup_task: ClassVar[asyncio.Task | None] = None
|
||||
_remove_tasks: ClassVar[set[asyncio.Task]] = set()
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize_id(cls, value: str | None) -> str | None:
|
||||
@@ -1788,7 +1980,7 @@ class BanMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1842,8 +2034,8 @@ class BanMemoryCache:
|
||||
cls._by_user = by_user
|
||||
cls._by_group = by_group
|
||||
cls._by_user_group = by_user_group
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(
|
||||
"ban cache refreshed: "
|
||||
f"user={len(by_user)} group={len(by_group)} "
|
||||
@@ -1855,7 +2047,7 @@ class BanMemoryCache:
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "ban")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -1873,13 +2065,13 @@ class BanMemoryCache:
|
||||
cls._by_user[entry.user_id] = entry
|
||||
elif entry.group_id:
|
||||
cls._by_group[entry.group_id] = entry
|
||||
cls._negative = {}
|
||||
RuntimeCacheSync.publish_event("ban", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.publish("ban", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, user_id: str | None, group_id: str | None) -> None:
|
||||
await cls._remove_local(user_id, group_id)
|
||||
RuntimeCacheSync.publish_event(
|
||||
RuntimeCacheMutation.publish(
|
||||
"ban", "delete", {"user_id": user_id, "group_id": group_id}
|
||||
)
|
||||
|
||||
@@ -1894,7 +2086,7 @@ class BanMemoryCache:
|
||||
cls._by_user.pop(user_id, None)
|
||||
elif group_id:
|
||||
cls._by_group.pop(group_id, None)
|
||||
cls._negative = {}
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
|
||||
@classmethod
|
||||
def _get_entry(cls, user_id: str | None, group_id: str | None) -> BanEntry | None:
|
||||
@@ -1995,7 +2187,7 @@ class BanMemoryCache:
|
||||
elif entry.group_id:
|
||||
cls._by_group.pop(entry.group_id, None)
|
||||
if expired:
|
||||
cls._negative = {}
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
if not delete_db or not expired:
|
||||
return
|
||||
from tortoise.expressions import Q
|
||||
@@ -2064,7 +2256,7 @@ class BanMemoryCache:
|
||||
cls._by_user[entry.user_id] = entry
|
||||
elif entry.group_id:
|
||||
cls._by_group[entry.group_id] = entry
|
||||
cls._negative = {}
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
@@ -2078,12 +2270,126 @@ class BanMemoryCache:
|
||||
|
||||
async def _safe_refresh(cache_cls: type, label: str) -> None:
|
||||
"""安全地刷新单个缓存,异常不影响其他缓存。"""
|
||||
if getattr(cache_cls, "_loaded", False):
|
||||
last_refresh = float(getattr(cache_cls, "_last_refresh", 0.0) or 0.0)
|
||||
if time.time() - last_refresh <= RUNTIME_CACHE_STARTUP_REFRESH_SKIP_SECONDS:
|
||||
logger.debug(f"{label} cache startup refresh skipped", LOG_COMMAND)
|
||||
return
|
||||
try:
|
||||
await cache_cls.refresh()
|
||||
except Exception as exc:
|
||||
RuntimeCacheMutation.mark_error(cache_cls, exc)
|
||||
logger.error(f"{label} cache init failed", LOG_COMMAND, e=exc)
|
||||
|
||||
|
||||
def _cache_health(
|
||||
cache_cls: type,
|
||||
*,
|
||||
entry_count: int,
|
||||
negative_count: int = 0,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"loaded": bool(getattr(cache_cls, "_loaded", False)),
|
||||
"entry_count": entry_count,
|
||||
"last_refresh": float(getattr(cache_cls, "_last_refresh", 0.0) or 0.0),
|
||||
"negative_count": negative_count,
|
||||
"last_error": getattr(cache_cls, "_last_error", None),
|
||||
}
|
||||
|
||||
|
||||
def health_snapshot() -> dict[str, dict[str, Any]]:
|
||||
"""Return in-memory runtime cache health without touching the database."""
|
||||
return {
|
||||
"plugin": _cache_health(
|
||||
PluginInfoMemoryCache,
|
||||
entry_count=len(PluginInfoMemoryCache._by_module),
|
||||
),
|
||||
"bot": _cache_health(
|
||||
BotMemoryCache,
|
||||
entry_count=len(BotMemoryCache._by_id),
|
||||
negative_count=len(BotMemoryCache._negative),
|
||||
),
|
||||
"group": _cache_health(
|
||||
GroupMemoryCache,
|
||||
entry_count=len(GroupMemoryCache._by_key),
|
||||
negative_count=len(GroupMemoryCache._negative),
|
||||
),
|
||||
"level": _cache_health(
|
||||
LevelUserMemoryCache,
|
||||
entry_count=len(LevelUserMemoryCache._by_key),
|
||||
negative_count=len(LevelUserMemoryCache._negative),
|
||||
),
|
||||
"task": _cache_health(
|
||||
TaskInfoMemoryCache,
|
||||
entry_count=len(TaskInfoMemoryCache._by_module),
|
||||
negative_count=len(TaskInfoMemoryCache._negative),
|
||||
),
|
||||
"plugin_limit": _cache_health(
|
||||
PluginLimitMemoryCache,
|
||||
entry_count=len(PluginLimitMemoryCache._by_id),
|
||||
negative_count=len(PluginLimitMemoryCache._negative),
|
||||
),
|
||||
"ban": _cache_health(
|
||||
BanMemoryCache,
|
||||
entry_count=(
|
||||
len(BanMemoryCache._by_user)
|
||||
+ len(BanMemoryCache._by_group)
|
||||
+ len(BanMemoryCache._by_user_group)
|
||||
),
|
||||
negative_count=len(BanMemoryCache._negative),
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def passive_status_snapshot(max_modules: int = 50) -> dict[str, Any]:
|
||||
"""Return passive-task state from in-memory caches only.
|
||||
|
||||
This is a local diagnostic helper: it does not query or write the database,
|
||||
and it is not used by runtime decisions.
|
||||
"""
|
||||
tasks = list(TaskInfoMemoryCache._by_module.values())
|
||||
disabled = sorted(task.module for task in tasks if not task.status)
|
||||
unloaded = sorted(task.module for task in tasks if not task.load_status)
|
||||
runtime_enabled = [
|
||||
task.module for task in tasks if task.status and task.load_status
|
||||
]
|
||||
bot_block_total = sum(
|
||||
len(_parse_block_modules(bot.block_tasks))
|
||||
for bot in BotMemoryCache._by_id.values()
|
||||
)
|
||||
group_block_total = sum(
|
||||
len(group.block_task_set) + len(group.superuser_block_task_set)
|
||||
for group in GroupMemoryCache._by_key.values()
|
||||
)
|
||||
return {
|
||||
"cache": health_snapshot(),
|
||||
"passive_tasks": {
|
||||
"total": len(tasks),
|
||||
"status_enabled": sum(1 for task in tasks if task.status),
|
||||
"load_status_enabled": sum(1 for task in tasks if task.load_status),
|
||||
"runtime_enabled": len(runtime_enabled),
|
||||
"disabled_modules": disabled[:max_modules],
|
||||
"disabled_modules_total": len(disabled),
|
||||
"unloaded_modules": unloaded[:max_modules],
|
||||
"unloaded_modules_total": len(unloaded),
|
||||
},
|
||||
"scoped_blocks": {
|
||||
"bot_block_tasks_total": bot_block_total,
|
||||
"group_block_tasks_total": group_block_total,
|
||||
},
|
||||
"semantics": {
|
||||
"available_tasks": "management_display_mirror_not_runtime_whitelist",
|
||||
"runtime_truth": [
|
||||
"TaskInfo.status",
|
||||
"TaskInfo.load_status",
|
||||
"BotConsole.block_tasks",
|
||||
"GroupConsole.block_task",
|
||||
"GroupConsole.superuser_block_task",
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=6)
|
||||
async def _init_runtime_cache():
|
||||
await RuntimeCacheSync.start()
|
||||
|
||||
@@ -12,6 +12,10 @@ T = TypeVar("T", bound=Model)
|
||||
class DataAccess(Generic[T]):
|
||||
"""数据访问兼容层,根据配置保留单点缓存读取和清理能力
|
||||
|
||||
边界说明:DataAccess 面向普通业务查询和低频管理链路。权限检查热路径
|
||||
必须优先使用 RuntimeCache/AuthSnapshot,避免高并发消息处理时触发
|
||||
DB/cache-aside 读放大。
|
||||
|
||||
新的高频运行态路径应优先使用 RuntimeCache 或 BoundedTTLCache。
|
||||
这里不再把 filter/all/create/update_or_create 结果写入通用缓存,
|
||||
create/update_or_create 只负责清理旧缓存,避免旧值残留。
|
||||
|
||||
@@ -5,6 +5,9 @@ from pydantic import BaseModel
|
||||
# 数据库操作超时设置(秒)
|
||||
DB_TIMEOUT_SECONDS = 3.0
|
||||
|
||||
# 启动期自动补齐字段/索引可能需要等待数据库锁或扫描较大的表,单独放宽超时
|
||||
DB_SCHEMA_GUARD_TIMEOUT_SECONDS = 30.0
|
||||
|
||||
# 性能监控阈值(秒)
|
||||
SLOW_QUERY_THRESHOLD = 0.5
|
||||
|
||||
|
||||
@@ -13,7 +13,7 @@ from tortoise.exceptions import OperationalError
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import DB_TIMEOUT_SECONDS, LOG_COMMAND
|
||||
from .config import DB_SCHEMA_GUARD_TIMEOUT_SECONDS, LOG_COMMAND
|
||||
|
||||
Dialect = Literal["sqlite", "postgres", "mysql", "unknown"]
|
||||
|
||||
@@ -466,7 +466,7 @@ async def repair_safe_schema_drift() -> SchemaGuardResult:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
connection.execute_query_dict(sql),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
timeout=DB_SCHEMA_GUARD_TIMEOUT_SECONDS,
|
||||
)
|
||||
columns[source] = ColumnInfo(name=source, data_type="")
|
||||
result.repaired_columns += 1
|
||||
@@ -521,7 +521,7 @@ async def repair_safe_schema_drift() -> SchemaGuardResult:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
connection.execute_query_dict(sql),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
timeout=DB_SCHEMA_GUARD_TIMEOUT_SECONDS,
|
||||
)
|
||||
existing_indexes.add(index_columns)
|
||||
result.repaired_indexes += 1
|
||||
@@ -529,6 +529,16 @@ async def repair_safe_schema_drift() -> SchemaGuardResult:
|
||||
f"SchemaGuard 已补齐索引: {table}.{index_columns}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
except TimeoutError as exc:
|
||||
result.warnings += 1
|
||||
result.skipped_indexes += 1
|
||||
logger.warning(
|
||||
"SchemaGuard 补齐索引超时,已跳过: "
|
||||
f"{table}.{index_columns} "
|
||||
f"({DB_SCHEMA_GUARD_TIMEOUT_SECONDS}s)",
|
||||
LOG_COMMAND,
|
||||
e=exc,
|
||||
)
|
||||
except OperationalError as exc:
|
||||
err = str(exc).lower()
|
||||
if any(
|
||||
|
||||
@@ -140,9 +140,6 @@ def register_runtime_bootstrap(_driver) -> None:
|
||||
global _thread_executor
|
||||
await _stop_launcher_watchdog()
|
||||
await stop_send_queue()
|
||||
from zhenxun.models._bot_message_buffer import stop_bot_message_store_buffer
|
||||
|
||||
await stop_bot_message_store_buffer()
|
||||
await stop_memory_governor()
|
||||
executor = _thread_executor
|
||||
_thread_executor = None
|
||||
|
||||
@@ -27,6 +27,7 @@ from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import should_pause_tasks
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.decorator.retry import Retry
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
from zhenxun.utils.pydantic_compat import parse_as
|
||||
|
||||
from .repository import ScheduleRepository
|
||||
@@ -37,6 +38,37 @@ SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit"
|
||||
_LAST_PRESSURE_SKIP = 0.0
|
||||
|
||||
|
||||
def _resolve_scheduler_bot(bot_id: str | None, log_target: str) -> Bot | None:
|
||||
if bot_id:
|
||||
try:
|
||||
return nonebot.get_bot(bot_id)
|
||||
except KeyError:
|
||||
logger.warning(f"{log_target} 需要的 Bot {bot_id} 不在线,本次执行跳过。")
|
||||
return None
|
||||
|
||||
bots = list(nonebot.get_bots().values())
|
||||
if not bots:
|
||||
logger.warning(f"{log_target} 当前没有可用 Bot,本次执行跳过。")
|
||||
return None
|
||||
if len(bots) == 1:
|
||||
return bots[0]
|
||||
|
||||
qq_client_bots = [
|
||||
bot for bot in bots if PlatformUtils.get_platform_scope(bot) == "qq_client"
|
||||
]
|
||||
if len(qq_client_bots) == 1:
|
||||
bot = qq_client_bots[0]
|
||||
logger.warning(
|
||||
f"{log_target} 未指定 Bot,多 Bot 在线,自动选择 OneBot {bot.self_id}。"
|
||||
)
|
||||
return bot
|
||||
|
||||
logger.warning(
|
||||
f"{log_target} 未指定 Bot 且多 Bot 在线,无法安全选择," "本次执行跳过。"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
class APSchedulerAdapter:
|
||||
"""封装对 APScheduler 的操作"""
|
||||
|
||||
@@ -343,7 +375,12 @@ async def _execute_job(
|
||||
return
|
||||
|
||||
try:
|
||||
bot = nonebot.get_bot()
|
||||
bot = _resolve_scheduler_bot(
|
||||
context_override.bot_id,
|
||||
f"临时任务 {plugin_name}",
|
||||
)
|
||||
if bot is None:
|
||||
return
|
||||
logger.info(f"开始执行临时任务: {plugin_name}")
|
||||
injected_params = {"context": context_override}
|
||||
state: T_State = {ScheduleContext: context_override}
|
||||
@@ -380,18 +417,9 @@ async def _execute_job(
|
||||
logger.warning(f"定时任务 {schedule_id} 不存在或已禁用,跳过执行。")
|
||||
return
|
||||
|
||||
try:
|
||||
bot = (
|
||||
nonebot.get_bot(schedule.bot_id)
|
||||
if schedule.bot_id
|
||||
else nonebot.get_bot()
|
||||
)
|
||||
except (KeyError, ValueError):
|
||||
logger.warning(
|
||||
f"任务 {schedule_id} 需要的 Bot {schedule.bot_id} "
|
||||
f"不在线,本次执行跳过。"
|
||||
)
|
||||
raise
|
||||
bot = _resolve_scheduler_bot(schedule.bot_id, f"任务 {schedule_id}")
|
||||
if bot is None:
|
||||
return
|
||||
|
||||
resolver = scheduler_manager._target_resolvers.get(schedule.target_type)
|
||||
if not resolver:
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable
|
||||
import contextlib
|
||||
import importlib
|
||||
from typing import Any, cast
|
||||
|
||||
from nonebot.adapters import Bot, Event
|
||||
@@ -10,6 +11,9 @@ from nonebot.log import logger
|
||||
_PATCHED = False
|
||||
_ORIGINAL_FETCH: Callable[..., Awaitable[Any]] | None = None
|
||||
_ORIGINAL_ONEBOT11_GROUP_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
|
||||
_ORIGINAL_QQ_C2C_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
|
||||
_ORIGINAL_QQ_GROUP_AT_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
|
||||
_ORIGINAL_QQ_GUILD_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
|
||||
|
||||
|
||||
def _sender_value(sender: Any, key: str, default: Any = None) -> Any:
|
||||
@@ -81,6 +85,88 @@ async def _fast_onebot11_group_message(bot: Bot, event: Event) -> dict[str, Any]
|
||||
}
|
||||
|
||||
|
||||
def _qq_bot_app_id(bot: Bot) -> str:
|
||||
bot_info = getattr(bot, "bot_info", None)
|
||||
app_id = getattr(bot_info, "id", None)
|
||||
return str(app_id or getattr(bot, "self_id", ""))
|
||||
|
||||
|
||||
async def _fast_qq_c2c_message(bot: Bot, event: Event) -> dict[str, Any]:
|
||||
"""Build Uninfo session for QQ official C2C messages from event fields."""
|
||||
|
||||
author = _event_value(event, "author")
|
||||
user_id = str(
|
||||
_sender_value(author, "user_openid")
|
||||
or _sender_value(author, "id")
|
||||
or _event_value(event, "user_id", "")
|
||||
)
|
||||
username = str(_sender_value(author, "username", "") or "")
|
||||
return {
|
||||
"user_id": user_id,
|
||||
"name": username,
|
||||
"nickname": username,
|
||||
"avatar": f"https://q.qlogo.cn/qqapp/{_qq_bot_app_id(bot)}/{user_id}/100",
|
||||
}
|
||||
|
||||
|
||||
async def _fast_qq_group_at_message(bot: Bot, event: Event) -> dict[str, Any]:
|
||||
"""Build Uninfo session for QQ official group-at messages from event fields."""
|
||||
|
||||
author = _event_value(event, "author")
|
||||
user_id = str(
|
||||
_sender_value(author, "member_openid")
|
||||
or _sender_value(author, "id")
|
||||
or _event_value(event, "user_id", "")
|
||||
)
|
||||
username = str(_sender_value(author, "username", "") or "")
|
||||
group_id = str(
|
||||
_event_value(event, "group_openid") or _event_value(event, "group_id") or ""
|
||||
)
|
||||
return {
|
||||
"user_id": user_id,
|
||||
"name": username,
|
||||
"nickname": username,
|
||||
"avatar": f"https://q.qlogo.cn/qqapp/{_qq_bot_app_id(bot)}/{user_id}/100",
|
||||
"group_id": group_id,
|
||||
}
|
||||
|
||||
|
||||
async def _fast_qq_guild_message(bot: Bot, event: Event) -> dict[str, Any]:
|
||||
"""Build Uninfo session for QQ official guild/channel messages locally.
|
||||
|
||||
nonebot-plugin-uninfo enriches guild messages through remote guild/channel
|
||||
APIs. Runtime auth only needs stable scene/user ids, so avoid remote calls
|
||||
during matcher fanout.
|
||||
"""
|
||||
|
||||
author = _event_value(event, "author")
|
||||
member = _event_value(event, "member")
|
||||
guild_id = str(_event_value(event, "guild_id", "") or "")
|
||||
channel_id = str(_event_value(event, "channel_id", "") or "")
|
||||
user_id = str(_sender_value(author, "id", "") or "")
|
||||
nickname = str(_sender_value(member, "nick", "") or "")
|
||||
username = str(_sender_value(author, "username", "") or "")
|
||||
base: dict[str, Any] = {
|
||||
"user_id": user_id,
|
||||
"name": username,
|
||||
"nickname": nickname or username,
|
||||
"avatar": _sender_value(author, "avatar"),
|
||||
"guild_id": guild_id,
|
||||
"channel_id": channel_id,
|
||||
"guild_name": "",
|
||||
"guild_avatar": None,
|
||||
"channel_name": "",
|
||||
"channel_type": -1,
|
||||
}
|
||||
roles = _sender_value(member, "roles")
|
||||
if roles is not None:
|
||||
base["roles"] = roles
|
||||
joined_at = _sender_value(member, "joined_at")
|
||||
if joined_at is not None:
|
||||
base["joined_at"] = joined_at
|
||||
return base
|
||||
|
||||
|
||||
async def _singleflight_fetch(self: Any, bot: Bot, event: Event) -> Any:
|
||||
original = _ORIGINAL_FETCH
|
||||
if original is None:
|
||||
@@ -114,6 +200,8 @@ async def _singleflight_fetch(self: Any, bot: Bot, event: Event) -> Any:
|
||||
|
||||
def apply_uninfo_onebot11_patch() -> None:
|
||||
global _ORIGINAL_FETCH, _ORIGINAL_ONEBOT11_GROUP_MESSAGE, _PATCHED
|
||||
global _ORIGINAL_QQ_C2C_MESSAGE, _ORIGINAL_QQ_GROUP_AT_MESSAGE
|
||||
global _ORIGINAL_QQ_GUILD_MESSAGE
|
||||
if _PATCHED:
|
||||
return
|
||||
|
||||
@@ -129,6 +217,59 @@ def apply_uninfo_onebot11_patch() -> None:
|
||||
setattr(_fast_onebot11_group_message, "__zhenxun_fast_onebot11__", True)
|
||||
fetcher.endpoint[GroupMessageEvent] = _fast_onebot11_group_message
|
||||
|
||||
with contextlib.suppress(Exception):
|
||||
qq_event_module = importlib.import_module("nonebot.adapters.qq.event")
|
||||
AtMessageCreateEvent = getattr(qq_event_module, "AtMessageCreateEvent")
|
||||
C2CMessageCreateEvent = getattr(qq_event_module, "C2CMessageCreateEvent")
|
||||
DirectMessageCreateEvent = getattr(qq_event_module, "DirectMessageCreateEvent")
|
||||
GroupAtMessageCreateEvent = getattr(
|
||||
qq_event_module,
|
||||
"GroupAtMessageCreateEvent",
|
||||
)
|
||||
GroupMessageCreateEvent = getattr(
|
||||
qq_event_module,
|
||||
"GroupMessageCreateEvent",
|
||||
)
|
||||
MessageCreateEvent = getattr(qq_event_module, "MessageCreateEvent")
|
||||
from nonebot_plugin_uninfo.adapters.qq.main import fetcher as qq_fetcher
|
||||
|
||||
original_c2c = qq_fetcher.endpoint.get(C2CMessageCreateEvent)
|
||||
if not getattr(original_c2c, "__zhenxun_fast_qq__", False):
|
||||
_ORIGINAL_QQ_C2C_MESSAGE = cast(
|
||||
Callable[..., Awaitable[dict[str, Any]]] | None,
|
||||
original_c2c,
|
||||
)
|
||||
setattr(_fast_qq_c2c_message, "__zhenxun_fast_qq__", True)
|
||||
qq_fetcher.endpoint[C2CMessageCreateEvent] = _fast_qq_c2c_message
|
||||
|
||||
for event_type in (GroupMessageCreateEvent, GroupAtMessageCreateEvent):
|
||||
original_group_at = qq_fetcher.endpoint.get(event_type)
|
||||
if getattr(original_group_at, "__zhenxun_fast_qq__", False):
|
||||
continue
|
||||
if _ORIGINAL_QQ_GROUP_AT_MESSAGE is None and original_group_at is not None:
|
||||
_ORIGINAL_QQ_GROUP_AT_MESSAGE = cast(
|
||||
Callable[..., Awaitable[dict[str, Any]]],
|
||||
original_group_at,
|
||||
)
|
||||
setattr(_fast_qq_group_at_message, "__zhenxun_fast_qq__", True)
|
||||
qq_fetcher.endpoint[event_type] = _fast_qq_group_at_message
|
||||
|
||||
for event_type in (
|
||||
MessageCreateEvent,
|
||||
AtMessageCreateEvent,
|
||||
DirectMessageCreateEvent,
|
||||
):
|
||||
original_guild = qq_fetcher.endpoint.get(event_type)
|
||||
if getattr(original_guild, "__zhenxun_fast_qq__", False):
|
||||
continue
|
||||
if _ORIGINAL_QQ_GUILD_MESSAGE is None and original_guild is not None:
|
||||
_ORIGINAL_QQ_GUILD_MESSAGE = cast(
|
||||
Callable[..., Awaitable[dict[str, Any]]],
|
||||
original_guild,
|
||||
)
|
||||
setattr(_fast_qq_guild_message, "__zhenxun_fast_qq__", True)
|
||||
qq_fetcher.endpoint[event_type] = _fast_qq_guild_message
|
||||
|
||||
try:
|
||||
from nonebot_plugin_uninfo.fetch import InfoFetcher
|
||||
except Exception as e:
|
||||
@@ -146,4 +287,4 @@ def apply_uninfo_onebot11_patch() -> None:
|
||||
setattr(_singleflight_fetch, "__zhenxun_singleflight__", True)
|
||||
setattr(InfoFetcher, "fetch", _singleflight_fetch)
|
||||
_PATCHED = True
|
||||
logger.debug("Uninfo OneBot11 fast fetch and singleflight patch applied")
|
||||
logger.debug("Uninfo fast fetch and singleflight patch applied")
|
||||
|
||||
@@ -23,29 +23,31 @@ class CommonUtils:
|
||||
async def task_is_block(
|
||||
cls, session: Uninfo | Bot, module: str, group_id: str | None = None
|
||||
) -> bool:
|
||||
"""判断被动技能是否可以发送
|
||||
"""判断被动技能是否被阻断。
|
||||
|
||||
运行真源固定为 TaskInfo.status/load_status、BotConsole.block_tasks、
|
||||
GroupConsole.block_task/superuser_block_task,以及 bot/group ban 状态。
|
||||
BotConsole.available_tasks 只用于管理展示,不作为运行白名单。
|
||||
|
||||
参数:
|
||||
module: 被动技能模块名
|
||||
group_id: 群组id
|
||||
|
||||
返回:
|
||||
bool: 是否可以发送
|
||||
bool: True 表示被动技能应被阻断,False 表示允许继续执行
|
||||
"""
|
||||
if isinstance(session, Bot):
|
||||
if interface := get_interface(session):
|
||||
info = interface.basic_info()
|
||||
if info["scope"] == SupportScope.qq_api:
|
||||
logger.info("q官bot放弃所有被动技能发言...")
|
||||
"""q官bot放弃所有被动技能发言"""
|
||||
return False
|
||||
if session.scene == SupportScope.qq_api:
|
||||
"""q官bot放弃所有被动技能发言"""
|
||||
logger.info("q官bot放弃所有被动技能发言...")
|
||||
return False
|
||||
logger.debug("q官bot放弃所有被动技能发言...")
|
||||
return True
|
||||
if isinstance(session, Session) and session.scope == SupportScope.qq_api:
|
||||
logger.debug("q官bot放弃所有被动技能发言...")
|
||||
return True
|
||||
if not group_id and isinstance(session, Session):
|
||||
group_id = session.group.id if session.group else None
|
||||
if await TaskInfoMemoryCache.is_disabled(module):
|
||||
if await TaskInfoMemoryCache.is_runtime_disabled(module):
|
||||
"""被动全局状态"""
|
||||
return True
|
||||
bot_snapshot = await BotMemoryCache.get(session.self_id)
|
||||
|
||||
@@ -13,11 +13,6 @@ class PriorityLifecycleType(StrEnum):
|
||||
"""关闭"""
|
||||
|
||||
|
||||
class BotSentType(StrEnum):
|
||||
GROUP = "GROUP"
|
||||
PRIVATE = "PRIVATE"
|
||||
|
||||
|
||||
class BankHandleType(StrEnum):
|
||||
DEPOSIT = "DEPOSIT"
|
||||
"""存款"""
|
||||
|
||||
+70
-17
@@ -1,5 +1,6 @@
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable
|
||||
import contextlib
|
||||
import random
|
||||
from typing import cast
|
||||
|
||||
@@ -37,6 +38,23 @@ def _adapter_name(bot: Bot) -> str:
|
||||
return adapter.__class__.__name__.lower()
|
||||
|
||||
|
||||
def _scope_name(scope: object) -> str:
|
||||
"""Normalize uninfo/alconna scope values without changing legacy platform."""
|
||||
if scope is None:
|
||||
return ""
|
||||
raw = getattr(scope, "name", None) or getattr(scope, "value", scope)
|
||||
text = str(raw or "").strip().lower()
|
||||
if not text:
|
||||
text = str(getattr(scope, "name", "") or "").strip().lower()
|
||||
text = text.replace("-", "_").replace(" ", "_")
|
||||
compact = "".join(ch for ch in text if ch.isalnum())
|
||||
if compact.endswith("qqclient"):
|
||||
return "qq_client"
|
||||
if compact.endswith("qqapi"):
|
||||
return "qq_api"
|
||||
return text
|
||||
|
||||
|
||||
class UserData(BaseModel):
|
||||
name: str
|
||||
"""昵称"""
|
||||
@@ -57,6 +75,31 @@ class UserData(BaseModel):
|
||||
|
||||
|
||||
class PlatformUtils:
|
||||
@classmethod
|
||||
def _resolve_unique_qq_client_bot(cls, log_cmd: str | None = None) -> Bot | None:
|
||||
bots = list(nonebot.get_bots().values())
|
||||
if not bots:
|
||||
logger.warning("当前没有可用的 OneBot 协议端 Bot,已跳过。", log_cmd)
|
||||
return None
|
||||
qq_client_bots = [
|
||||
bot for bot in bots if cls.get_platform_scope(bot) == "qq_client"
|
||||
]
|
||||
if len(qq_client_bots) == 1:
|
||||
bot = qq_client_bots[0]
|
||||
if len(bots) > 1:
|
||||
logger.warning(
|
||||
f"多 Bot 在线且未指定 Bot,自动选择 OneBot {bot.self_id}。",
|
||||
log_cmd,
|
||||
)
|
||||
return bot
|
||||
if not qq_client_bots:
|
||||
logger.warning("未找到 OneBot 协议端 Bot,已跳过。", log_cmd)
|
||||
else:
|
||||
logger.warning(
|
||||
"存在多个 OneBot 协议端 Bot,无法安全选择,已跳过。", log_cmd
|
||||
)
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def is_qbot(cls, session: Uninfo | Bot) -> bool:
|
||||
"""判断bot是否为qq官bot
|
||||
@@ -68,7 +111,11 @@ class PlatformUtils:
|
||||
bool: 是否为官bot
|
||||
"""
|
||||
if isinstance(session, Bot):
|
||||
if cls.get_platform_scope(session) == "qq_api":
|
||||
return True
|
||||
return bool(BotConfig.get_qbot_uid(session.self_id))
|
||||
if cls.get_platform_scope(session) == "qq_api":
|
||||
return True
|
||||
if BotConfig.get_qbot_uid(session.self_id):
|
||||
return True
|
||||
return session.scope == SupportScope.qq_api
|
||||
@@ -83,7 +130,7 @@ class PlatformUtils:
|
||||
group_id: 群组id
|
||||
duration: 禁言时长(分钟)
|
||||
"""
|
||||
if cls.get_platform(bot) == "qq":
|
||||
if cls.get_platform_scope(bot) == "qq_client":
|
||||
await bot.set_group_ban(
|
||||
group_id=int(group_id),
|
||||
user_id=int(user_id),
|
||||
@@ -111,7 +158,9 @@ class PlatformUtils:
|
||||
Receipt | None: Receipt
|
||||
"""
|
||||
if not bot:
|
||||
bot = nonebot.get_bot()
|
||||
bot = cls._resolve_unique_qq_client_bot("PlatformUtils:send_superuser")
|
||||
if bot is None:
|
||||
return []
|
||||
superuser_ids = []
|
||||
if superuser_id:
|
||||
superuser_ids.append(superuser_id)
|
||||
@@ -397,6 +446,11 @@ class PlatformUtils:
|
||||
def get_platform_scope(cls, t: Bot | Uninfo | object) -> str:
|
||||
"""获取细粒度平台作用域,不改变旧 get_platform 返回值。"""
|
||||
if isinstance(t, Bot):
|
||||
if interface := get_interface(t):
|
||||
with contextlib.suppress(Exception):
|
||||
scope = _scope_name(interface.basic_info().get("scope"))
|
||||
if scope:
|
||||
return scope
|
||||
adapter_name = _adapter_name(t)
|
||||
if "onebot" in adapter_name:
|
||||
return "qq_client"
|
||||
@@ -406,17 +460,13 @@ class PlatformUtils:
|
||||
return "qq_api"
|
||||
return adapter_name or cls.get_platform(t)
|
||||
|
||||
scope = str(getattr(t, "scope", "") or "").lower()
|
||||
scope = _scope_name(getattr(t, "scope", "") or "")
|
||||
if not scope:
|
||||
basic = getattr(t, "basic", None)
|
||||
if isinstance(basic, dict):
|
||||
scope = str(basic.get("scope") or "").lower()
|
||||
if "qq_client" in scope:
|
||||
return "qq_client"
|
||||
if "qq_api" in scope:
|
||||
return "qq_api"
|
||||
if scope.startswith("qq"):
|
||||
return "qq"
|
||||
scope = _scope_name(basic.get("scope"))
|
||||
if scope:
|
||||
return scope
|
||||
|
||||
adapter = getattr(t, "adapter", None)
|
||||
if adapter is not None:
|
||||
@@ -499,6 +549,9 @@ class PlatformUtils:
|
||||
返回:
|
||||
int: 更新个数
|
||||
"""
|
||||
if cls.get_platform_scope(bot) == "qq_api":
|
||||
logger.warning("QQ 官方适配器不支持旧好友同步,已跳过。", "更新好友信息")
|
||||
return 0
|
||||
create_list = []
|
||||
friend_list, platform = await cls.get_friend_list(bot)
|
||||
if friend_list:
|
||||
@@ -521,6 +574,11 @@ class PlatformUtils:
|
||||
返回:
|
||||
list[FriendUser]: 好友列表
|
||||
"""
|
||||
if cls.get_platform_scope(bot) == "qq_api":
|
||||
logger.warning(
|
||||
"QQ 官方适配器不支持旧好友列表查询,已返回空列表。", "好友列表"
|
||||
)
|
||||
return [], cls.get_platform(bot)
|
||||
if interface := get_interface(bot):
|
||||
user_list = await interface.get_users()
|
||||
return [
|
||||
@@ -602,14 +660,9 @@ class BroadcastEngine:
|
||||
except KeyError:
|
||||
logger.warning(f"Bot:{i} 对象未连接或不存在", log_cmd)
|
||||
if not self.bot_list:
|
||||
try:
|
||||
bot = nonebot.get_bot()
|
||||
bot = PlatformUtils._resolve_unique_qq_client_bot(log_cmd)
|
||||
if bot is not None:
|
||||
self.bot_list.append(bot)
|
||||
logger.warning(
|
||||
f"广播任务未传入Bot对象,使用默认Bot {bot.self_id}", log_cmd
|
||||
)
|
||||
except Exception as e:
|
||||
raise ValueError("当前没有可用的Bot对象...", log_cmd) from e
|
||||
|
||||
async def call_check(self, bot: Bot, group_id: str) -> bool:
|
||||
"""运行发送检测函数
|
||||
|
||||
Reference in New Issue
Block a user