mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-28 16:20:56 +08:00
bugfix:修复sqlite部分场景下锁竞争问题;新增插件恶意触发配置文件 (#2144)
* bugfix:修复sqlite部分场景下锁竞争问题;新增插件恶意触发配置文件 * 文件没同步完
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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]:
|
||||
"""验证路径是否安全
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 = "插件分群通用配置表"
|
||||
|
||||
@@ -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)
|
||||
|
||||
Vendored
+9
-21
@@ -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
@@ -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
@@ -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] = {}
|
||||
|
||||
@@ -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:"
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
|
||||
@@ -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}",
|
||||
|
||||
@@ -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()
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user