From f4d23426935756106fe0a6c344fff643a6e42ae1 Mon Sep 17 00:00:00 2001 From: Copaan <98086483+Copaan@users.noreply.github.com> Date: Wed, 24 Jun 2026 09:11:03 +0800 Subject: [PATCH] =?UTF-8?q?bugfix:=E4=BF=AE=E5=A4=8Dsqlite=E9=83=A8?= =?UTF-8?q?=E5=88=86=E5=9C=BA=E6=99=AF=E4=B8=8B=E9=94=81=E7=AB=9E=E4=BA=89?= =?UTF-8?q?=E9=97=AE=E9=A2=98=EF=BC=9B=E6=96=B0=E5=A2=9E=E6=8F=92=E4=BB=B6?= =?UTF-8?q?=E6=81=B6=E6=84=8F=E8=A7=A6=E5=8F=91=E9=85=8D=E7=BD=AE=E6=96=87?= =?UTF-8?q?=E4=BB=B6=20(#2144)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * bugfix:修复sqlite部分场景下锁竞争问题;新增插件恶意触发配置文件 * 文件没同步完 --- .../chat_history/chat_message.py | 90 ++--- zhenxun/builtin_plugins/help/data_source.py | 98 ++++-- zhenxun/builtin_plugins/hooks/__init__.py | 18 + zhenxun/builtin_plugins/hooks/auth_checker.py | 17 +- .../hooks/auth_event_selector.py | 49 ++- .../builtin_plugins/hooks/auth_pipeline.py | 4 + zhenxun/builtin_plugins/hooks/auth_profile.py | 3 +- .../builtin_plugins/hooks/auth_snapshot.py | 108 +++++- zhenxun/builtin_plugins/hooks/chkdsk_hook.py | 149 ++++++-- zhenxun/builtin_plugins/info/my_info.py | 65 +++- zhenxun/builtin_plugins/init/__init_cache.py | 17 +- .../plugin_store/data_source.py | 7 +- .../statistics/statistics_hook.py | 70 ++-- .../web_ui/api/tabs/dashboard/__init__.py | 10 +- .../web_ui/api/tabs/dashboard/data_source.py | 129 ++++--- .../web_ui/api/tabs/main/__init__.py | 21 +- .../web_ui/api/tabs/main/data_source.py | 150 +++++--- .../web_ui/api/tabs/manage/__init__.py | 6 +- .../web_ui/api/tabs/manage/data_source.py | 51 ++- .../web_ui/api/tabs/plugin_manage/store.py | 1 + zhenxun/builtin_plugins/web_ui/utils.py | 18 + zhenxun/models/ban_console.py | 31 +- zhenxun/models/group_console.py | 20 +- zhenxun/models/group_plugin_setting.py | 6 - zhenxun/services/buffered_writers.py | 125 ++----- zhenxun/services/cache/__init__.py | 30 +- zhenxun/services/cache/cache_containers.py | 282 --------------- zhenxun/services/cache/runtime_cache.py | 86 ++++- zhenxun/services/data_access.py | 50 +-- zhenxun/services/db_context/__init__.py | 1 + zhenxun/services/db_context/base_model.py | 16 +- zhenxun/services/db_context/config.py | 4 +- zhenxun/services/db_context/utils.py | 48 ++- zhenxun/services/db_context/watchdog.py | 123 +++++++ zhenxun/services/group_settings_service.py | 18 +- zhenxun/services/hot_query_cache.py | 126 +++++-- zhenxun/services/low_priority_writer.py | 329 ++++++++++++++++++ zhenxun/services/memory_governor.py | 4 +- zhenxun/services/message_load.py | 26 +- 39 files changed, 1577 insertions(+), 829 deletions(-) create mode 100644 zhenxun/services/db_context/watchdog.py create mode 100644 zhenxun/services/low_priority_writer.py diff --git a/zhenxun/builtin_plugins/chat_history/chat_message.py b/zhenxun/builtin_plugins/chat_history/chat_message.py index 60780964..90866c90 100644 --- a/zhenxun/builtin_plugins/chat_history/chat_message.py +++ b/zhenxun/builtin_plugins/chat_history/chat_message.py @@ -1,10 +1,6 @@ -import asyncio -import time - from nonebot import on_message from nonebot.plugin import PluginMetadata from nonebot_plugin_alconna import UniMsg -from nonebot_plugin_apscheduler import scheduler from nonebot_plugin_uninfo import Uninfo from zhenxun.configs.config import Config @@ -12,7 +8,12 @@ from zhenxun.configs.utils import PluginExtraData, RegisterConfig from zhenxun.models.chat_history import ChatHistory from zhenxun.services.db_context import with_db_timeout from zhenxun.services.log import logger -from zhenxun.services.message_load import is_overloaded, should_pause_tasks +from zhenxun.services.low_priority_writer import ( + LowPriorityWriterConfig, + append_low_priority_record, + register_low_priority_writer, +) +from zhenxun.services.message_load import is_overloaded from zhenxun.utils.enum import PluginType from zhenxun.utils.utils import get_entity_ids @@ -44,23 +45,45 @@ def rule(message: UniMsg) -> bool: chat_history = on_message(rule=rule, priority=1, block=False) -_HISTORY_QUEUE: asyncio.Queue[ChatHistory] = asyncio.Queue(maxsize=5000) -_DROP_COUNT = 0 -_LAST_DROP_LOG = 0.0 -_DROP_LOG_INTERVAL = 10.0 +_WRITER_NAME = "chat_history" _FLUSH_BATCH_SIZE = 200 _FLUSH_MAX_PER_TICK = 1000 _FLUSH_DB_TIMEOUT = 5.0 +async def _write_chat_history_batch(batch: list[ChatHistory], reason: str) -> None: + await with_db_timeout( + ChatHistory.bulk_create(batch, _FLUSH_BATCH_SIZE), + timeout=_FLUSH_DB_TIMEOUT, + operation=f"ChatHistory.bulk_create[{len(batch)}]", + source=f"chat_history:{reason}", + ) + + +register_low_priority_writer( + LowPriorityWriterConfig( + name=_WRITER_NAME, + write_batch=_write_chat_history_batch, + batch_size=_FLUSH_BATCH_SIZE, + trigger_size=_FLUSH_BATCH_SIZE, + max_retain=5000, + flush_interval_seconds=60.0, + max_items_per_cycle=_FLUSH_MAX_PER_TICK, + backoff_base_seconds=30.0, + backoff_max_seconds=600.0, + log_command="chat_history", + ) +) + + @chat_history.handle() async def _(message: UniMsg, session: Uninfo): entity = get_entity_ids(session) - now = time.time() if is_overloaded(): return try: - _HISTORY_QUEUE.put_nowait( + await append_low_priority_record( + _WRITER_NAME, ChatHistory( user_id=entity.user_id, group_id=entity.group_id, @@ -68,50 +91,7 @@ async def _(message: UniMsg, session: Uninfo): plain_text=message.extract_plain_text(), bot_id=session.self_id, platform=session.platform, - ) + ), ) - except asyncio.QueueFull: - global _DROP_COUNT, _LAST_DROP_LOG - _DROP_COUNT += 1 - if now - _LAST_DROP_LOG > _DROP_LOG_INTERVAL: - _LAST_DROP_LOG = now - logger.debug( - f"chat_history queue full, dropped {_DROP_COUNT} items", - "chat_history", - ) - - -@scheduler.scheduled_job( - "interval", - minutes=1, - max_instances=1, - coalesce=True, -) -async def _(): - try: - if should_pause_tasks(): - return - flushed = 0 - while flushed < _FLUSH_MAX_PER_TICK: - message_list: list[ChatHistory] = [] - limit = min(_FLUSH_BATCH_SIZE, _FLUSH_MAX_PER_TICK - flushed) - for _ in range(limit): - try: - message_list.append(_HISTORY_QUEUE.get_nowait()) - except asyncio.QueueEmpty: - break - if not message_list: - break - await with_db_timeout( - ChatHistory.bulk_create(message_list, _FLUSH_BATCH_SIZE), - timeout=_FLUSH_DB_TIMEOUT, - operation=f"ChatHistory.bulk_create[{len(message_list)}]", - source="chat_history", - ) - flushed += len(message_list) - if flushed: - backlog = _HISTORY_QUEUE.qsize() - suffix = f",剩余队列 {backlog} 条" if backlog else "" - logger.debug(f"批量添加聊天记录 {flushed} 条{suffix}", "定时任务") except Exception as e: logger.warning("存储聊天记录失败", "chat_history", e=e) diff --git a/zhenxun/builtin_plugins/help/data_source.py b/zhenxun/builtin_plugins/help/data_source.py index f7235feb..5515040b 100644 --- a/zhenxun/builtin_plugins/help/data_source.py +++ b/zhenxun/builtin_plugins/help/data_source.py @@ -15,7 +15,9 @@ from zhenxun.services import ( avatar_service, generate, ) +from zhenxun.services.db_context import with_db_timeout from zhenxun.services.log import logger +from zhenxun.services.message_load import is_db_unhealthy from zhenxun.services.renderer.result_cache import RenderResultMemoryCache from zhenxun.ui.models import PluginMenuCategory, PluginMenuData from zhenxun.utils.common_utils import format_usage_for_markdown @@ -25,6 +27,8 @@ from zhenxun.utils.platform import PlatformUtils from .utils import classify_plugin driver = nonebot.get_driver() +_DB_BUSY_MESSAGE = "数据库繁忙,请稍后再试" +_HELP_DB_TIMEOUT = 3.0 _HELP_MENU_IMAGE_CACHE = RenderResultMemoryCache( ttl_seconds=300, max_items=64, @@ -32,6 +36,24 @@ _HELP_MENU_IMAGE_CACHE = RenderResultMemoryCache( ) +class _DbBusyError(Exception): + pass + + +async def _read_db(factory, operation: str): + if is_db_unhealthy(): + raise _DbBusyError + try: + return await with_db_timeout( + factory(), + timeout=_HELP_DB_TIMEOUT, + operation=operation, + source="help", + ) + except TimeoutError as exc: + raise _DbBusyError from exc + + def _create_plugin_menu_item( bot: BotConsole | None, plugin: PluginInfo, @@ -79,11 +101,17 @@ def _create_plugin_menu_item( async def create_help_img( session: Uninfo, group_id: str | None, is_detail: bool -) -> bytes: +) -> str | bytes: """使用渲染服务生成帮助图片""" - classified_data = await classify_plugin( - session, group_id, is_detail, _create_plugin_menu_item - ) + try: + classified_data = await _read_db( + lambda: classify_plugin( + session, group_id, is_detail, _create_plugin_menu_item + ), + "Help.classify_plugin", + ) + except _DbBusyError: + return _DB_BUSY_MESSAGE sorted_categories = dict( sorted(classified_data.items(), key=lambda x: len(x[1]), reverse=True) @@ -157,9 +185,11 @@ async def get_user_allow_help(user_id: str) -> list[str]: list[str]: 插件类型列表 """ type_list = ["NORMAL", "DEPENDANT"] - for level in await LevelUser.filter(user_id=user_id).values_list( - "user_level", flat=True - ): + levels = await _read_db( + lambda: LevelUser.filter(user_id=user_id).values_list("user_level", flat=True), + "Help.user_allow_level", + ) + for level in levels: if level > 0: # type: ignore type_list.extend(("ADMIN", "ADMIN_SUPER")) break @@ -170,7 +200,7 @@ async def get_user_allow_help(user_id: str) -> list[str]: async def get_plugin_help( user_id: str, name: str, is_superuser: bool, variant: str | None = None -) -> bytes | None: +) -> str | bytes | None: """获取功能的帮助信息 参数: @@ -179,20 +209,35 @@ async def get_plugin_help( is_superuser: 是否为超级用户 variant: 使用的皮肤/变体名称 """ - type_list = await get_user_allow_help(user_id) - if name.isdigit(): - plugin = await PluginInfo.get_or_none(id=int(name), plugin_type__in=type_list) - else: - plugin = await PluginInfo.get_or_none( - name__iexact=name, load_status=True, plugin_type__in=type_list - ) + try: + type_list = await get_user_allow_help(user_id) + if name.isdigit(): + plugin = await _read_db( + lambda: PluginInfo.get_or_none(id=int(name), plugin_type__in=type_list), + "Help.plugin_by_id", + ) + else: + plugin = await _read_db( + lambda: PluginInfo.get_or_none( + name__iexact=name, load_status=True, plugin_type__in=type_list + ), + "Help.plugin_by_name", + ) + except _DbBusyError: + return _DB_BUSY_MESSAGE if plugin: _plugin = nonebot.get_plugin_by_module_name(plugin.module_path) if _plugin and _plugin.metadata: extra_data = PluginExtraData(**_plugin.metadata.extra) - call_count = await Statistics.filter(plugin_name=plugin.module).count() + try: + call_count = await _read_db( + lambda: Statistics.filter(plugin_name=plugin.module).count(), + "Help.plugin_call_count", + ) + except _DbBusyError: + return _DB_BUSY_MESSAGE usage = _plugin.metadata.usage metadata_items = [ @@ -263,14 +308,19 @@ async def get_llm_help(question: str, user_id: str) -> str | bytes: """ try: - allowed_types = await get_user_allow_help(user_id) - - plugins = await PluginInfo.get_plugins( - load_status=None, - filter_parent=False, - is_show=True, - plugin_type__in=allowed_types, - ) + try: + allowed_types = await get_user_allow_help(user_id) + plugins = await _read_db( + lambda: PluginInfo.get_plugins( + load_status=None, + filter_parent=False, + is_show=True, + plugin_type__in=allowed_types, + ), + "Help.llm_plugin_list", + ) + except _DbBusyError: + return _DB_BUSY_MESSAGE knowledge_base_parts = [] for p in plugins: diff --git a/zhenxun/builtin_plugins/hooks/__init__.py b/zhenxun/builtin_plugins/hooks/__init__.py index 3ad29d71..25c82a73 100644 --- a/zhenxun/builtin_plugins/hooks/__init__.py +++ b/zhenxun/builtin_plugins/hooks/__init__.py @@ -40,6 +40,24 @@ Config.add_plugin_config( type=int, ) +Config.add_plugin_config( + "hook", + "MALICIOUS_CHECK_MODE", + "off", + help="恶意触发检测模式:off=关闭,blacklist=仅列表插件检测,whitelist=列表插件跳过检测", + default_value="off", + type=str, +) + +Config.add_plugin_config( + "hook", + "MALICIOUS_CHECK_PLUGINS", + [], + help="恶意触发检测插件列表,按模式作为黑名单或白名单使用,填插件模块名", + default_value=[], + type=list, +) + Config.add_plugin_config( "hook", "IS_SEND_TIP_MESSAGE", diff --git a/zhenxun/builtin_plugins/hooks/auth_checker.py b/zhenxun/builtin_plugins/hooks/auth_checker.py index 9c8636fb..300a0f33 100644 --- a/zhenxun/builtin_plugins/hooks/auth_checker.py +++ b/zhenxun/builtin_plugins/hooks/auth_checker.py @@ -17,7 +17,7 @@ from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.user_console import UserConsole from zhenxun.services.cache.cache_containers import CacheDict from zhenxun.services.log import logger -from zhenxun.services.message_load import signal_overload +from zhenxun.services.message_load import is_db_unhealthy, signal_overload from zhenxun.utils.enum import GoldHandle, PluginType from zhenxun.utils.exception import InsufficientGold from zhenxun.utils.platform import PlatformUtils @@ -811,6 +811,8 @@ async def _run_selected_matcher( dependency_cache, lane: str = "command_exact", ) -> None: + # state is copied per matcher before dispatch; keep lane matcher-local. + state["_zx_dispatch_lane"] = lane async with _dispatch_lane_section(lane): await nb_message.check_and_run_matcher( matcher, @@ -938,6 +940,11 @@ async def _has_limits_cached( event_cache.setdefault("module_limit_entries", {})[module] = limit_entries event_cache.setdefault("module_limits_ready", {})[module] = True return has_limits + if is_db_unhealthy(): + module_limit_cache[module] = False + if event_cache is not None: + event_cache.setdefault("module_limits_ready", {})[module] = False + return False limits = await LimitManager.get_module_limits(module) has_limits = bool(limits) module_limit_cache[module] = has_limits @@ -1075,7 +1082,7 @@ async def _get_plugin_cache_first( 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: + if plugin is None and allow_cache_load and not is_db_unhealthy(): plugin = await provider.get_plugin(module) cache_miss = False if event_cache is not None: @@ -1324,6 +1331,9 @@ async def _prepare_auth_state_with_fallback( ) if prep is not None: return prep + if is_db_unhealthy(): + hook_recorder.set("auth_snapshot", "cache_miss_db_unhealthy") + return None hook_recorder.set("auth_snapshot", "cache_miss_fallback") return await _prepare_auth_state( module=module, @@ -1408,6 +1418,9 @@ async def _resolve_cost_gold( if prep.profile.cost_gold <= 0: hook_recorder.set("cost_gold", "skipped") return 0 + if is_db_unhealthy(): + hook_recorder.set("cost_gold", "db_unhealthy") + raise SkipPluginException("数据库繁忙,金币功能暂不可用...") cost_start = time.time() try: if prep.user is None: diff --git a/zhenxun/builtin_plugins/hooks/auth_event_selector.py b/zhenxun/builtin_plugins/hooks/auth_event_selector.py index c01fec31..5ecaefd1 100644 --- a/zhenxun/builtin_plugins/hooks/auth_event_selector.py +++ b/zhenxun/builtin_plugins/hooks/auth_event_selector.py @@ -39,6 +39,49 @@ class HandleEventSelectorDependencies: _HANDLE_EVENT_PATCHED = False _ORIGINAL_HANDLE_EVENT: Callable[..., Awaitable[None]] | None = None _ORIGINAL_ADAPTER_HANDLE_EVENTS: dict[object, object] = {} +_MATCHER_DEADLINE_BY_LANE = { + "system": 15.0, + "command": 12.0, + "temp": 15.0, + "passive_light": 3.0, + "fallback_ai": 5.0, +} +_DEFAULT_MATCHER_DEADLINE = 5.0 + + +def _matcher_deadline_for_lane(lane: str) -> float: + return _MATCHER_DEADLINE_BY_LANE.get(lane, _DEFAULT_MATCHER_DEADLINE) + + +def _matcher_name(matcher: type[Matcher]) -> str: + module = str(getattr(matcher, "module", "") or "") + lineno = str(getattr(matcher, "lineno", "") or "") + matcher_type = str(getattr(matcher, "type", "") or "") + name = module or matcher.__name__ + if lineno: + name = f"{name}:{lineno}" + if matcher_type: + name = f"{name}<{matcher_type}>" + return name + + +async def _run_matcher_with_deadline( + anyio_mod: Any, + coro: Awaitable[None], + matcher: type[Matcher], + lane: str, +) -> None: + timeout = _matcher_deadline_for_lane(lane) + try: + with anyio_mod.fail_after(timeout): + await coro + except TimeoutError: + signal_overload(20.0) + logger.warning( + "matcher dispatch timeout: " + f"matcher={_matcher_name(matcher)}, lane={lane}, timeout={timeout:.1f}s", + LOGGER_COMMAND, + ) def _trim_leading_text(message: Any) -> None: @@ -140,7 +183,6 @@ async def patched_handle_event( stop_propagation = getattr(nb_message, "StopPropagation") handle_exception = getattr(nb_message, "_handle_exception") anyio_mod = getattr(nb_message, "anyio") - run_coro_with_shield = getattr(nb_message, "run_coro_with_shield") async with async_exit_stack() as stack: if not await apply_event_preprocessors( @@ -256,7 +298,8 @@ async def patched_handle_event( continue matcher_state = deps.build_matcher_state(state) tg.start_soon( - run_coro_with_shield, + _run_matcher_with_deadline, + anyio_mod, deps.run_selected_matcher( matcher, bot, @@ -266,6 +309,8 @@ async def patched_handle_event( dependency_cache, lane, ), + matcher, + lane, ) if show_log: diff --git a/zhenxun/builtin_plugins/hooks/auth_pipeline.py b/zhenxun/builtin_plugins/hooks/auth_pipeline.py index e59f246d..11fc14aa 100644 --- a/zhenxun/builtin_plugins/hooks/auth_pipeline.py +++ b/zhenxun/builtin_plugins/hooks/auth_pipeline.py @@ -10,6 +10,7 @@ from nonebot.adapters import Bot, Event from nonebot.matcher import Matcher from nonebot_plugin_uninfo import Uninfo +from zhenxun.services.message_load import is_db_unhealthy from zhenxun.utils.utils import EntityIDs from .auth.context import ( @@ -296,6 +297,9 @@ async def policy_precheck_stage( ctx.flags = apply_policy_precheck(ctx, deps) except PermissionExemption as exc: _recorder(ctx).set("policy_fallback", str(exc)) + if is_db_unhealthy(): + ctx.stop(allowed=True, effect="allow", reason="db_unhealthy_cache_miss") + return ctx.prep = await deps.prepare_auth_state( module=ctx.module, context=ctx.event_context, diff --git a/zhenxun/builtin_plugins/hooks/auth_profile.py b/zhenxun/builtin_plugins/hooks/auth_profile.py index 16128d2a..aa68c388 100644 --- a/zhenxun/builtin_plugins/hooks/auth_profile.py +++ b/zhenxun/builtin_plugins/hooks/auth_profile.py @@ -2,6 +2,7 @@ from __future__ import annotations from dataclasses import dataclass +from zhenxun.services.message_load import is_db_unhealthy from zhenxun.utils.enum import BlockType, PluginType from .auth.data_provider import ( @@ -103,7 +104,7 @@ async def get_plugin_auth_profile( if limits is None: limits = provider.get_module_limits_if_ready(module) limits_ready = limits is not None - if limits is None and allow_cache_load: + if limits is None and allow_cache_load and not is_db_unhealthy(): limits = await provider.get_module_limits(module) limits_ready = True if limits is None: diff --git a/zhenxun/builtin_plugins/hooks/auth_snapshot.py b/zhenxun/builtin_plugins/hooks/auth_snapshot.py index 9cc6e81c..4034e8ca 100644 --- a/zhenxun/builtin_plugins/hooks/auth_snapshot.py +++ b/zhenxun/builtin_plugins/hooks/auth_snapshot.py @@ -10,7 +10,9 @@ from zhenxun.services.cache.runtime_cache import ( GroupSnapshot, LevelUserSnapshot, ) +from zhenxun.services.db_context import with_db_timeout from zhenxun.services.log import logger +from zhenxun.services.message_load import is_db_unhealthy from .auth.config import LOGGER_COMMAND from .auth.context import EventContext @@ -50,6 +52,41 @@ def _build_runtime_group_snapshot(context: EventContext) -> GroupSnapshot | None ) +def _build_default_bot_snapshot(context: EventContext) -> BotSnapshot: + """Fail-open bot snapshot used only while DB cold-path is unhealthy.""" + return BotSnapshot( + bot_id=context.bot_id, + status=True, + platform=context.platform, + block_plugins="", + block_tasks="", + available_plugins="", + available_tasks="", + ) + + +def _build_default_group_snapshot(context: EventContext) -> GroupSnapshot | None: + """Fail-open group snapshot used only while DB cold-path is unhealthy.""" + if 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 @@ -72,6 +109,8 @@ async def _repair_missing_qq_client_group( provider: PermissionDataProvider, ) -> GroupSnapshot | None: """Persist a minimal OneBot group when startup group sync returned empty.""" + if is_db_unhealthy(): + return None key = _qq_client_group_repair_key(context) if key is None or not provider.group_cache_loaded(): return None @@ -97,9 +136,14 @@ async def _repair_missing_qq_client_group( "group_flag": 1, "platform": context.platform, } - group, _ = await GroupConsole.get_or_create_root_group( - group_id=group_id, - defaults=defaults, + group, _ = await with_db_timeout( + GroupConsole.get_or_create_root_group( + group_id=group_id, + defaults=defaults, + ), + timeout=2.0, + operation="GroupConsole.get_or_create_root_group", + source="auth_snapshot.repair_missing_group", ) from zhenxun.services.cache.runtime_cache import GroupMemoryCache @@ -129,6 +173,7 @@ class AuthSnapshot: ban_state: bool | None = None user_balance_loaded: bool = False user_balance: int | None = None + db_unhealthy: bool = False cache_misses: frozenset[str] = field(default_factory=frozenset) @property @@ -173,43 +218,63 @@ async def build_auth_snapshot( event_cache = context.event_cache entity = context.entity cache_misses: set[str] = set() + db_unhealthy = is_db_unhealthy() + can_load_cache = allow_cache_load and not db_unhealthy bot_data: BotSnapshot | None = None if ( event_cache is not None and "bot_data" in event_cache - and (event_cache.get("bot_cache_ready") or not allow_cache_load) + and (event_cache.get("bot_cache_ready") or not can_load_cache) ): bot_data = event_cache.get("bot_data") else: bot_data = provider.get_bot_if_ready(bot.self_id) if bot_data is None: - if allow_cache_load: + if can_load_cache: bot_data = await provider.get_bot(bot.self_id) + elif db_unhealthy: + bot_data = _build_default_bot_snapshot(context) 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"] = provider.bot_cache_loaded() + event_cache["bot_cache_ready"] = provider.bot_cache_loaded() or db_unhealthy + if bot_data is None and db_unhealthy: + bot_data = _build_default_bot_snapshot(context) + if event_cache is not None: + event_cache["bot_data"] = bot_data + event_cache["bot_cache_ready"] = True group = None if entity.group_id: if ( event_cache is not None and "group" in event_cache - and (event_cache.get("group_cache_ready") or not allow_cache_load) + and (event_cache.get("group_cache_ready") or not can_load_cache) ): group = event_cache.get("group") else: 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: + elif group is None and can_load_cache: group = await provider.get_group(entity.group_id, entity.channel_id) + if group is None and db_unhealthy: + group = _build_default_group_snapshot(context) + cache_misses.discard("group") if event_cache is not None: event_cache["group"] = group - event_cache["group_cache_ready"] = provider.group_cache_loaded() - if group is None: + event_cache["group_cache_ready"] = ( + provider.group_cache_loaded() or db_unhealthy + ) + if group is None and db_unhealthy: + group = _build_default_group_snapshot(context) + cache_misses.discard("group") + if event_cache is not None: + event_cache["group"] = group + event_cache["group_cache_ready"] = True + if group is None and not db_unhealthy: group = await _repair_missing_qq_client_group( context, provider=provider, @@ -232,7 +297,7 @@ async def build_auth_snapshot( if ( event_cache is not None and "admin_levels" in event_cache - and (event_cache.get("admin_cache_ready") or not allow_cache_load) + and (event_cache.get("admin_cache_ready") or not can_load_cache) ): admin_levels = event_cache.get("admin_levels") else: @@ -241,16 +306,26 @@ async def build_auth_snapshot( entity.group_id, ) if admin_levels is None: - if allow_cache_load: + if can_load_cache: admin_levels = await provider.get_admin_levels( entity.user_id, entity.group_id, ) + elif db_unhealthy: + admin_levels = (None, None) else: cache_misses.add("admin_levels") if event_cache is not None: event_cache["admin_levels"] = admin_levels - event_cache["admin_cache_ready"] = provider.admin_cache_loaded() + event_cache["admin_cache_ready"] = ( + provider.admin_cache_loaded() or db_unhealthy + ) + if admin_levels is None and db_unhealthy and not provider.admin_cache_loaded(): + admin_levels = (None, None) + cache_misses.discard("admin_levels") + if event_cache is not None: + event_cache["admin_levels"] = admin_levels + event_cache["admin_cache_ready"] = True ban_state = None if not skip_ban: @@ -260,11 +335,15 @@ async def build_auth_snapshot( 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: + elif can_load_cache: 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 + elif db_unhealthy: + ban_state = False + if event_cache is not None: + event_cache["ban_state"] = ban_state else: cache_misses.add("ban") @@ -276,6 +355,7 @@ async def build_auth_snapshot( group=group, admin_levels=admin_levels, ban_state=ban_state, + db_unhealthy=db_unhealthy, cache_misses=frozenset(cache_misses), ) diff --git a/zhenxun/builtin_plugins/hooks/chkdsk_hook.py b/zhenxun/builtin_plugins/hooks/chkdsk_hook.py index c50b663f..d0447821 100644 --- a/zhenxun/builtin_plugins/hooks/chkdsk_hook.py +++ b/zhenxun/builtin_plugins/hooks/chkdsk_hook.py @@ -61,6 +61,88 @@ _blmt = BanCheckLimiter( 4, ) +_MALICIOUS_CHECK_MODES = {"off", "blacklist", "whitelist"} +_EVENT_PLUGIN_DEDUPE_TTL = 30.0 +_EVENT_PLUGIN_DEDUPE_MAX = 4096 +_event_plugin_seen: dict[str, float] = {} + + +def _malicious_check_mode() -> str: + mode = str(Config.get_config("hook", "MALICIOUS_CHECK_MODE") or "off") + mode = mode.strip().lower() + return mode if mode in _MALICIOUS_CHECK_MODES else "off" + + +def _malicious_plugin_set() -> set[str]: + value = Config.get_config("hook", "MALICIOUS_CHECK_PLUGINS") + if value is None: + return set() + if isinstance(value, str): + items = value.replace("\n", ",").split(",") + elif isinstance(value, list | tuple | set): + items = value + else: + items = [value] + return {str(item).strip().casefold() for item in items if str(item).strip()} + + +def _should_check_plugin(module: str, lane: str) -> bool: + mode = _malicious_check_mode() + if mode == "off": + return False + + normalized_module = str(module or "").strip().casefold() + if not normalized_module: + return False + + plugin_set = _malicious_plugin_set() + in_plugin_set = normalized_module in plugin_set + is_passive = str(lane or "").startswith("passive_") + + if mode == "blacklist": + return in_plugin_set + if is_passive: + return False + if mode == "whitelist": + return not in_plugin_set + return False + + +def _event_plugin_key(event: Event, user_id: str, module: str) -> str: + message_id = getattr(event, "message_id", None) or getattr(event, "id", None) + if message_id is None: + message_id = id(event) + return f"{message_id}:{user_id}:{module}" + + +def _remember_event_plugin_once(key: str) -> bool: + now = time.monotonic() + expires_at = _event_plugin_seen.get(key) + if expires_at is not None and expires_at > now: + return False + + _event_plugin_seen[key] = now + _EVENT_PLUGIN_DEDUPE_TTL + if len(_event_plugin_seen) > _EVENT_PLUGIN_DEDUPE_MAX: + target_size = _EVENT_PLUGIN_DEDUPE_MAX // 2 + for cache_key, cache_expires_at in list(_event_plugin_seen.items()): + if cache_expires_at <= now or len(_event_plugin_seen) > target_size: + _event_plugin_seen.pop(cache_key, None) + if len(_event_plugin_seen) <= target_size: + break + return True + + +def _mark_event_plugin_checked( + state: T_State, event: Event, user_id: str, module: str +) -> bool: + checked = state.setdefault("_zx_malicious_checked_plugins", set()) + if isinstance(checked, set): + if module in checked: + return False + checked.add(module) + + return _remember_event_plugin_once(_event_plugin_key(event, user_id, module)) + def _get_positive_config(key: str, cast_type: type[int] | type[float]) -> int | float: value = Config.get_config("hook", key) @@ -102,6 +184,10 @@ async def _( else: return + lane = state.get("_zx_dispatch_lane") + if not _should_check_plugin(module, lane if isinstance(lane, str) else ""): + return + user_id = resolve_actor_user_id(event, session.id1) group_id = resolve_event_group_id(event, session.id3 or session.id2) # 超级用户豁免恶意检测(A6):与权威权限路径保持一致,避免误封管理者。 @@ -111,35 +197,42 @@ async def _( is_superuser = user_id in bot.config.superusers if is_superuser: return + else: + return + + if not _mark_event_plugin_checked(state, event, user_id, module): + return + + # 只统计通过模式/lane过滤且同事件同插件去重后的有效触发。 + limiter_key = f"{user_id}__{module}" malicious_check_time = float(_get_positive_config("MALICIOUS_CHECK_TIME", float)) malicious_ban_count = int(_get_positive_config("MALICIOUS_BAN_COUNT", int)) malicious_ban_time = int(_get_positive_config("MALICIOUS_BAN_TIME", int)) _blmt.configure(malicious_check_time, malicious_ban_count) - if user_id and module: - if _blmt.check(f"{user_id}__{module}"): - await BanConsole.ban( - user_id, - group_id, - 9, - "恶意触发命令检测", - malicious_ban_time * 60, - bot.self_id, - ) - logger.info( - f"触发了恶意触发检测: {matcher.plugin_name}", - "HOOK", - session=session, - ) - await MessageUtils.build_message( - [ - At(flag="user", target=user_id), - "检测到恶意触发命令,您将被封禁 30 分钟", - ] - ).send() - logger.debug( - f"触发了恶意触发检测: {matcher.plugin_name}", - "HOOK", - session=session, - ) - raise IgnoredException("检测到恶意触发命令") - _blmt.add(f"{user_id}__{module}") + if _blmt.check(limiter_key): + await BanConsole.ban( + user_id, + group_id, + 9, + "恶意触发命令检测", + malicious_ban_time * 60, + bot.self_id, + ) + logger.info( + f"触发了恶意触发检测: {matcher.plugin_name}", + "HOOK", + session=session, + ) + await MessageUtils.build_message( + [ + At(flag="user", target=user_id), + "检测到恶意触发命令,您将被封禁 30 分钟", + ] + ).send() + logger.debug( + f"触发了恶意触发检测: {matcher.plugin_name}", + "HOOK", + session=session, + ) + raise IgnoredException("检测到恶意触发命令") + _blmt.add(limiter_key) diff --git a/zhenxun/builtin_plugins/info/my_info.py b/zhenxun/builtin_plugins/info/my_info.py index 8e827231..bbdf5e59 100644 --- a/zhenxun/builtin_plugins/info/my_info.py +++ b/zhenxun/builtin_plugins/info/my_info.py @@ -12,6 +12,8 @@ from zhenxun.models.sign_user import SignUser from zhenxun.models.statistics import Statistics from zhenxun.models.user_console import UserConsole from zhenxun.services import avatar_service +from zhenxun.services.db_context import with_db_timeout +from zhenxun.services.message_load import is_db_unhealthy from zhenxun.utils.platform import PlatformUtils RACE = [ @@ -81,6 +83,21 @@ lik2level = { 10: 1, 0: 0, } +_INFO_DB_TIMEOUT = 3.0 + + +async def _read_db(factory, operation: str, default): + if is_db_unhealthy(): + return default + try: + return await with_db_timeout( + factory(), + timeout=_INFO_DB_TIMEOUT, + operation=operation, + source="my_info", + ) + except Exception: + return default def get_level(impression: float) -> int: @@ -103,13 +120,17 @@ async def get_chat_history( """ now = datetime.now() filter_date = now - timedelta(days=7) - date_list = ( - await ChatHistory.filter( - user_id=user_id, group_id=group_id, create_time__gte=filter_date + date_list = await _read_db( + lambda: ChatHistory.filter( + user_id=user_id, + group_id=group_id, + create_time__gte=filter_date, ) .annotate(date=RawSQL("DATE(create_time)"), count=Count("id")) .group_by("date") - .values("date", "count") + .values("date", "count"), + "MyInfo.chat_history_chart", + [], ) chart_date: list[str] = [] count_list: list[int] = [] @@ -143,20 +164,40 @@ async def get_user_info( avatar_path = await avatar_service.get_avatar_path(platform, user_id) avatar_url = avatar_path.as_uri() if avatar_path else "" - user = await UserConsole.get_user(user_id, platform) - permission_level = await LevelUser.get_user_level(user_id, group_id) + user = await _read_db( + lambda: UserConsole.get_user(user_id, platform), + "MyInfo.user_console", + None, + ) + permission_level = await _read_db( + lambda: LevelUser.get_user_level(user_id, group_id), + "MyInfo.level_user", + 0, + ) sign_level = 0 - if sign_user := await SignUser.get_or_none(user_id=user_id): + if sign_user := await _read_db( + lambda: SignUser.get_or_none(user_id=user_id), + "MyInfo.sign_user", + None, + ): sign_level = get_level(float(sign_user.impression)) - chat_count = await ChatHistory.filter(user_id=user_id, group_id=group_id).count() - stat_count = await Statistics.filter(user_id=user_id, group_id=group_id).count() + chat_count = await _read_db( + lambda: ChatHistory.filter(user_id=user_id, group_id=group_id).count(), + "MyInfo.chat_count", + 0, + ) + stat_count = await _read_db( + lambda: Statistics.filter(user_id=user_id, group_id=group_id).count(), + "MyInfo.stat_count", + 0, + ) selected_indices = [""] * 9 selected_indices[sign_level] = "select" - uid = f"{user.uid}".rjust(8, "0") + uid = f"{getattr(user, 'uid', 0)}".rjust(8, "0") uid_formatted = f"{uid[:4]} {uid[4:]}" now = datetime.now() @@ -182,8 +223,8 @@ async def get_user_info( ), }, "stats": { - "gold": user.gold, - "prop_count": len(user.props), + "gold": getattr(user, "gold", 0), + "prop_count": len(getattr(user, "props", {}) or {}), "call_count": stat_count, "chat_count": chat_count, }, diff --git a/zhenxun/builtin_plugins/init/__init_cache.py b/zhenxun/builtin_plugins/init/__init_cache.py index 8938c0a0..94c24911 100644 --- a/zhenxun/builtin_plugins/init/__init_cache.py +++ b/zhenxun/builtin_plugins/init/__init_cache.py @@ -4,10 +4,7 @@ 负责注册各种缓存类型,实现按需缓存机制 """ -from zhenxun.models.ban_console import BanConsole from zhenxun.models.bot_console import BotConsole -from zhenxun.models.group_console import GroupConsole -from zhenxun.models.group_plugin_setting import GroupPluginSetting from zhenxun.models.level_user import LevelUser from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.user_console import UserConsole @@ -21,21 +18,11 @@ from zhenxun.utils.enum import CacheType def register_cache_types(): """注册所有缓存类型""" CacheRegistry.register(CacheType.PLUGINS, PluginInfo) - CacheRegistry.register(CacheType.GROUPS, GroupConsole) CacheRegistry.register(CacheType.BOT, BotConsole) CacheRegistry.register(CacheType.USERS, UserConsole) - CacheRegistry.register( - CacheType.GROUP_PLUGIN_SETTINGS, - GroupPluginSetting, - key_format="{group_id}_{plugin_name}_{key}", - ) CacheRegistry.register( CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}" ) - CacheRegistry.register(CacheType.BAN, BanConsole, key_format="{user_id}_{group_id}") - if cache_config.cache_mode == CacheMode.NONE: - logger.info("缓存功能已禁用,将直接从数据库获取数据") - else: - logger.info(f"已注册所有缓存类型,缓存模式: {cache_config.cache_mode}") - logger.info("使用增量缓存模式,数据将按需加载到缓存中") + if cache_config.cache_mode == CacheMode.REDIS and cache_config.redis_host: + logger.info(f"已注册 Redis 模型缓存类型,缓存模式: {cache_config.cache_mode}") diff --git a/zhenxun/builtin_plugins/plugin_store/data_source.py b/zhenxun/builtin_plugins/plugin_store/data_source.py index 3bf60f37..7ae89805 100644 --- a/zhenxun/builtin_plugins/plugin_store/data_source.py +++ b/zhenxun/builtin_plugins/plugin_store/data_source.py @@ -314,7 +314,12 @@ class StoreManager: is_external, source, ) - return f"插件 {plugin_info.name} 安装成功! 重启后生效" + return ( + f"插件 {plugin_info.name} 安装完成\n" + "- 已下载插件文件\n" + "- 已处理依赖文件\n" + "- 重启后生效" + ) @classmethod async def install_plugin_with_repo( diff --git a/zhenxun/builtin_plugins/statistics/statistics_hook.py b/zhenxun/builtin_plugins/statistics/statistics_hook.py index 1c323a56..a6051edd 100644 --- a/zhenxun/builtin_plugins/statistics/statistics_hook.py +++ b/zhenxun/builtin_plugins/statistics/statistics_hook.py @@ -7,14 +7,19 @@ from nonebot.adapters.onebot.v11 import PokeNotifyEvent from nonebot.matcher import Matcher from nonebot.message import run_postprocessor from nonebot.plugin import PluginMetadata -from nonebot_plugin_apscheduler import scheduler from nonebot_plugin_uninfo import Uninfo from zhenxun.configs.utils import PluginExtraData from zhenxun.models.statistics import Statistics from zhenxun.services.cache.runtime_cache import PluginInfoMemoryCache +from zhenxun.services.db_context import with_db_timeout from zhenxun.services.log import logger -from zhenxun.services.message_load import should_pause_tasks +from zhenxun.services.low_priority_writer import ( + LowPriorityWriterConfig, + append_low_priority_record, + flush_low_priority_writer, + register_low_priority_writer, +) from zhenxun.utils.enum import PluginType from zhenxun.utils.utils import get_entity_ids @@ -29,37 +34,44 @@ __plugin_meta__ = PluginMetadata( STATS_BUFFER_FLUSH_SIZE = 5000 STATS_BUFFER_MAX_RETAIN = 10000 -TEMP_LIST: list[Statistics] = [] _STATS_FLUSH_LOCK = asyncio.Lock() +_WRITER_NAME = "statistics" driver = get_driver() +async def _write_statistics_batch(batch: list[Statistics], reason: str) -> None: + await with_db_timeout( + Statistics.bulk_create(batch), + timeout=5.0, + operation=f"Statistics.bulk_create[{len(batch)}]", + source=f"statistics:{reason}", + ) + + +register_low_priority_writer( + LowPriorityWriterConfig( + name=_WRITER_NAME, + write_batch=_write_statistics_batch, + batch_size=STATS_BUFFER_FLUSH_SIZE, + trigger_size=STATS_BUFFER_FLUSH_SIZE, + max_retain=STATS_BUFFER_MAX_RETAIN, + flush_interval_seconds=30 * 60, + max_items_per_cycle=STATS_BUFFER_FLUSH_SIZE, + backoff_base_seconds=30.0, + backoff_max_seconds=600.0, + log_command="定时任务", + ) +) + + async def _flush_statistics_buffer(reason: str) -> int: + """Compatibility entry used by memory governor and shutdown hooks.""" async with _STATS_FLUSH_LOCK: - call_list = TEMP_LIST.copy() - TEMP_LIST.clear() - if not call_list: - return 0 - try: - await Statistics.bulk_create(call_list) - except Exception as e: - logger.error(f"{reason}批量添加调用记录失败", "定时任务", e=e) - retain_count = max(STATS_BUFFER_MAX_RETAIN - len(TEMP_LIST), 0) - if retain_count: - TEMP_LIST[:0] = call_list[-retain_count:] - return 0 - logger.debug(f"{reason}批量添加调用记录 {len(call_list)} 条", "定时任务") - return len(call_list) + return await flush_low_priority_writer(_WRITER_NAME, reason, force=True) async def _append_statistics(record: Statistics) -> None: - # 在锁内追加(B8),与 flush 的 copy+clear 串行,消除逻辑窗口; - # flush 自身再次获取同一把锁,故在锁外触发避免重入。 - async with _STATS_FLUSH_LOCK: - TEMP_LIST.append(record) - should_flush = len(TEMP_LIST) >= STATS_BUFFER_FLUSH_SIZE - if should_flush: - await _flush_statistics_buffer("缓冲区触发") + await append_low_priority_record(_WRITER_NAME, record) @run_postprocessor @@ -95,16 +107,6 @@ async def _( ) -@scheduler.scheduled_job("interval", minutes=30, max_instances=1, coalesce=True) -async def _(): - try: - if should_pause_tasks(): - return - await _flush_statistics_buffer("定时") - except Exception as e: - logger.error("定时批量添加调用记录", "定时任务", e=e) - - @driver.on_shutdown async def _flush_statistics_on_shutdown(): await _flush_statistics_buffer("关闭") diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/dashboard/__init__.py b/zhenxun/builtin_plugins/web_ui/api/tabs/dashboard/__init__.py index fa719cdf..ea5a699c 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/dashboard/__init__.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/dashboard/__init__.py @@ -7,7 +7,7 @@ from nonebot.config import Config from zhenxun.services.log import logger from ....base_model import BaseResultModel, QueryModel, Result -from ....utils import authentication +from ....utils import DB_BUSY_MESSAGE, authentication from .data_source import ApiDataSource from .model import AllChatAndCallCount, BotInfo, ChatCallMonthCount, QueryChatCallCount @@ -28,6 +28,8 @@ driver = nonebot.get_driver() async def _() -> Result[list[BotInfo]]: try: return Result.ok(await ApiDataSource.get_bot_list(), "拿到信息啦!") + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error(f"{router.prefix}/get_bot_list 调用错误", "WebUi", e=e) return Result.fail(f"发生了一点错误捏 {type(e)}: {e}") @@ -45,6 +47,8 @@ async def _(bot_id: str | None = None) -> Result[QueryChatCallCount]: return Result.ok( await ApiDataSource.get_chat_and_call_count(bot_id), "拿到信息啦!" ) + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error(f"{router.prefix}/get_chat_and_call_count 调用错误", "WebUi", e=e) return Result.fail(f"发生了一点错误捏 {type(e)}: {e}") @@ -62,6 +66,8 @@ async def _(bot_id: str | None = None) -> Result[AllChatAndCallCount]: return Result.ok( await ApiDataSource.get_all_chat_and_call_count(bot_id), "拿到信息啦!" ) + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error( f"{router.prefix}/get_all_chat_and_call_count 调用错误", "WebUi", e=e @@ -81,6 +87,8 @@ async def _(bot_id: str | None = None) -> Result[ChatCallMonthCount]: return Result.ok( await ApiDataSource.get_chat_and_call_month(bot_id), "拿到信息啦!" ) + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error(f"{router.prefix}/get_chat_and_call_month 调用错误", "WebUi", e=e) return Result.fail(f"发生了一点错误捏 {type(e)}: {e}") diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/dashboard/data_source.py b/zhenxun/builtin_plugins/web_ui/api/tabs/dashboard/data_source.py index 87011c93..77a4db83 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/dashboard/data_source.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/dashboard/data_source.py @@ -17,6 +17,7 @@ from zhenxun.utils.manager.priority_manager import PriorityLifecycle from zhenxun.utils.platform import PlatformUtils from ....base_model import BaseResultModel, QueryModel +from ....utils import webui_db_call from ..main.data_source import bot_live from .model import ( AllChatAndCallCount, @@ -77,14 +78,20 @@ class ApiDataSource: logger.warning("获取bot好友/群组信息失败...", "WebUi", e=e) bot_info.group_count = 0 bot_info.friend_count = 0 - bot_info.day_call = await Statistics.filter( - create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute), - bot_id=bot.self_id, - ).count() - bot_info.received_messages = await ChatHistory.filter( - bot_id=bot_info.self_id, - create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute), - ).count() + bot_info.day_call = await webui_db_call( + Statistics.filter( + create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute), + bot_id=bot.self_id, + ).count(), + "Dashboard.bot_day_call", + ) + bot_info.received_messages = await webui_db_call( + ChatHistory.filter( + bot_id=bot_info.self_id, + create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute), + ).count(), + "Dashboard.bot_received_messages", + ) bot_info.connect_time = bot_live.get(bot.self_id) or 0 if bot_info.connect_time: connect_date = datetime.fromtimestamp(CONNECT_TIME) @@ -117,17 +124,29 @@ class ApiDataSource: query = ChatHistory if bot_id: query = query.filter(bot_id=bot_id) - chat_all_count = await query.annotate().count() - chat_day_count = await query.filter( - create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute) - ).count() + chat_all_count = await webui_db_call( + query.annotate().count(), + "Dashboard.chat_all_count", + ) + chat_day_count = await webui_db_call( + query.filter( + create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute) + ).count(), + "Dashboard.chat_day_count", + ) query = Statistics if bot_id: query = query.filter(bot_id=bot_id) - call_all_count = await query.annotate().count() - call_day_count = await query.filter( - create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute) - ).count() + call_all_count = await webui_db_call( + query.annotate().count(), + "Dashboard.call_all_count", + ) + call_day_count = await webui_db_call( + query.filter( + create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute) + ).count(), + "Dashboard.call_day_count", + ) return QueryChatCallCount( chat_num=chat_all_count, chat_day=chat_day_count, @@ -151,31 +170,51 @@ class ApiDataSource: query = ChatHistory if bot_id: query = query.filter(bot_id=bot_id) - chat_week_count = await query.filter( - create_time__gte=now - timedelta(days=7, hours=now.hour, minutes=now.minute) - ).count() - chat_month_count = await query.filter( - create_time__gte=now - - timedelta(days=30, hours=now.hour, minutes=now.minute) - ).count() - chat_year_count = await query.filter( - create_time__gte=now - - timedelta(days=365, hours=now.hour, minutes=now.minute) - ).count() + chat_week_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=7, hours=now.hour, minutes=now.minute) + ).count(), + "Dashboard.chat_week_count", + ) + chat_month_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=30, hours=now.hour, minutes=now.minute) + ).count(), + "Dashboard.chat_month_count", + ) + chat_year_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=365, hours=now.hour, minutes=now.minute) + ).count(), + "Dashboard.chat_year_count", + ) query = Statistics if bot_id: query = query.filter(bot_id=bot_id) - call_week_count = await query.filter( - create_time__gte=now - timedelta(days=7, hours=now.hour, minutes=now.minute) - ).count() - call_month_count = await query.filter( - create_time__gte=now - - timedelta(days=30, hours=now.hour, minutes=now.minute) - ).count() - call_year_count = await query.filter( - create_time__gte=now - - timedelta(days=365, hours=now.hour, minutes=now.minute) - ).count() + call_week_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=7, hours=now.hour, minutes=now.minute) + ).count(), + "Dashboard.call_week_count", + ) + call_month_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=30, hours=now.hour, minutes=now.minute) + ).count(), + "Dashboard.call_month_count", + ) + call_year_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=365, hours=now.hour, minutes=now.minute) + ).count(), + "Dashboard.call_year_count", + ) return AllChatAndCallCount( chat_week=chat_week_count, chat_month=chat_month_count, @@ -202,17 +241,19 @@ class ApiDataSource: if bot_id: chat_query = chat_query.filter(bot_id=bot_id) call_query = call_query.filter(bot_id=bot_id) - chat_date_list = ( - await chat_query.filter(create_time__gte=filter_date) + chat_date_list = await webui_db_call( + chat_query.filter(create_time__gte=filter_date) .annotate(date=RawSQL("DATE(create_time)"), count=Count("id")) .group_by("date") - .values("date", "count") + .values("date", "count"), + "Dashboard.chat_month_series", ) - call_date_list = ( - await call_query.filter(create_time__gte=filter_date) + call_date_list = await webui_db_call( + call_query.filter(create_time__gte=filter_date) .annotate(date=RawSQL("DATE(create_time)"), count=Count("id")) .group_by("date") - .values("date", "count") + .values("date", "count"), + "Dashboard.call_month_series", ) date_list = [] chat_count_list = [] diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/main/__init__.py b/zhenxun/builtin_plugins/web_ui/api/tabs/main/__init__.py index 9756a470..6f36a43a 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/main/__init__.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/main/__init__.py @@ -16,7 +16,7 @@ from zhenxun.utils.platform import PlatformUtils from ....base_model import Result from ....config import QueryDateType -from ....utils import authentication, get_system_status +from ....utils import DB_BUSY_MESSAGE, authentication, get_system_status from .data_source import ApiDataSource from .model import ( ActiveGroup, @@ -82,6 +82,8 @@ async def _(bot_id: str | None = None) -> Result[list[BaseInfo]]: if not result: Result.warning_("无Bot连接...") return Result.ok(result, "拿到信息啦!") + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error(f"{router.prefix}/get_base_info 调用错误", "WebUi", e=e) return Result.fail(f"发生了一点错误捏 {type(e)}: {e}") @@ -97,6 +99,8 @@ async def _(bot_id: str | None = None) -> Result[list[BaseInfo]]: async def _(bot_id: str | None = None) -> Result[QueryCount]: try: return Result.ok(await ApiDataSource.get_all_chat_count(bot_id), "拿到信息啦!") + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error(f"{router.prefix}/get_all_chat_count 调用错误", "WebUi", e=e) return Result.fail(f"发生了一点错误捏 {type(e)}: {e}") @@ -112,6 +116,8 @@ async def _(bot_id: str | None = None) -> Result[QueryCount]: async def _(bot_id: str | None = None) -> Result[QueryCount]: try: return Result.ok(await ApiDataSource.get_all_call_count(bot_id), "拿到信息啦!") + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error(f"{router.prefix}/get_all_call_count 调用错误", "WebUi", e=e) return Result.fail(f"发生了一点错误捏 {type(e)}: {e}") @@ -188,6 +194,8 @@ async def _( return Result.ok( await ApiDataSource.get_active_group(date_type, bot_id), "拿到信息啦!" ) + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error(f"{router.prefix}/get_active_group 调用错误", "WebUi", e=e) return Result.fail(f"发生了一点错误捏 {type(e)}: {e}") @@ -207,6 +215,8 @@ async def _( return Result.ok( await ApiDataSource.get_hot_plugin(date_type, bot_id), "拿到信息啦!" ) + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error(f"{router.prefix}/get_hot_plugin 调用错误", "WebUi", e=e) return Result.fail(f"发生了一点错误捏 {type(e)}: {e}") @@ -294,9 +304,14 @@ async def system_logs_realtime(websocket: WebSocket, sleep: int = 5): await asyncio.wait_for(websocket.send_text(system_status.json()), timeout=5) try: await asyncio.wait_for(disconnect_event.wait(), timeout=max(sleep, 1)) - except TimeoutError: + except asyncio.TimeoutError: pass - except (WebSocketDisconnect, ConnectionClosedError, ConnectionClosedOK): + except ( + asyncio.CancelledError, + WebSocketDisconnect, + ConnectionClosedError, + ConnectionClosedOK, + ): pass finally: _SYSTEM_STATUS_CONNECTIONS.discard(websocket) diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/main/data_source.py b/zhenxun/builtin_plugins/web_ui/api/tabs/main/data_source.py index ec30a0ae..1c43b91c 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/main/data_source.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/main/data_source.py @@ -20,6 +20,7 @@ from zhenxun.utils.enum import PluginType from zhenxun.utils.platform import PlatformUtils from ....config import AVA_URL, GROUP_AVA_URL, QueryDateType +from ....utils import webui_db_call from .model import ( ActiveGroup, BaseInfo, @@ -105,10 +106,13 @@ class ApiDataSource: """ now = datetime.now() # 今日累计接收消息 - select_bot.received_messages = await ChatHistory.filter( - bot_id=select_bot.self_id, - create_time__gte=now - timedelta(hours=now.hour), - ).count() + select_bot.received_messages = await webui_db_call( + ChatHistory.filter( + bot_id=select_bot.self_id, + create_time__gte=now - timedelta(hours=now.hour), + ).count(), + "Main.received_messages", + ) # 群聊数量 try: select_bot.group_count = len( @@ -129,13 +133,15 @@ class ApiDataSource: connect_date = datetime.fromtimestamp(select_bot.connect_time) select_bot.connect_date = connect_date.strftime("%Y-%m-%d %H:%M:%S") select_bot.version = cls.__get_bot_version() - day_call = await Statistics.filter( - create_time__gte=now - timedelta(hours=now.hour) - ).count() + day_call = await webui_db_call( + Statistics.filter(create_time__gte=now - timedelta(hours=now.hour)).count(), + "Main.day_call", + ) select_bot.day_call = day_call - select_bot.connect_count = await BotConnectLog.filter( - bot_id=select_bot.self_id - ).count() + select_bot.connect_count = await webui_db_call( + BotConnectLog.filter(bot_id=select_bot.self_id).count(), + "Main.connect_count", + ) @classmethod async def get_base_info(cls, bot_id: str | None) -> list[BaseInfo] | None: @@ -177,21 +183,37 @@ class ApiDataSource: query = ChatHistory if bot_id: query = query.filter(bot_id=bot_id) - all_count = await query.annotate().count() - day_count = await query.filter( - create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute) - ).count() - week_count = await query.filter( - create_time__gte=now - timedelta(days=7, hours=now.hour, minutes=now.minute) - ).count() - month_count = await query.filter( - create_time__gte=now - - timedelta(days=30, hours=now.hour, minutes=now.minute) - ).count() - year_count = await query.filter( - create_time__gte=now - - timedelta(days=365, hours=now.hour, minutes=now.minute) - ).count() + all_count = await webui_db_call( + query.annotate().count(), + "Main.chat_all_count", + ) + day_count = await webui_db_call( + query.filter( + create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute) + ).count(), + "Main.chat_day_count", + ) + week_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=7, hours=now.hour, minutes=now.minute) + ).count(), + "Main.chat_week_count", + ) + month_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=30, hours=now.hour, minutes=now.minute) + ).count(), + "Main.chat_month_count", + ) + year_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=365, hours=now.hour, minutes=now.minute) + ).count(), + "Main.chat_year_count", + ) return QueryCount( num=all_count, day=day_count, @@ -214,21 +236,37 @@ class ApiDataSource: query = Statistics if bot_id: query = query.filter(bot_id=bot_id) - all_count = await query.annotate().count() - day_count = await query.filter( - create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute) - ).count() - week_count = await query.filter( - create_time__gte=now - timedelta(days=7, hours=now.hour, minutes=now.minute) - ).count() - month_count = await query.filter( - create_time__gte=now - - timedelta(days=30, hours=now.hour, minutes=now.minute) - ).count() - year_count = await query.filter( - create_time__gte=now - - timedelta(days=365, hours=now.hour, minutes=now.minute) - ).count() + all_count = await webui_db_call( + query.annotate().count(), + "Main.call_all_count", + ) + day_count = await webui_db_call( + query.filter( + create_time__gte=now - timedelta(hours=now.hour, minutes=now.minute) + ).count(), + "Main.call_day_count", + ) + week_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=7, hours=now.hour, minutes=now.minute) + ).count(), + "Main.call_week_count", + ) + month_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=30, hours=now.hour, minutes=now.minute) + ).count(), + "Main.call_month_count", + ) + year_count = await webui_db_call( + query.filter( + create_time__gte=now + - timedelta(days=365, hours=now.hour, minutes=now.minute) + ).count(), + "Main.call_year_count", + ) return QueryCount( num=all_count, day=day_count, @@ -296,19 +334,21 @@ class ApiDataSource: list[ActiveGroup]: 活跃群组列表 """ query = cls.__get_query(ChatHistory, date_type, bot_id) - data_list = ( - await query.annotate(count=Count("id")) + data_list = await webui_db_call( + query.annotate(count=Count("id")) .filter(group_id__not_isnull=True) .group_by("group_id") .order_by("-count") .limit(5) - .values_list("group_id", "count") + .values_list("group_id", "count"), + "Main.active_group", ) id2name = {} if data_list: - if info_list := await GroupConsole.filter( - group_id__in=[x[0] for x in data_list] - ).all(): + if info_list := await webui_db_call( + GroupConsole.filter(group_id__in=[x[0] for x in data_list]).all(), + "Main.active_group_names", + ): for group_info in info_list: id2name[group_info.group_id] = group_info.group_name active_group_list = [ @@ -341,19 +381,23 @@ class ApiDataSource: list[HotPlugin]: 热门插件列表 """ query = cls.__get_query(Statistics, date_type, bot_id) - data_list = ( - await query.annotate(count=Count("id")) + data_list = await webui_db_call( + query.annotate(count=Count("id")) .group_by("plugin_name") .order_by("-count") .limit(5) - .values_list("plugin_name", "count") + .values_list("plugin_name", "count"), + "Main.hot_plugin", ) hot_plugin_list = [] module_list = [x[0] for x in data_list] - plugins = await PluginInfo.get_plugins( - load_status=None, - filter_parent=False, - module__in=module_list, + plugins = await webui_db_call( + PluginInfo.get_plugins( + load_status=None, + filter_parent=False, + module__in=module_list, + ), + "Main.hot_plugin_names", ) module2name = {p.module: p.name for p in plugins} for data in data_list: diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py b/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py index 94a23d05..b01dc8f6 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/manage/__init__.py @@ -12,7 +12,7 @@ from zhenxun.utils.platform import PlatformUtils from ....base_model import Result from ....config import AVA_URL, GROUP_AVA_URL -from ....utils import authentication +from ....utils import DB_BUSY_MESSAGE, authentication from .data_source import ApiDataSource from .model import ( ClearRequest, @@ -294,6 +294,8 @@ async def _(bot_id: str, user_id: str) -> Result[UserDetail]: if result else Result.warning_("未找到该好友...") ) + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except (ValueError, KeyError): return Result.warning_("指定Bot未连接...") except Exception as e: @@ -311,6 +313,8 @@ async def _(bot_id: str, user_id: str) -> Result[UserDetail]: async def _(group_id: str) -> Result[GroupDetail]: try: return Result.ok(await ApiDataSource.get_group_detail(group_id), "拿到信息啦!") + except TimeoutError: + return Result.fail(DB_BUSY_MESSAGE) except Exception as e: logger.error(f"{router.prefix}/get_group_detail 调用错误", "WebUi", e=e) return Result.fail(f"发生了一点错误捏 {type(e)}: {e}") diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/manage/data_source.py b/zhenxun/builtin_plugins/web_ui/api/tabs/manage/data_source.py index 0b18db6b..626c6c07 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/manage/data_source.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/manage/data_source.py @@ -13,6 +13,7 @@ from zhenxun.utils.enum import RequestType from zhenxun.utils.platform import PlatformUtils from ....config import AVA_URL, GROUP_AVA_URL +from ....utils import webui_db_call from .model import ( FriendRequestResult, GroupDetail, @@ -110,20 +111,24 @@ class ApiDataSource: fd = [x for x in friend_list if x.user_id == user_id] if not fd: return None - like_plugin_list = ( - await Statistics.filter(user_id=user_id) + like_plugin_list = await webui_db_call( + Statistics.filter(user_id=user_id) .annotate(count=Count("id")) .group_by("plugin_name") .order_by("-count") .limit(5) - .values_list("plugin_name", "count") + .values_list("plugin_name", "count"), + "Manage.friend_like_plugin", ) like_plugin = {} module_list = [x[0] for x in like_plugin_list] - plugins = await PluginInfo.get_plugins( - load_status=None, - filter_parent=False, - module__in=module_list, + plugins = await webui_db_call( + PluginInfo.get_plugins( + load_status=None, + filter_parent=False, + module__in=module_list, + ), + "Manage.friend_like_plugin_names", ) module2name = {p.module: p.name for p in plugins} for data in like_plugin_list: @@ -136,8 +141,14 @@ class ApiDataSource: nickname=user.user_name, remark="", is_ban=await BanConsole.is_ban(user_id), - chat_count=await ChatHistory.filter(user_id=user_id).count(), - call_count=await Statistics.filter(user_id=user_id).count(), + chat_count=await webui_db_call( + ChatHistory.filter(user_id=user_id).count(), + "Manage.friend_chat_count", + ), + call_count=await webui_db_call( + Statistics.filter(user_id=user_id).count(), + "Manage.friend_call_count", + ), like_plugin=like_plugin, ) @@ -151,16 +162,20 @@ class ApiDataSource: 返回: dict[str, int]: 插件与调用次数 """ - like_plugin_list = ( - await Statistics.filter(group_id=group_id) + like_plugin_list = await webui_db_call( + Statistics.filter(group_id=group_id) .annotate(count=Count("id")) .group_by("plugin_name") .order_by("-count") .limit(5) - .values_list("plugin_name", "count") + .values_list("plugin_name", "count"), + "Manage.group_like_plugin", ) like_plugin = {} - plugins = await PluginInfo.get_plugins() + plugins = await webui_db_call( + PluginInfo.get_plugins(), + "Manage.group_like_plugin_names", + ) module2name = {p.module: p.name for p in plugins} for data in like_plugin_list: name = module2name.get(data[0]) or data[0] @@ -268,8 +283,14 @@ class ApiDataSource: name=group.group_name, member_count=group.member_count, max_member_count=group.max_member_count, - chat_count=await ChatHistory.filter(group_id=group_id).count(), - call_count=await Statistics.filter(group_id=group_id).count(), + chat_count=await webui_db_call( + ChatHistory.filter(group_id=group_id).count(), + "Manage.group_chat_count", + ), + call_count=await webui_db_call( + Statistics.filter(group_id=group_id).count(), + "Manage.group_call_count", + ), like_plugin=like_plugin, level=group.level, status=group.status, diff --git a/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/store.py b/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/store.py index 44f67b86..9e70dc79 100644 --- a/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/store.py +++ b/zhenxun/builtin_plugins/web_ui/api/tabs/plugin_manage/store.py @@ -49,6 +49,7 @@ async def _(param: PluginIr) -> Result: from zhenxun.builtin_plugins.plugin_store import StoreManager result = await StoreManager.add_plugin(str(param.id)) # type: ignore + logger.info(result.replace("\n", ";"), "插件商店") return Result.ok(info=result) except Exception as e: return Result.fail(f"安装插件失败: {type(e)}: {e}") diff --git a/zhenxun/builtin_plugins/web_ui/utils.py b/zhenxun/builtin_plugins/web_ui/utils.py index e2884af2..0f38ca9b 100644 --- a/zhenxun/builtin_plugins/web_ui/utils.py +++ b/zhenxun/builtin_plugins/web_ui/utils.py @@ -12,11 +12,15 @@ import ujson as json from zhenxun.configs.config import Config from zhenxun.configs.path_config import DATA_PATH +from zhenxun.services.db_context import with_db_timeout +from zhenxun.services.message_load import is_db_unhealthy from .base_model import SystemFolderSize, SystemStatus, User ALGORITHM = "HS256" ACCESS_TOKEN_EXPIRE_MINUTES = 30 +DB_BUSY_MESSAGE = "数据库繁忙,请稍后再试" +WEBUI_DB_TIMEOUT = 3.0 oauth2_scheme = OAuth2PasswordBearer(tokenUrl="api/login") @@ -28,6 +32,20 @@ if token_file.exists(): token_data = json.load(open(token_file, encoding="utf8")) +async def webui_db_call(coro, operation: str): + if is_db_unhealthy(): + close = getattr(coro, "close", None) + if callable(close): + close() + raise TimeoutError(DB_BUSY_MESSAGE) + return await with_db_timeout( + coro, + timeout=WEBUI_DB_TIMEOUT, + operation=operation, + source="web_ui", + ) + + def validate_path(path_str: str | None) -> tuple[Path | None, str | None]: """验证路径是否安全 diff --git a/zhenxun/models/ban_console.py b/zhenxun/models/ban_console.py index 12302096..504b7ae2 100644 --- a/zhenxun/models/ban_console.py +++ b/zhenxun/models/ban_console.py @@ -5,12 +5,10 @@ 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 -from zhenxun.services.data_access import DataAccess from zhenxun.services.db_context import Model from zhenxun.services.log import logger -from zhenxun.utils.enum import CacheType, DbLockType +from zhenxun.utils.enum import DbLockType from zhenxun.utils.exception import UserAndGroupIsNone @@ -38,28 +36,8 @@ class BanConsole(Model): unique_together = ("user_id", "group_id") indexes = [("user_id",), ("group_id",)] # noqa: RUF012 - cache_type = CacheType.BAN - """缓存类型""" - cache_key_field = ("user_id", "group_id") - """缓存键字段""" enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT] """开启锁""" - _cache_checked: ClassVar[bool] = False - - @classmethod - def _ensure_cache_registered(cls): - """兜底注册 BAN 缓存,避免启动时序导致的未注册问题。""" - if cls._cache_checked: - return - try: - CacheRoot.get_model(CacheType.BAN) - except CacheException: - CacheRegistry.register( - CacheType.BAN, - cls, - key_format="{user_id}_{group_id}", - ) - cls._cache_checked = True @classmethod async def create(cls, *args, **kwargs) -> Self: @@ -87,15 +65,13 @@ class BanConsole(Model): 返回: Self | None: Self """ - cls._ensure_cache_registered() if not user_id and not group_id: raise UserAndGroupIsNone() if user_id: - dao = DataAccess(cls) return ( - await dao.safe_get_or_none(user_id=user_id, group_id=group_id) + await cls.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 await cls.safe_get_or_none(user_id=user_id, group_id__isnull=True) ) else: return await cls.safe_get_or_none( @@ -177,7 +153,6 @@ class BanConsole(Model): ) if not user_id and not group_id: raise UserAndGroupIsNone() - cls._ensure_cache_registered() target, _ = await cls.update_or_create( user_id=user_id, group_id=group_id, diff --git a/zhenxun/models/group_console.py b/zhenxun/models/group_console.py index 19001f13..2d3e2f89 100644 --- a/zhenxun/models/group_console.py +++ b/zhenxun/models/group_console.py @@ -9,10 +9,9 @@ from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.task_info import TaskInfo from zhenxun.services.cache import CacheRoot from zhenxun.services.cache.runtime_cache import GroupMemoryCache -from zhenxun.services.data_access import DataAccess from zhenxun.services.db_context import Model from zhenxun.services.db_context.schema_ops import AlterColumnType, CreateIndex -from zhenxun.utils.enum import CacheType, DbLockType, PluginType +from zhenxun.utils.enum import DbLockType, PluginType if TYPE_CHECKING: from zhenxun.services.cache.runtime_cache import GroupSnapshot @@ -98,10 +97,6 @@ class GroupConsole(Model): ("group_id",) ] - cache_type = CacheType.GROUPS - """缓存类型""" - cache_key_field = ("group_id", "channel_id") - """缓存键字段""" enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT] """开启锁""" _root_group_locks: ClassVar[dict[str, asyncio.Lock]] = {} @@ -331,9 +326,11 @@ class GroupConsole(Model): if update_fields: await keep.save(update_fields=update_fields) - for group in groups: - if group.id != keep.id: - await group.delete() + duplicate_ids = [ + group.id for group in groups if group.id and group.id != keep.id + ] + if duplicate_ids: + await cls.filter(id__in=duplicate_ids).delete() await GroupMemoryCache.upsert_from_model(keep) return keep @@ -417,14 +414,13 @@ class GroupConsole(Model): clean_duplicates: bool = True, ) -> Self | None: """获取群组(数据库)""" - dao = DataAccess(cls) if channel_id: - return await dao.safe_get_or_none( + return await cls.safe_get_or_none( group_id=group_id, channel_id=channel_id, clean_duplicates=clean_duplicates, ) - return await dao.safe_get_or_none( + return await cls.safe_get_or_none( group_id=group_id, channel_id__isnull=True, clean_duplicates=clean_duplicates, diff --git a/zhenxun/models/group_plugin_setting.py b/zhenxun/models/group_plugin_setting.py index e6005707..e05e5087 100644 --- a/zhenxun/models/group_plugin_setting.py +++ b/zhenxun/models/group_plugin_setting.py @@ -1,7 +1,6 @@ from tortoise import fields from zhenxun.services.db_context import Model -from zhenxun.utils.enum import CacheType class GroupPluginSetting(Model): @@ -18,11 +17,6 @@ class GroupPluginSetting(Model): updated_at = fields.DatetimeField(auto_now=True, description="最后更新时间") """最后更新时间""" - cache_type = CacheType.GROUP_PLUGIN_SETTINGS - """缓存类型""" - cache_key_field = ("group_id", "plugin_name") - """缓存键字段""" - class Meta: # pyright: ignore [reportIncompatibleVariableOverride] table = "group_plugin_settings" table_description = "插件分群通用配置表" diff --git a/zhenxun/services/buffered_writers.py b/zhenxun/services/buffered_writers.py index b2dd7842..2f646531 100644 --- a/zhenxun/services/buffered_writers.py +++ b/zhenxun/services/buffered_writers.py @@ -1,61 +1,49 @@ from __future__ import annotations -import asyncio -from collections import deque -import contextlib -import time - from zhenxun.models.user_gold_log import UserGoldLog -from zhenxun.services.log import logger +from zhenxun.services.low_priority_writer import ( + LowPriorityWriterConfig, + append_low_priority_record, + flush_low_priority_writer, + register_low_priority_writer, +) from zhenxun.utils.enum import GoldHandle -from zhenxun.utils.manager.priority_manager import PriorityLifecycle LOG_COMMAND = "BufferedWriters" +_WRITER_NAME = "user_gold_log" _USER_GOLD_LOG_BUFFER_MAX_RETAIN = 10_000 _USER_GOLD_LOG_FLUSH_TRIGGER_SIZE = 128 _USER_GOLD_LOG_FLUSH_BATCH_SIZE = 500 _USER_GOLD_LOG_FLUSH_INTERVAL_SECONDS = 60.0 -_USER_GOLD_LOG_DROP_LOG_INTERVAL_SECONDS = 10.0 - -_user_gold_log_buffer: deque[UserGoldLog] = deque() -_user_gold_log_buffer_lock = asyncio.Lock() -_user_gold_log_flush_lock = asyncio.Lock() -_user_gold_log_flush_task: asyncio.Task[None] | None = None -_user_gold_log_dropped = 0 -_user_gold_log_last_drop_log_at = 0.0 -def _ensure_user_gold_log_flush_task() -> None: - global _user_gold_log_flush_task - if _user_gold_log_flush_task is not None and not _user_gold_log_flush_task.done(): - return - _user_gold_log_flush_task = asyncio.create_task(_user_gold_log_flush_loop()) +async def _write_user_gold_log_batch( + batch: list[UserGoldLog], + reason: str, +) -> None: + from zhenxun.services.db_context import with_db_timeout - -def _record_user_gold_log_drop() -> None: - global _user_gold_log_dropped, _user_gold_log_last_drop_log_at - _user_gold_log_dropped += 1 - now = time.monotonic() - if now - _user_gold_log_last_drop_log_at < _USER_GOLD_LOG_DROP_LOG_INTERVAL_SECONDS: - return - _user_gold_log_last_drop_log_at = now - logger.warning( - "user_gold_log buffer full, dropped " - f"{_user_gold_log_dropped} records, backlog={len(_user_gold_log_buffer)}", - LOG_COMMAND, + await with_db_timeout( + UserGoldLog.bulk_create(batch, _USER_GOLD_LOG_FLUSH_BATCH_SIZE), + timeout=5.0, + operation=f"UserGoldLog.bulk_create[{len(batch)}]", + source=f"user_gold_log:{reason}", ) -async def _user_gold_log_flush_loop() -> None: - while True: - await asyncio.sleep(_USER_GOLD_LOG_FLUSH_INTERVAL_SECONDS) - try: - await flush_user_gold_log_buffer("定时") - except asyncio.CancelledError: - raise - except Exception as exc: - logger.warning("定时批量写入金币流水失败", LOG_COMMAND, e=exc) +register_config = LowPriorityWriterConfig( + name=_WRITER_NAME, + write_batch=_write_user_gold_log_batch, + batch_size=_USER_GOLD_LOG_FLUSH_BATCH_SIZE, + trigger_size=_USER_GOLD_LOG_FLUSH_TRIGGER_SIZE, + max_retain=_USER_GOLD_LOG_BUFFER_MAX_RETAIN, + flush_interval_seconds=_USER_GOLD_LOG_FLUSH_INTERVAL_SECONDS, + max_items_per_cycle=_USER_GOLD_LOG_FLUSH_BATCH_SIZE, + backoff_base_seconds=30.0, + backoff_max_seconds=600.0, + log_command=LOG_COMMAND, +) async def append_user_gold_log( @@ -64,65 +52,16 @@ async def append_user_gold_log( handle: GoldHandle, source: str | None = None, ) -> None: - _ensure_user_gold_log_flush_task() record = UserGoldLog(user_id=user_id, gold=gold, handle=handle, source=source) - async with _user_gold_log_buffer_lock: - if len(_user_gold_log_buffer) >= _USER_GOLD_LOG_BUFFER_MAX_RETAIN: - _user_gold_log_buffer.popleft() - _record_user_gold_log_drop() - _user_gold_log_buffer.append(record) - should_flush = ( - len(_user_gold_log_buffer) >= _USER_GOLD_LOG_FLUSH_TRIGGER_SIZE - and not _user_gold_log_flush_lock.locked() - ) - if should_flush: - await flush_user_gold_log_buffer("缓冲区触发") + await append_low_priority_record(_WRITER_NAME, record) async def flush_user_gold_log_buffer(reason: str) -> int: - async with _user_gold_log_flush_lock: - written = 0 - while True: - batch: list[UserGoldLog] = [] - async with _user_gold_log_buffer_lock: - if not _user_gold_log_buffer: - break - while ( - _user_gold_log_buffer - and len(batch) < _USER_GOLD_LOG_FLUSH_BATCH_SIZE - ): - batch.append(_user_gold_log_buffer.popleft()) - if not batch: - break - try: - await UserGoldLog.bulk_create(batch, _USER_GOLD_LOG_FLUSH_BATCH_SIZE) - except Exception as exc: - async with _user_gold_log_buffer_lock: - retain_count = max( - _USER_GOLD_LOG_BUFFER_MAX_RETAIN - len(_user_gold_log_buffer), - 0, - ) - for record in reversed(batch[-retain_count:]): - _user_gold_log_buffer.appendleft(record) - logger.error(f"{reason}批量写入金币流水失败", LOG_COMMAND, e=exc) - return written - written += len(batch) - if written: - logger.debug(f"{reason}批量写入金币流水 {written} 条", LOG_COMMAND) - return written + return await flush_low_priority_writer(_WRITER_NAME, reason, force=True) async def stop_user_gold_log_buffer() -> int: - global _user_gold_log_flush_task - task = _user_gold_log_flush_task - _user_gold_log_flush_task = None - if task is not None: - task.cancel() - with contextlib.suppress(BaseException): - await task return await flush_user_gold_log_buffer("关闭") -@PriorityLifecycle.on_shutdown(priority=90) -async def _flush_user_gold_log_buffer_on_shutdown() -> None: - await stop_user_gold_log_buffer() +register_low_priority_writer(register_config) diff --git a/zhenxun/services/cache/__init__.py b/zhenxun/services/cache/__init__.py index 44d1b9c8..7e83871a 100644 --- a/zhenxun/services/cache/__init__.py +++ b/zhenxun/services/cache/__init__.py @@ -52,7 +52,6 @@ from typing import Any, ClassVar, Generic, TypeVar, cast, get_type_hints from typing_extensions import Self from aiocache import Cache as AioCache -from aiocache import SimpleMemoryCache from aiocache.base import BaseCache from aiocache.serializers import JsonSerializer import nonebot @@ -109,6 +108,10 @@ driver = nonebot.get_driver() cache_config = nonebot.get_plugin_config(Config) +def _redis_cache_enabled() -> bool: + return cache_config.cache_mode == CacheMode.REDIS and bool(cache_config.redis_host) + + class CacheException(Exception): """缓存相关异常""" @@ -190,15 +193,10 @@ class CacheManager: def cache_backend(self) -> BaseCache | AioCache: """获取缓存后端""" if self._cache_backend is None: - ttl = cache_config.redis_expire - if cache_config.cache_mode == CacheMode.NONE: - ttl = 0 - logger.info("缓存功能已禁用,使用非持久化内存缓存", LOG_COMMAND) - elif cache_config.cache_mode == CacheMode.REDIS and cache_config.redis_host: + if _redis_cache_enabled(): try: from aiocache import RedisCache - # 使用Redis缓存 self._cache_backend = RedisCache( serializer=JsonSerializer(), namespace=CACHE_KEY_PREFIX, @@ -215,19 +213,11 @@ class CacheManager: return self._cache_backend except ImportError as e: logger.error( - "导入aiocache[redis]失败,将默认使用内存缓存...", + "导入aiocache[redis]失败,CacheRoot 模型缓存已禁用", LOG_COMMAND, e=e, ) - else: - logger.info("使用内存缓存", LOG_COMMAND) - # 默认使用内存缓存 - self._cache_backend = SimpleMemoryCache( - serializer=JsonSerializer(), - namespace=CACHE_KEY_PREFIX, - timeout=30, - ttl=ttl, - ) + raise CacheException("CacheRoot 后端未启用") return self._cache_backend async def invalidate_cache( @@ -815,11 +805,9 @@ class Cache(Generic[T]): @driver.on_startup async def _(): - CacheRoot.enabled = cache_config.cache_mode != CacheMode.NONE + CacheRoot.enabled = _redis_cache_enabled() if CacheRoot.enabled: - logger.info("缓存系统已启用", LOG_COMMAND) - else: - logger.info("缓存系统已禁用", LOG_COMMAND) + logger.info("CacheRoot Redis 模型缓存已启用", LOG_COMMAND) @driver.on_shutdown diff --git a/zhenxun/services/cache/cache_containers.py b/zhenxun/services/cache/cache_containers.py index 47e393b9..14d5b3fc 100644 --- a/zhenxun/services/cache/cache_containers.py +++ b/zhenxun/services/cache/cache_containers.py @@ -248,285 +248,3 @@ class CacheDict(Generic[T]): # 清理过期的键 self._clean_expired() return f"CacheDict({self.name}, {len(self._data)} items)" - - -class CacheList(Generic[T]): - """缓存列表类,提供类似普通列表的接口,数据只存储在内存中""" - - _instances: weakref.WeakSet = weakref.WeakSet() - - def __init__(self, name: str, expire: int = 0, max_items: int | None = None): - """初始化缓存列表 - - 参数: - name: 列表名称 - expire: 过期时间(秒),默认为0表示永不过期 - max_items: 最大缓存项数,None 使用统一默认值,0 表示不限制 - """ - self.name = name.upper() - self.expire = expire - self.max_items = DEFAULT_CACHE_MAX_ITEMS if max_items is None else max_items - self._data: list[CacheData[T]] = [] - self._expire_time = 0 - self.__class__._instances.add(self) - - # 如果设置了过期时间,计算整个列表的过期时间 - if self.expire > 0: - self._expire_time = time.time() + self.expire - - def __getitem__(self, index: int) -> T: - """获取列表项 - - 参数: - index: 列表索引 - - 返回: - T: 列表值 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - raise IndexError(f"列表索引 {index} 超出范围") - - if 0 <= index < len(self._data): - return self._data[index].value - raise IndexError(f"列表索引 {index} 超出范围") - - def __setitem__(self, index: int, value: T): - """设置列表项 - - 参数: - index: 列表索引 - value: 列表值 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - - # 确保索引有效 - while len(self._data) <= index: - raise IndexError(f"列表索引 {index} 超出范围") - self._data[index] = CacheData(value=value) - - # 更新过期时间 - self._update_expire_time() - - def __delitem__(self, index: int): - """删除列表项 - - 参数: - index: 列表索引 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - raise IndexError(f"列表索引 {index} 超出范围") - - if not 0 <= index < len(self._data): - raise IndexError(f"列表索引 {index} 超出范围") - del self._data[index] - # 更新过期时间 - self._update_expire_time() - - def __len__(self) -> int: - """获取列表长度 - - 返回: - int: 列表长度 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - return len(self._data) - - def append(self, value: T): - """添加列表项 - - 参数: - value: 列表值 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - - self._data.append(CacheData(value=value)) - self._enforce_limit() - - # 更新过期时间 - self._update_expire_time() - - def extend(self, values: list[T]): - """扩展列表 - - 参数: - values: 要添加的值列表 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - - self._data.extend([CacheData(value=v) for v in values]) - self._enforce_limit() - - # 更新过期时间 - self._update_expire_time() - - def insert(self, index: int, value: T): - """插入列表项 - - 参数: - index: 插入位置 - value: 列表值 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - - self._data.insert(index, CacheData(value=value)) - self._enforce_limit() - - # 更新过期时间 - self._update_expire_time() - - def pop(self, index: int = -1) -> T: - """删除并返回列表项 - - 参数: - index: 列表索引,默认为最后一项 - - 返回: - Any: 列表值 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - raise IndexError("从空列表中弹出") - - if not self._data: - raise IndexError("从空列表中弹出") - - item = self._data.pop(index) - - # 更新过期时间 - self._update_expire_time() - - return item.value - - def remove(self, value: T): - """删除第一个匹配的列表项 - - 参数: - value: 要删除的值 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - raise ValueError(f"{value} 不在列表中") - - # 查找匹配的项 - for i, item in enumerate(self._data): - if item.value == value: - del self._data[i] - # 更新过期时间 - self._update_expire_time() - return - - raise ValueError(f"{value} 不在列表中") - - def clear(self) -> None: - """清空列表""" - self._data.clear() - # 重置过期时间 - self._update_expire_time() - - def stats(self) -> dict[str, int]: - """返回当前缓存条目统计。""" - if self._is_expired(): - self.clear() - return {"items": len(self._data), "max_items": self.max_items} - - @classmethod - def stats_all(cls) -> dict[str, dict[str, int]]: - """返回所有 CacheList 实例的条目统计。""" - result: dict[str, dict[str, int]] = {} - for cache in list(cls._instances): - stats = cache.stats() - if stats["items"]: - result[cache.name] = stats - return result - - @classmethod - def clear_all(cls) -> dict[str, int]: - result: dict[str, int] = {} - for cache in list(cls._instances): - size = len(cache._data) - if size: - cache.clear() - result[cache.name] = result.get(cache.name, 0) + size - return result - - def index(self, value: T, start: int = 0, end: int | None = None) -> int: - """查找值的索引 - - 参数: - value: 要查找的值 - start: 起始索引 - end: 结束索引 - - 返回: - int: 索引位置 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - raise ValueError(f"{value} 不在列表中") - - end = end if end is not None else len(self._data) - - for i in range(start, min(end, len(self._data))): - if self._data[i].value == value: - return i - - raise ValueError(f"{value} 不在列表中") - - def count(self, value: T) -> int: - """计算值出现的次数 - - 参数: - value: 要计数的值 - - 返回: - int: 出现次数 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - return 0 - - # sourcery skip: simplify-constant-sum - return sum(1 for item in self._data if item.value == value) - - def _is_expired(self) -> bool: - """检查整个列表是否过期""" - return self._expire_time > 0 and self._expire_time < time.time() - - def _update_expire_time(self): - """更新过期时间""" - self._expire_time = time.time() + self.expire if self.expire > 0 else 0 - - def _enforce_limit(self) -> None: - if self.max_items <= 0: - return - if len(self._data) > self.max_items: - del self._data[: len(self._data) - self.max_items] - - def __str__(self) -> str: - """字符串表示 - - 返回: - str: 字符串表示 - """ - # 检查整个列表是否过期 - if self._is_expired(): - self.clear() - return f"CacheList({self.name}, {len(self._data)} items)" diff --git a/zhenxun/services/cache/runtime_cache.py b/zhenxun/services/cache/runtime_cache.py index d43cc1b7..3e9e3d4f 100644 --- a/zhenxun/services/cache/runtime_cache.py +++ b/zhenxun/services/cache/runtime_cache.py @@ -11,6 +11,7 @@ import uuid from zhenxun.services.cache.config import CacheMode from zhenxun.services.log import logger +from zhenxun.services.message_load import is_db_unhealthy from zhenxun.utils.enum import ( BlockType, LimitCheckType, @@ -57,6 +58,7 @@ 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 +RUNTIME_CACHE_DB_TIMEOUT_SECONDS = 3.0 INSTANCE_ID = uuid.uuid4().hex @@ -749,6 +751,9 @@ class RuntimeCacheMutation: @classmethod async def ensure_loaded(cls, cache_cls: type, label: str) -> None: + if is_db_unhealthy(): + cls.mark_error(cache_cls, RuntimeError("database unhealthy")) + return if getattr(cache_cls, "_loaded", False): return now = time.monotonic() @@ -756,6 +761,9 @@ class RuntimeCacheMutation: return lock = cls._load_locks.setdefault(label, asyncio.Lock()) async with lock: + if is_db_unhealthy(): + cls.mark_error(cache_cls, RuntimeError("database unhealthy")) + return if getattr(cache_cls, "_loaded", False): return now = time.monotonic() @@ -768,7 +776,27 @@ class RuntimeCacheMutation: cls._retry_after[label] = ( time.monotonic() + RUNTIME_CACHE_LOAD_RETRY_SECONDS ) - raise + return + + @staticmethod + async def read_db(cache_cls: type, coro, *, operation: str): + from zhenxun.services.db_context import with_db_timeout + + if is_db_unhealthy(): + RuntimeCacheMutation.mark_error( + cache_cls, RuntimeError("database unhealthy") + ) + return None + try: + return await with_db_timeout( + coro, + timeout=RUNTIME_CACHE_DB_TIMEOUT_SECONDS, + operation=operation, + source="runtime_cache", + ) + except Exception as exc: + RuntimeCacheMutation.mark_error(cache_cls, exc) + return None @staticmethod def mark_refreshed(cache_cls: type) -> None: @@ -834,7 +862,13 @@ class PluginInfoMemoryCache: from zhenxun.models.plugin_info import PluginInfo async with cls._lock: - plugins = await PluginInfo.all() + plugins = await RuntimeCacheMutation.read_db( + cls, + PluginInfo.all(), + operation="PluginInfoMemoryCache.refresh", + ) + if plugins is None: + return by_module: dict[str, PluginInfoSnapshot] = {} by_module_path: dict[str, PluginInfoSnapshot] = {} for plugin in plugins: @@ -1016,7 +1050,13 @@ class BotMemoryCache: from zhenxun.models.bot_console import BotConsole async with cls._lock: - records = await BotConsole.all() + records = await RuntimeCacheMutation.read_db( + cls, + BotConsole.all(), + operation="BotMemoryCache.refresh", + ) + if records is None: + return cls._by_id = {str(r.bot_id): BotSnapshot.from_model(r) for r in records} RuntimeCacheMutation.clear_negative_all(cls) RuntimeCacheMutation.mark_refreshed(cls) @@ -1201,7 +1241,13 @@ class GroupMemoryCache: from zhenxun.models.group_console import GroupConsole async with cls._lock: - records = await GroupConsole.all() + records = await RuntimeCacheMutation.read_db( + cls, + GroupConsole.all(), + operation="GroupMemoryCache.refresh", + ) + if records is None: + return by_key: dict[tuple[str, str], GroupSnapshot] = {} for record in records: entry = GroupSnapshot.from_model(record) @@ -1374,7 +1420,13 @@ class LevelUserMemoryCache: from zhenxun.models.level_user import LevelUser async with cls._lock: - records = await LevelUser.all() + records = await RuntimeCacheMutation.read_db( + cls, + LevelUser.all(), + operation="LevelUserMemoryCache.refresh", + ) + if records is None: + return by_key: dict[tuple[str, str], LevelUserSnapshot] = {} by_user_max: dict[str, int] = {} for record in records: @@ -1607,7 +1659,13 @@ class TaskInfoMemoryCache: from zhenxun.models.task_info import TaskInfo async with cls._lock: - records = await TaskInfo.all() + records = await RuntimeCacheMutation.read_db( + cls, + TaskInfo.all(), + operation="TaskInfoMemoryCache.refresh", + ) + if records is None: + return by_module: dict[str, TaskInfoSnapshot] = {} by_name: dict[str, TaskInfoSnapshot] = {} for record in records: @@ -1790,7 +1848,13 @@ class PluginLimitMemoryCache: from zhenxun.models.plugin_limit import PluginLimit async with cls._lock: - records = await PluginLimit.filter(status=True).all() + records = await RuntimeCacheMutation.read_db( + cls, + PluginLimit.filter(status=True).all(), + operation="PluginLimitMemoryCache.refresh", + ) + if records is None: + return by_id: dict[int, PluginLimitSnapshot] = {} by_module: dict[str, list[PluginLimitSnapshot]] = {} for record in records: @@ -2015,7 +2079,13 @@ class BanMemoryCache: async with cls._lock: now_ts = time.time() - records = await BanConsole.all() + records = await RuntimeCacheMutation.read_db( + cls, + BanConsole.all(), + operation="BanMemoryCache.refresh", + ) + if records is None: + return by_user: dict[str, BanEntry] = {} by_group: dict[str, BanEntry] = {} by_user_group: dict[tuple[str, str], BanEntry] = {} diff --git a/zhenxun/services/data_access.py b/zhenxun/services/data_access.py index 22b3ad7d..ccef2a66 100644 --- a/zhenxun/services/data_access.py +++ b/zhenxun/services/data_access.py @@ -48,6 +48,7 @@ class DataAccess(Generic[T]): # 添加缓存统计信息 _cache_stats: ClassVar[dict] = {} + _ENABLE_CACHE_STATS: ClassVar[bool] = False # 空结果标记 _NULL_RESULT = "__NULL_RESULT_PLACEHOLDER__" # 默认空结果缓存时间(秒)- 设置为5分钟,避免频繁查询数据库 @@ -87,12 +88,10 @@ class DataAccess(Generic[T]): self.key_field = getattr(model_cls, "cache_key_field", key_field) self.cache_type = getattr(model_cls, "cache_type", cache_type) - if not self.cache_type: - raise ValueError("缓存类型不能为空") - self.cache = Cache(self.cache_type) + self.cache = Cache(self.cache_type) if self.cache_type else None # 初始化缓存统计 - if self.cache_type not in self._cache_stats: + if self.cache_type and self.cache_type not in self._cache_stats: self._cache_stats[self.cache_type] = { "hits": 0, # 缓存命中次数 "misses": 0, # 缓存未命中次数 @@ -137,6 +136,12 @@ class DataAccess(Generic[T]): stats["null_sets"] = 0 stats["deletes"] = 0 + @classmethod + def _bump_cache_stat(cls, cache_type: str | None, key: str) -> None: + if not cls._ENABLE_CACHE_STATS or not cache_type: + return + cls._cache_stats[cache_type][key] += 1 + def _build_cache_key_from_kwargs(self, **kwargs) -> str | None: """从关键字参数构建缓存键 @@ -191,7 +196,7 @@ class DataAccess(Generic[T]): # 如果成功构建缓存键,尝试从缓存获取 if cache_key is not None: - data = await self.cache.get(cache_key) + data = await self.cache.get(cache_key) if self.cache else None logger.debug( f"{self.model_cls.__name__} key: {cache_key}" f" 从缓存获取到的数据 {type(data)}: {data}" @@ -199,7 +204,7 @@ class DataAccess(Generic[T]): if data == self._NULL_RESULT: # 空结果缓存命中 - self._cache_stats[self.cache_type]["null_hits"] += 1 + self._bump_cache_stat(self.cache_type, "null_hits") logger.debug( f"{self.model_cls.__name__} 从缓存获取到空结果: {cache_key}" ) @@ -211,14 +216,14 @@ class DataAccess(Generic[T]): return None elif data: # 缓存命中 - self._cache_stats[self.cache_type]["hits"] += 1 + self._bump_cache_stat(self.cache_type, "hits") logger.debug( f"{self.model_cls.__name__} 从缓存获取数据成功: {cache_key}" ) return cast(T, data) else: # 缓存未命中 - self._cache_stats[self.cache_type]["misses"] += 1 + self._bump_cache_stat(self.cache_type, "misses") logger.debug(f"{self.model_cls.__name__} 缓存未命中: {cache_key}") except Exception as e: logger.error(f"{self.model_cls.__name__} 从缓存获取数据失败: {kwargs}", e=e) @@ -234,8 +239,9 @@ class DataAccess(Generic[T]): cache_key = self._build_cache_key_for_item(data) if cache_key is not None: # 存入缓存 - await self.cache.set(cache_key, data) - self._cache_stats[self.cache_type]["sets"] += 1 + if self.cache: + await self.cache.set(cache_key, data) + self._bump_cache_stat(self.cache_type, "sets") logger.debug( f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}" ) @@ -249,8 +255,8 @@ class DataAccess(Generic[T]): # 存入空结果缓存,使用较短的过期时间 await self.cache.set( cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL - ) - self._cache_stats[self.cache_type]["null_sets"] += 1 + ) if self.cache else None + self._bump_cache_stat(self.cache_type, "null_sets") logger.debug( f"{self.model_cls.__name__} 空结果已存入缓存: {cache_key}," f" TTL={self._NULL_RESULT_TTL}秒" @@ -343,8 +349,9 @@ class DataAccess(Generic[T]): return False # 删除缓存 - await self.cache.delete(cache_key) - self._cache_stats[self.cache_type]["deletes"] += 1 + if self.cache: + await self.cache.delete(cache_key) + self._bump_cache_stat(self.cache_type, "deletes") logger.debug(f"已清除{self.model_cls.__name__}缓存: {cache_key}") return True except Exception as e: @@ -416,8 +423,9 @@ class DataAccess(Generic[T]): if cache_key is None: return - await self.cache.delete(cache_key) - self._cache_stats[self.cache_type]["deletes"] += 1 + if self.cache: + await self.cache.delete(cache_key) + self._bump_cache_stat(self.cache_type, "deletes") logger.debug( f"{self.model_cls.__name__} {action}: 已失效兼容缓存: {cache_key}" ) @@ -543,8 +551,9 @@ class DataAccess(Generic[T]): if cache_key is not None: # 如果成功构建缓存键,直接删除缓存 - await self.cache.delete(cache_key) - self._cache_stats[self.cache_type]["deletes"] += 1 + if self.cache: + await self.cache.delete(cache_key) + self._bump_cache_stat(self.cache_type, "deletes") logger.debug( f"{self.model_cls.__name__} delete: 已删除缓存: {cache_key}" ) @@ -558,8 +567,9 @@ class DataAccess(Generic[T]): for item in items: item_cache_key = self._build_cache_key_for_item(item) if item_cache_key is not None: - await self.cache.delete(item_cache_key) - self._cache_stats[self.cache_type]["deletes"] += 1 + if self.cache: + await self.cache.delete(item_cache_key) + self._bump_cache_stat(self.cache_type, "deletes") if items: logger.debug( f"{self.model_cls.__name__} delete:" diff --git a/zhenxun/services/db_context/__init__.py b/zhenxun/services/db_context/__init__.py index b144dae0..7f3cd8e8 100644 --- a/zhenxun/services/db_context/__init__.py +++ b/zhenxun/services/db_context/__init__.py @@ -18,6 +18,7 @@ from zhenxun.configs.config import BotConfig from zhenxun.services.log import logger from zhenxun.utils.manager.priority_manager import PriorityLifecycle +from . import watchdog as _watchdog # noqa: F401 from .base_model import Model from .config import ( DB_TIMEOUT_SECONDS, diff --git a/zhenxun/services/db_context/base_model.py b/zhenxun/services/db_context/base_model.py index 91c998a2..63b419db 100644 --- a/zhenxun/services/db_context/base_model.py +++ b/zhenxun/services/db_context/base_model.py @@ -27,7 +27,7 @@ class Model(TortoiseModel): """ sem_data: ClassVar[dict[str, dict[str, asyncio.Semaphore]]] = {} - _current_locks: ClassVar[dict[int, DbLockType]] = {} # 跟踪当前协程持有的锁 + _current_locks: ClassVar[dict[tuple[str, int], DbLockType]] = {} def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) @@ -95,21 +95,21 @@ class Model(TortoiseModel): @classmethod def _require_lock(cls, lock_type: DbLockType) -> bool: """检查是否需要真正加锁""" - task_id = id(asyncio.current_task()) - return cls._current_locks.get(task_id) != lock_type + lock_key = (cls.__name__, id(asyncio.current_task())) + return cls._current_locks.get(lock_key) != lock_type @classmethod @contextlib.asynccontextmanager async def _lock_context(cls, lock_type: DbLockType): """带重入检查的锁上下文""" - task_id = id(asyncio.current_task()) + lock_key = (cls.__name__, id(asyncio.current_task())) need_lock = cls._require_lock(lock_type) if need_lock and (sem := cls.get_semaphore(lock_type)): - cls._current_locks[task_id] = lock_type + cls._current_locks[lock_key] = lock_type async with sem: yield - cls._current_locks.pop(task_id, None) + cls._current_locks.pop(lock_key, None) else: yield @@ -171,8 +171,8 @@ class Model(TortoiseModel): await obj.save() result = (obj, False) else: - # 创建时不重复加锁 - result = await cls.create(**kwargs, **(defaults or {})), True + obj = await super().create(**kwargs, **(defaults or {})) + result = (obj, True) if cache_type := cls.get_cache_type(): await CacheRoot.invalidate_cache( diff --git a/zhenxun/services/db_context/config.py b/zhenxun/services/db_context/config.py index 98d7d8c0..38a19b41 100644 --- a/zhenxun/services/db_context/config.py +++ b/zhenxun/services/db_context/config.py @@ -27,7 +27,9 @@ MYSQL_CONFIG = { SQLITE_CONFIG = { "journal_mode": "WAL", # 提高并发写入性能 - "busy_timeout": 30000, # SQLite 锁等待超时,单位毫秒 + # SQLite 的底层锁等待应接近业务超时,避免上层放弃后 worker 仍长时间占线。 + # Windows bind mount / Docker 场景下 SQLite 不适合高写并发。 + "busy_timeout": 5000, "foreign_keys": "ON", } diff --git a/zhenxun/services/db_context/utils.py b/zhenxun/services/db_context/utils.py index a1bb3824..8873f188 100644 --- a/zhenxun/services/db_context/utils.py +++ b/zhenxun/services/db_context/utils.py @@ -1,7 +1,9 @@ import asyncio +import contextlib import time from zhenxun.services.log import logger +from zhenxun.services.message_load import signal_db_unhealthy from .config import ( DB_TIMEOUT_SECONDS, @@ -9,6 +11,40 @@ from .config import ( SLOW_QUERY_THRESHOLD, ) +_SQLITE_STALL_UNTIL = 0.0 +_SQLITE_STALL_REASON = "" +_DB_UNHEALTHY_TIMEOUT_SECONDS = 30.0 +_SQLITE_STALL_TIMEOUT_SECONDS = 60.0 + + +def _is_sqlite_connection() -> bool: + with contextlib.suppress(Exception): + from tortoise import Tortoise + + connection = Tortoise.get_connection("default") + capabilities = getattr(connection, "capabilities", None) + dialect = str(getattr(capabilities, "dialect", "") or "").lower() + return dialect.startswith("sqlite") + return False + + +def _mark_sqlite_stall(reason: str, duration: float) -> None: + global _SQLITE_STALL_REASON, _SQLITE_STALL_UNTIL + until = time.monotonic() + max(duration, 0.0) + if until > _SQLITE_STALL_UNTIL: + _SQLITE_STALL_UNTIL = until + _SQLITE_STALL_REASON = str(reason or "")[:200] + + +def is_sqlite_stall_suspected() -> bool: + return time.monotonic() < _SQLITE_STALL_UNTIL + + +def sqlite_stall_reason() -> str: + if not is_sqlite_stall_suspected(): + return "" + return _SQLITE_STALL_REASON + async def with_db_timeout( coro, @@ -19,13 +55,23 @@ async def with_db_timeout( """带超时控制的数据库操作""" start_time = time.time() try: - logger.debug(f"开始执行数据库操作: {operation} 来源: {source}") result = await asyncio.wait_for(coro, timeout=timeout) elapsed = time.time() - start_time if elapsed > SLOW_QUERY_THRESHOLD and operation: logger.warning(f"慢查询: {operation} 耗时 {elapsed:.3f}s", LOG_COMMAND) return result except asyncio.TimeoutError: + timeout_reason = f"{operation or 'database_operation'} from {source or '-'}" + unhealthy_duration = _DB_UNHEALTHY_TIMEOUT_SECONDS + if _is_sqlite_connection(): + unhealthy_duration = _SQLITE_STALL_TIMEOUT_SECONDS + _mark_sqlite_stall(timeout_reason, unhealthy_duration) + logger.warning( + "SQLite 数据库操作超时,疑似 aiosqlite worker/连接被锁等待卡住;" + "已暂停低优先级数据库任务", + LOG_COMMAND, + ) + signal_db_unhealthy(unhealthy_duration, reason=timeout_reason) if operation: logger.error( f"数据库操作超时: {operation} (>{timeout}s) 来源: {source}", diff --git a/zhenxun/services/db_context/watchdog.py b/zhenxun/services/db_context/watchdog.py new file mode 100644 index 00000000..010134f8 --- /dev/null +++ b/zhenxun/services/db_context/watchdog.py @@ -0,0 +1,123 @@ +from __future__ import annotations + +import asyncio +import contextlib +import time + +from tortoise import Tortoise +from tortoise.connection import connections + +from zhenxun.services.log import logger +from zhenxun.services.low_priority_writer import low_priority_writer_active_count +from zhenxun.services.message_load import ( + signal_db_unhealthy, +) +from zhenxun.utils.manager.priority_manager import PriorityLifecycle + +from .config import LOG_COMMAND + +_CHECK_INTERVAL_SECONDS = 15.0 +_CHECK_TIMEOUT_SECONDS = 2.0 +_FAIL_THRESHOLD = 3 +_UNHEALTHY_SECONDS = 60.0 +_RECONNECT_COOLDOWN_SECONDS = 60.0 +_RECONNECT_WAIT_IDLE_SECONDS = 3.0 + +_WATCHDOG_TASK: asyncio.Task[None] | None = None +_RECONNECT_LOCK = asyncio.Lock() +_LAST_RECONNECT_AT = 0.0 + + +def _is_sqlite_connection() -> bool: + with contextlib.suppress(Exception): + connection = Tortoise.get_connection("default") + capabilities = getattr(connection, "capabilities", None) + dialect = str(getattr(capabilities, "dialect", "") or "").lower() + return dialect.startswith("sqlite") + return False + + +async def _select_one() -> None: + connection = Tortoise.get_connection("default") + await connection.execute_query("SELECT 1") + + +async def _try_reconnect(reason: str) -> None: + global _LAST_RECONNECT_AT + now = time.monotonic() + if now - _LAST_RECONNECT_AT < _RECONNECT_COOLDOWN_SECONDS: + return + if low_priority_writer_active_count() > 0: + return + async with _RECONNECT_LOCK: + now = time.monotonic() + if now - _LAST_RECONNECT_AT < _RECONNECT_COOLDOWN_SECONDS: + return + await asyncio.sleep(_RECONNECT_WAIT_IDLE_SECONDS) + if low_priority_writer_active_count() > 0: + return + try: + await connections.close_all(discard=True) + # ConnectionHandler lazily recreates default connection from db_config. + Tortoise.get_connection("default") + _LAST_RECONNECT_AT = time.monotonic() + logger.warning( + f"SQLite watchdog rebuilt default connection: {reason}", + LOG_COMMAND, + ) + except Exception as exc: + _LAST_RECONNECT_AT = time.monotonic() + signal_db_unhealthy( + _UNHEALTHY_SECONDS, + reason=f"watchdog reconnect:{reason}", + ) + logger.warning("SQLite watchdog reconnect failed", LOG_COMMAND, e=exc) + + +async def _watchdog_loop() -> None: + failures = 0 + while True: + await asyncio.sleep(_CHECK_INTERVAL_SECONDS) + if not _is_sqlite_connection(): + failures = 0 + continue + try: + await asyncio.wait_for(_select_one(), timeout=_CHECK_TIMEOUT_SECONDS) + failures = 0 + except asyncio.CancelledError: + raise + except Exception as exc: + failures += 1 + reason = f"sqlite watchdog SELECT 1 failed x{failures}: {exc}" + signal_db_unhealthy(_UNHEALTHY_SECONDS, reason=reason) + logger.warning(reason, LOG_COMMAND) + if failures >= _FAIL_THRESHOLD: + await _try_reconnect(reason) + + +def start_db_watchdog() -> None: + global _WATCHDOG_TASK + if _WATCHDOG_TASK is not None and not _WATCHDOG_TASK.done(): + return + _WATCHDOG_TASK = asyncio.create_task(_watchdog_loop()) + + +async def stop_db_watchdog() -> None: + global _WATCHDOG_TASK + task = _WATCHDOG_TASK + _WATCHDOG_TASK = None + if task is None: + return + task.cancel() + with contextlib.suppress(BaseException): + await task + + +@PriorityLifecycle.on_startup(priority=8) +async def _start_db_watchdog() -> None: + start_db_watchdog() + + +@PriorityLifecycle.on_shutdown(priority=10) +async def _stop_db_watchdog() -> None: + await stop_db_watchdog() diff --git a/zhenxun/services/group_settings_service.py b/zhenxun/services/group_settings_service.py index 366f71d8..de79ec0e 100644 --- a/zhenxun/services/group_settings_service.py +++ b/zhenxun/services/group_settings_service.py @@ -6,7 +6,6 @@ import ujson as json from zhenxun.configs.config import Config from zhenxun.models.group_plugin_setting import GroupPluginSetting from zhenxun.services.cache import BoundedTTLCache -from zhenxun.services.data_access import DataAccess from zhenxun.services.log import logger from zhenxun.utils.pydantic_compat import model_dump, model_validate, parse_as @@ -20,7 +19,6 @@ class GroupSettingsService: """ def __init__(self): - self.dao = DataAccess(GroupPluginSetting) self._cache = BoundedTTLCache[str, dict[str, Any]]( "GROUP_PLUGIN_SETTINGS_VIEW", ttl_seconds=600, @@ -48,13 +46,12 @@ class GroupSettingsService: settings_dict = model_dump(settings_model) json_value = json.dumps(settings_dict, ensure_ascii=False) - await self.dao.update_or_create( + await GroupPluginSetting.update_or_create( defaults={"settings": json_value}, # type: ignore group_id=group_id, plugin_name=plugin_name, ) - await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name) await self._clear_merged_cache(group_id, plugin_name) async def set_key_value( @@ -72,19 +69,19 @@ class GroupSettingsService: setting_entry.settings[key] = value await setting_entry.save(update_fields=["settings"]) - await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name) await self._clear_merged_cache(group_id, plugin_name) async def reset_key(self, group_id: str, plugin_name: str, key: str) -> bool: """重置单个配置项""" - setting = await self.dao.get_or_none(group_id=group_id, plugin_name=plugin_name) + setting = await GroupPluginSetting.safe_get_or_none( + group_id=group_id, plugin_name=plugin_name + ) if setting and isinstance(setting.settings, dict) and key in setting.settings: del setting.settings[key] if not setting.settings: await setting.delete() else: await setting.save(update_fields=["settings"]) - await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name) await self._clear_merged_cache(group_id, plugin_name) return True return False @@ -119,12 +116,11 @@ class GroupSettingsService: 返回: bool: 如果成功删除了一个条目,则返回 True,否则返回 False。 """ - deleted_count = await self.dao.delete( + deleted_count = await GroupPluginSetting.filter( group_id=group_id, plugin_name=plugin_name - ) + ).delete() if deleted_count > 0: - await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name) await self._clear_merged_cache(group_id, plugin_name) logger.debug(f"已重置插件 '{plugin_name}' 在群组 '{group_id}' 的配置。") return True @@ -176,7 +172,7 @@ class GroupSettingsService: for key in global_config_group.configs.keys() } - group_setting_entry = await self.dao.get_or_none( + group_setting_entry = await GroupPluginSetting.safe_get_or_none( group_id=group_id, plugin_name=plugin_name ) if group_setting_entry: diff --git a/zhenxun/services/hot_query_cache.py b/zhenxun/services/hot_query_cache.py index 55bf67b4..ec4079c0 100644 --- a/zhenxun/services/hot_query_cache.py +++ b/zhenxun/services/hot_query_cache.py @@ -9,6 +9,8 @@ from typing import Any, Literal from tortoise.functions import Count from zhenxun.services.cache.bounded_ttl import BoundedTTLCache +from zhenxun.services.db_context import with_db_timeout +from zhenxun.services.message_load import is_db_unhealthy @dataclass(frozen=True, slots=True) @@ -81,6 +83,26 @@ _CHAT_RANK_LOCKS: dict[str, asyncio.Lock] = {} _CHAT_FIRST_MSG_LOCKS: dict[str, asyncio.Lock] = {} _STATISTICS_LOCKS: dict[str, asyncio.Lock] = {} _MAX_LOCK_POOL_SIZE = 4096 +_MEMBER_DB_TIMEOUT = 2.0 +_AGGREGATE_DB_TIMEOUT = 3.0 + + +async def _read_or_default( + coro, + *, + timeout: float, + operation: str, + default, +): + try: + return await with_db_timeout( + coro, + timeout=timeout, + operation=operation, + source="hot_query_cache", + ) + except TimeoutError: + return default def _get_lock(pool: dict[str, asyncio.Lock], key: str) -> asyncio.Lock: @@ -126,13 +148,20 @@ async def get_group_members( from zhenxun.models.group_member_info import GroupInfoUser - rows = await GroupInfoUser.filter(group_id=group_key).values_list( - "id", - "user_id", - "user_name", - "user_join_time", - "uid", - "platform", + if is_db_unhealthy(): + return () + rows = await _read_or_default( + GroupInfoUser.filter(group_id=group_key).values_list( + "id", + "user_id", + "user_name", + "user_join_time", + "uid", + "platform", + ), + timeout=_MEMBER_DB_TIMEOUT, + operation="hot_query_cache.get_group_members", + default=(), ) members = tuple( GroupMemberSnapshot( @@ -192,15 +221,20 @@ async def get_group_member_map( if missing: from zhenxun.models.group_member_info import GroupInfoUser - rows = await GroupInfoUser.filter( - group_id=group_key, user_id__in=missing - ).values_list( - "id", - "user_id", - "user_name", - "user_join_time", - "uid", - "platform", + if is_db_unhealthy(): + return result + rows = await _read_or_default( + GroupInfoUser.filter(group_id=group_key, user_id__in=missing).values_list( + "id", + "user_id", + "user_name", + "user_join_time", + "uid", + "platform", + ), + timeout=_MEMBER_DB_TIMEOUT, + operation="hot_query_cache.get_group_member_map", + default=(), ) found: set[str] = set() for row in rows: @@ -258,8 +292,13 @@ async def get_group_user_ids(group_id: str | int | None) -> set[str]: from zhenxun.models.group_member_info import GroupInfoUser - rows = await GroupInfoUser.filter(group_id=group_key).values_list( - "user_id", flat=True + if is_db_unhealthy(): + return set() + rows = await _read_or_default( + GroupInfoUser.filter(group_id=group_key).values_list("user_id", flat=True), + timeout=_MEMBER_DB_TIMEOUT, + operation="hot_query_cache.get_group_user_ids", + default=(), ) user_ids = tuple(str(user_id) for user_id in rows if user_id) await _GROUP_USER_IDS_CACHE.set(group_key, user_ids) @@ -283,8 +322,13 @@ async def get_user_group_ids(user_id: str | int | None) -> list[str]: from zhenxun.models.group_member_info import GroupInfoUser - rows = await GroupInfoUser.filter(user_id=user_key).values_list( - "group_id", flat=True + if is_db_unhealthy(): + return [] + rows = await _read_or_default( + GroupInfoUser.filter(user_id=user_key).values_list("group_id", flat=True), + timeout=_MEMBER_DB_TIMEOUT, + operation="hot_query_cache.get_user_group_ids", + default=(), ) group_ids = tuple(str(group_id) for group_id in rows if group_id) await _USER_GROUP_CACHE.set(user_key, group_ids) @@ -314,8 +358,15 @@ async def get_member_names( if missing: from zhenxun.models.group_member_info import GroupInfoUser - rows = await GroupInfoUser.filter(user_id__in=missing).values_list( - "user_id", "user_name" + if is_db_unhealthy(): + return result + rows = await _read_or_default( + GroupInfoUser.filter(user_id__in=missing).values_list( + "user_id", "user_name" + ), + timeout=_MEMBER_DB_TIMEOUT, + operation="hot_query_cache.get_member_names", + default=(), ) for user_id, user_name in rows: user_key = str(user_id) @@ -395,6 +446,8 @@ async def get_chat_history_rank_cached( if cached is not None: return list(cached) + if is_db_unhealthy(): + return [] order_prefix = "-" if order == "DESC" else "" query: Any = model.filter(group_id=gid) if gid else model if date_scope: @@ -403,12 +456,15 @@ async def get_chat_history_rank_cached( date_scope[1].isoformat(" "), ) query = query.filter(create_time__range=filter_scope) - rows = ( - await query.annotate(count=Count("user_id")) + rows = await _read_or_default( + query.annotate(count=Count("user_id")) .order_by(f"{order_prefix}count") .group_by("user_id") .limit(limit) - .values_list("user_id", "count") + .values_list("user_id", "count"), + timeout=_AGGREGATE_DB_TIMEOUT, + operation="hot_query_cache.get_chat_history_rank", + default=(), ) result = tuple((str(user_id), int(count)) for user_id, count in rows) await _CHAT_RANK_CACHE.set(key, result) @@ -430,8 +486,15 @@ async def get_chat_history_first_msg_datetime_cached( if cached is not None: return cached[0] + if is_db_unhealthy(): + return None query: Any = model.filter(group_id=group_id) if group_id else model.all() - message = await query.order_by("create_time").first() + message = await _read_or_default( + query.order_by("create_time").first(), + timeout=_AGGREGATE_DB_TIMEOUT, + operation="hot_query_cache.get_chat_history_first_msg", + default=None, + ) result = getattr(message, "create_time", None) if message else None await _CHAT_FIRST_MSG_CACHE.set(key, (result,)) return result @@ -459,6 +522,8 @@ async def get_statistics_plugin_counts_cached( if cached is not None: return list(cached) + if is_db_unhealthy(): + return [] from zhenxun.models.statistics import Statistics query: Any = Statistics @@ -472,10 +537,13 @@ async def get_statistics_plugin_counts_cached( query = query.filter(plugin_name=plugin_name) if start_time: query = query.filter(create_time__gte=start_time) - rows = ( - await query.annotate(count=Count("id")) + rows = await _read_or_default( + query.annotate(count=Count("id")) .group_by("plugin_name") - .values_list("plugin_name", "count") + .values_list("plugin_name", "count"), + timeout=_AGGREGATE_DB_TIMEOUT, + operation="hot_query_cache.get_statistics_plugin_counts", + default=(), ) result = tuple((str(plugin), int(count)) for plugin, count in rows) await _STATISTICS_COUNT_CACHE.set(key, result) diff --git a/zhenxun/services/low_priority_writer.py b/zhenxun/services/low_priority_writer.py new file mode 100644 index 00000000..83286168 --- /dev/null +++ b/zhenxun/services/low_priority_writer.py @@ -0,0 +1,329 @@ +from __future__ import annotations + +import asyncio +from collections import deque +from collections.abc import Awaitable, Callable +from dataclasses import dataclass, field +import time +from typing import Any + +from zhenxun.services.log import logger +from zhenxun.services.message_load import should_pause_tasks, signal_db_unhealthy +from zhenxun.utils.manager.priority_manager import PriorityLifecycle + +LOG_COMMAND = "LowPriorityWriter" + +WriteBatch = Callable[[list[Any], str], Awaitable[None]] + +_POLL_INTERVAL_SECONDS = 1.0 +_DB_UNHEALTHY_SECONDS = 30.0 + + +@dataclass(slots=True) +class LowPriorityWriterConfig: + name: str + write_batch: WriteBatch + batch_size: int = 500 + trigger_size: int = 500 + max_retain: int = 10_000 + flush_interval_seconds: float = 60.0 + max_items_per_cycle: int = 1_000 + backoff_base_seconds: float = 30.0 + backoff_max_seconds: float = 600.0 + log_command: str = LOG_COMMAND + + +@dataclass(slots=True) +class _WriterState: + config: LowPriorityWriterConfig + buffer: deque[Any] = field(default_factory=deque) + lock: asyncio.Lock = field(default_factory=asyncio.Lock) + dropped: int = 0 + last_drop_log_at: float = 0.0 + last_flush_at: float = field(default_factory=time.monotonic) + failures: int = 0 + backoff_until: float = 0.0 + + +_WRITERS: dict[str, _WriterState] = {} +_WORKER_TASK: asyncio.Task[None] | None = None +_WAKE_EVENT: asyncio.Event | None = None +_FLUSH_LOCK = asyncio.Lock() +_ACTIVE_FLUSHES = 0 +_STOPPING = False + + +def _wake() -> None: + if _WAKE_EVENT is not None: + _WAKE_EVENT.set() + + +def _ensure_worker() -> None: + global _STOPPING, _WAKE_EVENT, _WORKER_TASK + if _STOPPING: + _STOPPING = False + try: + loop = asyncio.get_running_loop() + except RuntimeError: + return + if _WAKE_EVENT is None: + _WAKE_EVENT = asyncio.Event() + if _WORKER_TASK is not None and not _WORKER_TASK.done(): + return + _WORKER_TASK = loop.create_task(_worker_loop()) + + +def register_low_priority_writer(config: LowPriorityWriterConfig) -> None: + """Register or update an append-only low-priority DB writer.""" + if config.batch_size <= 0: + raise ValueError("batch_size must be positive") + if config.trigger_size <= 0: + raise ValueError("trigger_size must be positive") + if config.max_retain <= 0: + raise ValueError("max_retain must be positive") + state = _WRITERS.get(config.name) + if state is None: + _WRITERS[config.name] = _WriterState(config=config) + else: + state.config = config + _ensure_worker() + + +async def append_low_priority_record(name: str, record: Any) -> bool: + """Append a record without doing DB work in the caller's hot path.""" + state = _WRITERS.get(name) + if state is None: + raise KeyError(f"low priority writer not registered: {name}") + _ensure_worker() + should_wake = False + async with state.lock: + if len(state.buffer) >= state.config.max_retain: + state.buffer.popleft() + state.dropped += 1 + _log_drop_if_needed(state) + state.buffer.append(record) + should_wake = len(state.buffer) >= state.config.trigger_size + if should_wake: + _wake() + return True + + +def _log_drop_if_needed(state: _WriterState) -> None: + now = time.monotonic() + if now - state.last_drop_log_at < 10.0: + return + state.last_drop_log_at = now + logger.warning( + f"{state.config.name} low priority buffer full, " + f"dropped={state.dropped}, backlog={len(state.buffer)}", + state.config.log_command, + ) + + +async def flush_low_priority_writer( + name: str, + reason: str, + *, + force: bool = False, +) -> int: + state = _WRITERS.get(name) + if state is None: + return 0 + async with _FLUSH_LOCK: + return await _flush_state(state, reason, force=force) + + +async def flush_all_low_priority_writers( + reason: str, + *, + force: bool = False, +) -> int: + total = 0 + async with _FLUSH_LOCK: + for state in list(_WRITERS.values()): + total += await _flush_state(state, reason, force=force) + return total + + +async def _worker_loop() -> None: + while not _STOPPING: + event = _WAKE_EVENT + if event is None: + await asyncio.sleep(_POLL_INTERVAL_SECONDS) + else: + try: + await asyncio.wait_for( + event.wait(), + timeout=_POLL_INTERVAL_SECONDS, + ) + except asyncio.TimeoutError: + pass + event.clear() + if should_pause_tasks(): + continue + async with _FLUSH_LOCK: + for state in list(_WRITERS.values()): + if _state_due_for_flush(state): + await _flush_state(state, "低优先队列") + + +def _state_due_for_flush(state: _WriterState) -> bool: + if not state.buffer: + return False + now = time.monotonic() + if now < state.backoff_until: + return False + return ( + len(state.buffer) >= state.config.trigger_size + or now - state.last_flush_at >= state.config.flush_interval_seconds + ) + + +async def _flush_state( + state: _WriterState, + reason: str, + *, + force: bool = False, +) -> int: + if not force: + if should_pause_tasks(): + return 0 + if time.monotonic() < state.backoff_until: + return 0 + written = 0 + max_items = state.config.max_items_per_cycle if not force else float("inf") + while written < max_items: + batch = await _take_batch(state) + if not batch: + break + try: + await _write_batch(state, batch, reason) + except (TimeoutError, asyncio.TimeoutError) as exc: + # asyncio.wait_for timeout only cancels the awaiter. With SQLite/aiosqlite + # the worker thread may still finish the SQL later, so restoring this + # append-only low priority batch can duplicate rows. Prefer dropping the + # uncertain batch; chat history/statistics/logs are lossy by design here. + _mark_uncertain_timeout(state, reason, exc, len(batch)) + break + except Exception as exc: + await _restore_batch(state, batch) + _mark_failure(state, reason, exc) + break + written += len(batch) + state.failures = 0 + state.backoff_until = 0.0 + state.last_flush_at = time.monotonic() + if written: + logger.debug( + f"{reason}写入 {state.config.name} {written} 条, " + f"backlog={len(state.buffer)}", + state.config.log_command, + ) + return written + + +async def _take_batch(state: _WriterState) -> list[Any]: + batch: list[Any] = [] + async with state.lock: + while state.buffer and len(batch) < state.config.batch_size: + batch.append(state.buffer.popleft()) + return batch + + +async def _restore_batch(state: _WriterState, batch: list[Any]) -> None: + if not batch: + return + async with state.lock: + retain_count = max(state.config.max_retain - len(state.buffer), 0) + restore_items = batch[-retain_count:] if retain_count else [] + for record in reversed(restore_items): + state.buffer.appendleft(record) + dropped = len(batch) - len(restore_items) + if dropped: + state.dropped += dropped + _log_drop_if_needed(state) + + +async def _write_batch( + state: _WriterState, + batch: list[Any], + reason: str, +) -> None: + global _ACTIVE_FLUSHES + _ACTIVE_FLUSHES += 1 + try: + await state.config.write_batch(batch, reason) + finally: + _ACTIVE_FLUSHES = max(_ACTIVE_FLUSHES - 1, 0) + + +def _mark_failure(state: _WriterState, reason: str, exc: Exception) -> None: + state.failures += 1 + backoff = min( + state.config.backoff_base_seconds * (2 ** (state.failures - 1)), + state.config.backoff_max_seconds, + ) + state.backoff_until = time.monotonic() + backoff + signal_db_unhealthy(_DB_UNHEALTHY_SECONDS, reason=f"{state.config.name}:{reason}") + logger.warning( + f"{reason}写入 {state.config.name} 失败, " + f"backoff={backoff:.0f}s, backlog={len(state.buffer)}", + state.config.log_command, + e=exc, + ) + + +def _mark_uncertain_timeout( + state: _WriterState, + reason: str, + exc: BaseException, + batch_size: int, +) -> None: + state.failures += 1 + backoff = min( + state.config.backoff_base_seconds * (2 ** (state.failures - 1)), + state.config.backoff_max_seconds, + ) + state.backoff_until = time.monotonic() + backoff + state.dropped += batch_size + signal_db_unhealthy(_DB_UNHEALTHY_SECONDS, reason=f"{state.config.name}:{reason}") + log_exc = exc if isinstance(exc, Exception) else None + logger.warning( + f"{reason}写入 {state.config.name} 超时, " + f"dropped_uncertain={batch_size}, backoff={backoff:.0f}s, " + f"backlog={len(state.buffer)}", + state.config.log_command, + e=log_exc, + ) + + +def low_priority_writer_active_count() -> int: + return _ACTIVE_FLUSHES + + +def low_priority_writer_backlog() -> dict[str, int]: + return {name: len(state.buffer) for name, state in _WRITERS.items()} + + +async def stop_low_priority_writer() -> int: + global _STOPPING, _WORKER_TASK + _STOPPING = True + task = _WORKER_TASK + _WORKER_TASK = None + if task is not None: + task.cancel() + try: + await task + except BaseException: + pass + return await flush_all_low_priority_writers("关闭", force=True) + + +@PriorityLifecycle.on_startup(priority=3) +async def _start_low_priority_writer() -> None: + _ensure_worker() + + +@PriorityLifecycle.on_shutdown(priority=95) +async def _stop_low_priority_writer() -> None: + await stop_low_priority_writer() diff --git a/zhenxun/services/memory_governor.py b/zhenxun/services/memory_governor.py index b186c676..b4a3dc44 100644 --- a/zhenxun/services/memory_governor.py +++ b/zhenxun/services/memory_governor.py @@ -12,7 +12,7 @@ from aiocache import SimpleMemoryCache from zhenxun.services.cache import CacheRoot from zhenxun.services.cache.bounded_ttl import BoundedTTLCache -from zhenxun.services.cache.cache_containers import CacheDict, CacheList +from zhenxun.services.cache.cache_containers import CacheDict from zhenxun.services.log import logger from zhenxun.services.message_load import idle_seconds, is_overloaded @@ -131,7 +131,6 @@ async def _run_reclaim() -> None: cleared: dict[str, Any] = {} cache_stats_before = { "cache_dict": CacheDict.stats_all(), - "cache_list": CacheList.stats_all(), "bounded_ttl": await BoundedTTLCache.stats_all(), } @@ -139,7 +138,6 @@ async def _run_reclaim() -> None: cleared["user_gold_logs"] = await _flush_user_gold_log_buffer() cleared["bounded_ttl_clear"] = await BoundedTTLCache.clear_all() cleared["cache_dict_clear"] = CacheDict.clear_all() - cleared["cache_list_clear"] = CacheList.clear_all() cleared["runtime_negative"] = _clear_runtime_negative_caches() cleared["auth_local"] = _clear_auth_local_caches() cleared["avatar_l1"] = _clear_avatar_memory_cache() diff --git a/zhenxun/services/message_load.py b/zhenxun/services/message_load.py index 6b7294d0..3a320110 100644 --- a/zhenxun/services/message_load.py +++ b/zhenxun/services/message_load.py @@ -3,6 +3,8 @@ from __future__ import annotations import time _OVERLOAD_UNTIL = 0.0 +_DB_UNHEALTHY_UNTIL = 0.0 +_DB_UNHEALTHY_REASON = "" _LAST_ACTIVITY = time.monotonic() @@ -31,5 +33,27 @@ def is_overloaded() -> bool: return time.monotonic() < _OVERLOAD_UNTIL +def signal_db_unhealthy(duration: float = 30.0, reason: str = "") -> None: + """Mark database-dependent low-priority tasks as unsafe to run briefly.""" + global _DB_UNHEALTHY_REASON, _DB_UNHEALTHY_UNTIL + if duration <= 0: + return + now = time.monotonic() + until = now + duration + if until > _DB_UNHEALTHY_UNTIL: + _DB_UNHEALTHY_UNTIL = until + _DB_UNHEALTHY_REASON = str(reason or "")[:200] + + +def is_db_unhealthy() -> bool: + return time.monotonic() < _DB_UNHEALTHY_UNTIL + + +def db_unhealthy_reason() -> str: + if not is_db_unhealthy(): + return "" + return _DB_UNHEALTHY_REASON + + def should_pause_tasks() -> bool: - return is_overloaded() + return is_overloaded() or is_db_unhealthy()