bugfix:修复sqlite部分场景下锁竞争问题;新增插件恶意触发配置文件 (#2144)

* bugfix:修复sqlite部分场景下锁竞争问题;新增插件恶意触发配置文件

* 文件没同步完
This commit is contained in:
Copaan
2026-06-24 09:11:03 +08:00
committed by GitHub
parent 73cbe2a609
commit f4d2342693
39 changed files with 1577 additions and 829 deletions
@@ -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)
+74 -24
View File
@@ -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:
+18
View File
@@ -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",
+15 -2
View File
@@ -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:
@@ -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:
@@ -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,
@@ -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:
+94 -14
View File
@@ -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),
)
+121 -28
View File
@@ -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)
+53 -12
View File
@@ -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,
},
+2 -15
View File
@@ -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}")
@@ -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(
@@ -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("关闭")
@@ -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}")
@@ -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 = []
@@ -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)
@@ -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:
@@ -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}")
@@ -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,
@@ -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}")
+18
View File
@@ -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]:
"""验证路径是否安全
+3 -28
View File
@@ -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,
+8 -12
View File
@@ -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,
-6
View File
@@ -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 = "插件分群通用配置表"
+32 -93
View File
@@ -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)
+9 -21
View File
@@ -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
-282
View File
@@ -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)"
+78 -8
View File
@@ -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] = {}
+30 -20
View File
@@ -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:"
+1
View File
@@ -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,
+8 -8
View File
@@ -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(
+3 -1
View File
@@ -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",
}
+47 -1
View File
@@ -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}",
+123
View File
@@ -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()
+7 -11
View File
@@ -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:
+97 -29
View File
@@ -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)
+329
View File
@@ -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()
+1 -3
View File
@@ -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()
+25 -1
View File
@@ -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()