mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-07 04:40:00 +08:00
bugfix:更换数据库初始化超时路径以修复连接超时问题 (#2137)
* bugfix:更换数据库初始化超时路径以修复连接超时问题 * 移除部分观测链路 * 细节修改 * 权限检查细节修改2 * bugfix:修复金币懒加载造成插件金币消耗不了的问题 * bugfix:整理鉴权逻辑 * 完善缓存系统 * 优化官端使用 * bugfix:增加官机和频道的前导统一剥离,新增获取群列表自愈 * bugfix:修复导入问题
This commit is contained in:
@@ -1,547 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections import deque
|
||||
import contextlib
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta
|
||||
import json
|
||||
import random
|
||||
import time
|
||||
from typing import Any, TypeVar
|
||||
|
||||
from tortoise import Tortoise
|
||||
|
||||
from zhenxun.builtin_plugins.hooks.auth_runtime_config import (
|
||||
AUTH_OBSERVABILITY_RUNTIME_CONFIG,
|
||||
)
|
||||
from zhenxun.models.auth_decision_log import AuthDecisionLog
|
||||
from zhenxun.models.runtime_backpressure_log import RuntimeBackpressureLog
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
|
||||
LOG_COMMAND = "AuthObservability"
|
||||
|
||||
_BUFFER_MAX_RETAIN = AUTH_OBSERVABILITY_RUNTIME_CONFIG.buffer_max_retain
|
||||
_FLUSH_TRIGGER_SIZE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_trigger_size
|
||||
_FLUSH_BATCH_SIZE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_batch_size
|
||||
_FLUSH_INTERVAL_SECONDS = AUTH_OBSERVABILITY_RUNTIME_CONFIG.flush_interval_seconds
|
||||
_DROP_LOG_INTERVAL_SECONDS = AUTH_OBSERVABILITY_RUNTIME_CONFIG.drop_log_interval_seconds
|
||||
_ALLOW_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.allow_sample_rate
|
||||
_OVERLOADED_ALLOW_SAMPLE_RATE = (
|
||||
AUTH_OBSERVABILITY_RUNTIME_CONFIG.overloaded_allow_sample_rate
|
||||
)
|
||||
_NON_ALLOW_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.non_allow_sample_rate
|
||||
_BACKPRESSURE_SAMPLE_RATE = AUTH_OBSERVABILITY_RUNTIME_CONFIG.backpressure_sample_rate
|
||||
_BACKPRESSURE_SEVERE_ACTIVE_THRESHOLD = (
|
||||
AUTH_OBSERVABILITY_RUNTIME_CONFIG.backpressure_severe_active_threshold
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class AuthDecisionLogRecord:
|
||||
bot_id: str | None
|
||||
platform: str | None
|
||||
group_id: str | None
|
||||
user_id: str | None
|
||||
module: str | None
|
||||
effect: str
|
||||
reason: str | None = None
|
||||
shadow_effect: str | None = None
|
||||
shadow_reason: str | None = None
|
||||
side_effect_state: dict[str, Any] | None = None
|
||||
latency_ms: float = 0.0
|
||||
overloaded: bool = False
|
||||
|
||||
def to_model(self) -> AuthDecisionLog:
|
||||
return AuthDecisionLog(
|
||||
bot_id=self.bot_id,
|
||||
platform=self.platform,
|
||||
group_id=self.group_id,
|
||||
user_id=self.user_id,
|
||||
module=self.module,
|
||||
effect=self.effect,
|
||||
reason=self.reason,
|
||||
shadow_effect=self.shadow_effect,
|
||||
shadow_reason=self.shadow_reason,
|
||||
side_effect_state=json.dumps(
|
||||
self.side_effect_state,
|
||||
ensure_ascii=False,
|
||||
separators=(",", ":"),
|
||||
)[:4000]
|
||||
if self.side_effect_state
|
||||
else None,
|
||||
latency_ms=self.latency_ms,
|
||||
overloaded=self.overloaded,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RuntimeBackpressureLogRecord:
|
||||
scope_key: str | None
|
||||
reason: str | None
|
||||
lane: str | None
|
||||
action: str
|
||||
queue_size: int = 0
|
||||
active_count: int = 0
|
||||
duration_ms: float = 0.0
|
||||
|
||||
def to_model(self) -> RuntimeBackpressureLog:
|
||||
return RuntimeBackpressureLog(
|
||||
scope_key=self.scope_key,
|
||||
reason=self.reason,
|
||||
lane=self.lane,
|
||||
action=self.action,
|
||||
queue_size=self.queue_size,
|
||||
active_count=self.active_count,
|
||||
duration_ms=self.duration_ms,
|
||||
)
|
||||
|
||||
|
||||
_auth_decision_buffer: deque[AuthDecisionLogRecord] = deque()
|
||||
_backpressure_buffer: deque[RuntimeBackpressureLogRecord] = deque()
|
||||
_buffer_lock = asyncio.Lock()
|
||||
_flush_lock = asyncio.Lock()
|
||||
_flush_task: asyncio.Task[None] | None = None
|
||||
_dropped = 0
|
||||
_last_drop_log_at = 0.0
|
||||
_last_schema_repair_at = 0.0
|
||||
_SCHEMA_REPAIR_INTERVAL_SECONDS = 300.0
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def _ensure_flush_task() -> None:
|
||||
global _flush_task
|
||||
if _flush_task is not None and not _flush_task.done():
|
||||
return
|
||||
_flush_task = asyncio.create_task(_flush_loop())
|
||||
|
||||
|
||||
def _record_drop() -> None:
|
||||
global _dropped, _last_drop_log_at
|
||||
_dropped += 1
|
||||
now = time.monotonic()
|
||||
if now - _last_drop_log_at < _DROP_LOG_INTERVAL_SECONDS:
|
||||
return
|
||||
_last_drop_log_at = now
|
||||
logger.warning(
|
||||
"auth observability buffer full, dropped "
|
||||
f"{_dropped} records, auth_backlog={len(_auth_decision_buffer)}, "
|
||||
f"backpressure_backlog={len(_backpressure_buffer)}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
|
||||
|
||||
def _sample(rate: float) -> bool:
|
||||
if rate >= 1:
|
||||
return True
|
||||
if rate <= 0:
|
||||
return False
|
||||
return random.random() < rate
|
||||
|
||||
|
||||
def _auth_decision_sample_rate(effect: str, overloaded: bool) -> float:
|
||||
if effect != "allow":
|
||||
return _NON_ALLOW_SAMPLE_RATE
|
||||
if overloaded:
|
||||
return _OVERLOADED_ALLOW_SAMPLE_RATE
|
||||
return _ALLOW_SAMPLE_RATE
|
||||
|
||||
|
||||
def _backpressure_sample_rate(record: RuntimeBackpressureLogRecord) -> float:
|
||||
if record.reason and record.reason.startswith("hooks_"):
|
||||
return 1.0
|
||||
if record.active_count >= _BACKPRESSURE_SEVERE_ACTIVE_THRESHOLD:
|
||||
return 1.0
|
||||
if record.action in {"skip", "defer"}:
|
||||
return _BACKPRESSURE_SAMPLE_RATE
|
||||
return min(_BACKPRESSURE_SAMPLE_RATE, 0.02)
|
||||
|
||||
|
||||
async def _append_auth_decision_record(record: AuthDecisionLogRecord) -> None:
|
||||
_ensure_flush_task()
|
||||
async with _buffer_lock:
|
||||
total = len(_auth_decision_buffer) + len(_backpressure_buffer)
|
||||
if total >= _BUFFER_MAX_RETAIN:
|
||||
if len(_auth_decision_buffer) >= len(_backpressure_buffer):
|
||||
with contextlib.suppress(IndexError):
|
||||
_auth_decision_buffer.popleft()
|
||||
else:
|
||||
with contextlib.suppress(IndexError):
|
||||
_backpressure_buffer.popleft()
|
||||
_record_drop()
|
||||
_auth_decision_buffer.append(record)
|
||||
should_flush = (
|
||||
len(_auth_decision_buffer) + len(_backpressure_buffer)
|
||||
>= _FLUSH_TRIGGER_SIZE
|
||||
and not _flush_lock.locked()
|
||||
)
|
||||
if should_flush:
|
||||
# Fire-and-forget keeps auth hot path independent of database stalls.
|
||||
asyncio.create_task(flush_auth_observability_buffer("缓冲区触发")) # noqa: RUF006
|
||||
|
||||
|
||||
async def _append_backpressure_record(record: RuntimeBackpressureLogRecord) -> None:
|
||||
_ensure_flush_task()
|
||||
async with _buffer_lock:
|
||||
total = len(_auth_decision_buffer) + len(_backpressure_buffer)
|
||||
if total >= _BUFFER_MAX_RETAIN:
|
||||
if len(_auth_decision_buffer) >= len(_backpressure_buffer):
|
||||
with contextlib.suppress(IndexError):
|
||||
_auth_decision_buffer.popleft()
|
||||
else:
|
||||
with contextlib.suppress(IndexError):
|
||||
_backpressure_buffer.popleft()
|
||||
_record_drop()
|
||||
_backpressure_buffer.append(record)
|
||||
should_flush = (
|
||||
len(_auth_decision_buffer) + len(_backpressure_buffer)
|
||||
>= _FLUSH_TRIGGER_SIZE
|
||||
and not _flush_lock.locked()
|
||||
)
|
||||
if should_flush:
|
||||
# Fire-and-forget keeps auth hot path independent of database stalls.
|
||||
asyncio.create_task(flush_auth_observability_buffer("缓冲区触发")) # noqa: RUF006
|
||||
|
||||
|
||||
async def append_auth_decision_log(
|
||||
*,
|
||||
bot_id: str | None,
|
||||
platform: str | None,
|
||||
group_id: str | None,
|
||||
user_id: str | None,
|
||||
module: str | None,
|
||||
effect: str,
|
||||
reason: str | None = None,
|
||||
shadow_effect: str | None = None,
|
||||
shadow_reason: str | None = None,
|
||||
side_effect_state: dict[str, Any] | None = None,
|
||||
latency_ms: float = 0.0,
|
||||
overloaded: bool = False,
|
||||
) -> None:
|
||||
if shadow_effect is None and not _sample(
|
||||
_auth_decision_sample_rate(effect, overloaded)
|
||||
):
|
||||
return
|
||||
record = AuthDecisionLogRecord(
|
||||
bot_id=bot_id,
|
||||
platform=platform,
|
||||
group_id=group_id,
|
||||
user_id=user_id,
|
||||
module=module,
|
||||
effect=effect,
|
||||
reason=(reason or "")[:255] or None,
|
||||
shadow_effect=(shadow_effect or "")[:32] or None,
|
||||
shadow_reason=(shadow_reason or "")[:255] or None,
|
||||
side_effect_state=side_effect_state,
|
||||
latency_ms=latency_ms,
|
||||
overloaded=overloaded,
|
||||
)
|
||||
await _append_auth_decision_record(record)
|
||||
|
||||
|
||||
async def append_runtime_backpressure_log(
|
||||
*,
|
||||
scope_key: str | None,
|
||||
reason: str | None,
|
||||
lane: str | None,
|
||||
action: str,
|
||||
queue_size: int = 0,
|
||||
active_count: int = 0,
|
||||
duration_ms: float = 0.0,
|
||||
) -> None:
|
||||
record = RuntimeBackpressureLogRecord(
|
||||
scope_key=(scope_key or "")[:255] or None,
|
||||
reason=(reason or "")[:255] or None,
|
||||
lane=(lane or "")[:64] or None,
|
||||
action=action,
|
||||
queue_size=queue_size,
|
||||
active_count=active_count,
|
||||
duration_ms=duration_ms,
|
||||
)
|
||||
if not _sample(_backpressure_sample_rate(record)):
|
||||
return
|
||||
await _append_backpressure_record(record)
|
||||
|
||||
|
||||
async def _flush_loop() -> None:
|
||||
while True:
|
||||
await asyncio.sleep(_FLUSH_INTERVAL_SECONDS)
|
||||
try:
|
||||
await flush_auth_observability_buffer("定时")
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.warning("定时批量写入权限观测日志失败", LOG_COMMAND, e=exc)
|
||||
|
||||
|
||||
async def _drain_batch(buffer: deque[T]) -> list[T]:
|
||||
batch: list[T] = []
|
||||
async with _buffer_lock:
|
||||
while buffer and len(batch) < _FLUSH_BATCH_SIZE:
|
||||
batch.append(buffer.popleft())
|
||||
return batch
|
||||
|
||||
|
||||
async def _restore_batch(buffer: deque[T], batch: list[T]) -> None:
|
||||
async with _buffer_lock:
|
||||
retain_count = max(_BUFFER_MAX_RETAIN - len(buffer), 0)
|
||||
for record in reversed(batch[-retain_count:]):
|
||||
buffer.appendleft(record)
|
||||
|
||||
|
||||
def _is_schema_mismatch_error(exc: Exception) -> bool:
|
||||
message = str(exc).lower()
|
||||
return any(
|
||||
marker in message
|
||||
for marker in (
|
||||
"no column named",
|
||||
"unknown column",
|
||||
"column does not exist",
|
||||
"no such column",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def _try_repair_auth_schema_once() -> bool:
|
||||
global _last_schema_repair_at
|
||||
now = time.monotonic()
|
||||
if now - _last_schema_repair_at < _SCHEMA_REPAIR_INTERVAL_SECONDS:
|
||||
return False
|
||||
_last_schema_repair_at = now
|
||||
try:
|
||||
from zhenxun.services.db_context.schema_guard import repair_table_schema
|
||||
|
||||
await repair_table_schema("auth_decision_log")
|
||||
await repair_table_schema("runtime_backpressure_log")
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.warning("权限观测日志表结构自修复失败", LOG_COMMAND, e=exc)
|
||||
return False
|
||||
|
||||
|
||||
async def flush_auth_observability_buffer(reason: str) -> int:
|
||||
async with _flush_lock:
|
||||
written = 0
|
||||
while True:
|
||||
auth_batch = await _drain_batch(_auth_decision_buffer)
|
||||
backpressure_batch = await _drain_batch(_backpressure_buffer)
|
||||
if not auth_batch and not backpressure_batch:
|
||||
break
|
||||
try:
|
||||
if auth_batch:
|
||||
await AuthDecisionLog.bulk_create(
|
||||
[record.to_model() for record in auth_batch],
|
||||
_FLUSH_BATCH_SIZE,
|
||||
)
|
||||
written += len(auth_batch)
|
||||
if backpressure_batch:
|
||||
await RuntimeBackpressureLog.bulk_create(
|
||||
[record.to_model() for record in backpressure_batch],
|
||||
_FLUSH_BATCH_SIZE,
|
||||
)
|
||||
written += len(backpressure_batch)
|
||||
except Exception as exc:
|
||||
if _is_schema_mismatch_error(exc):
|
||||
if await _try_repair_auth_schema_once():
|
||||
try:
|
||||
if auth_batch:
|
||||
await AuthDecisionLog.bulk_create(
|
||||
[record.to_model() for record in auth_batch],
|
||||
_FLUSH_BATCH_SIZE,
|
||||
)
|
||||
written += len(auth_batch)
|
||||
if backpressure_batch:
|
||||
await RuntimeBackpressureLog.bulk_create(
|
||||
[
|
||||
record.to_model()
|
||||
for record in backpressure_batch
|
||||
],
|
||||
_FLUSH_BATCH_SIZE,
|
||||
)
|
||||
written += len(backpressure_batch)
|
||||
continue
|
||||
except Exception as retry_exc:
|
||||
exc = retry_exc
|
||||
dropped = len(auth_batch) + len(backpressure_batch)
|
||||
logger.warning(
|
||||
f"{reason}批量写入权限观测日志遇到表结构不匹配,"
|
||||
f"已丢弃低优先级观测日志 {dropped} 条,等待下次启动修复",
|
||||
LOG_COMMAND,
|
||||
e=exc,
|
||||
)
|
||||
return written
|
||||
await _restore_batch(_auth_decision_buffer, auth_batch)
|
||||
await _restore_batch(_backpressure_buffer, backpressure_batch)
|
||||
logger.error(f"{reason}批量写入权限观测日志失败", LOG_COMMAND, e=exc)
|
||||
return written
|
||||
if written:
|
||||
logger.debug(f"{reason}批量写入权限观测日志 {written} 条", LOG_COMMAND)
|
||||
return written
|
||||
|
||||
|
||||
async def stop_auth_observability_buffer() -> int:
|
||||
global _flush_task
|
||||
task = _flush_task
|
||||
_flush_task = None
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
with contextlib.suppress(BaseException):
|
||||
await task
|
||||
return await flush_auth_observability_buffer("关闭")
|
||||
|
||||
|
||||
def _percentile(values: list[float], ratio: float) -> float:
|
||||
if not values:
|
||||
return 0.0
|
||||
ordered = sorted(values)
|
||||
index = min(max(round((len(ordered) - 1) * ratio), 0), len(ordered) - 1)
|
||||
return round(ordered[index], 3)
|
||||
|
||||
|
||||
def _bucket_counts(rows: list[dict[str, Any]], field: str) -> dict[str, int]:
|
||||
counts: dict[str, int] = {}
|
||||
for row in rows:
|
||||
key = str(row.get(field) or "<none>")
|
||||
counts[key] = counts.get(key, 0) + 1
|
||||
return counts
|
||||
|
||||
|
||||
def _lane_budget_advice(rows: list[dict[str, Any]]) -> dict[str, dict[str, Any]]:
|
||||
buckets: dict[str, list[dict[str, Any]]] = {}
|
||||
for row in rows:
|
||||
lane = str(row.get("lane") or "<unknown>")
|
||||
buckets.setdefault(lane, []).append(row)
|
||||
advice: dict[str, dict[str, Any]] = {}
|
||||
for lane, items in buckets.items():
|
||||
if lane == "<unknown>":
|
||||
continue
|
||||
durations = [float(item.get("duration_ms") or 0.0) for item in items]
|
||||
slow_waits = sum(1 for value in durations if value >= 200.0)
|
||||
active_max = max(
|
||||
(int(item.get("active_count") or 0) for item in items), default=0
|
||||
)
|
||||
total = len(items)
|
||||
if not total:
|
||||
continue
|
||||
pressure_ratio = slow_waits / total
|
||||
if pressure_ratio >= 0.2 or active_max >= 5:
|
||||
action = "increase_or_split"
|
||||
elif pressure_ratio == 0 and active_max <= 1 and total >= 20:
|
||||
action = "can_reduce"
|
||||
else:
|
||||
action = "keep"
|
||||
advice[lane] = {
|
||||
"samples": total,
|
||||
"slow_waits": slow_waits,
|
||||
"pressure_ratio": round(pressure_ratio, 3),
|
||||
"active_max": active_max,
|
||||
"p95_duration_ms": _percentile(durations, 0.95),
|
||||
"action": action,
|
||||
}
|
||||
return advice
|
||||
|
||||
|
||||
def _query_placeholder() -> str:
|
||||
try:
|
||||
connection = Tortoise.get_connection("default")
|
||||
if (
|
||||
getattr(connection, "capabilities", None)
|
||||
and getattr(
|
||||
connection.capabilities,
|
||||
"dialect",
|
||||
"",
|
||||
)
|
||||
== "postgres"
|
||||
):
|
||||
return "$1"
|
||||
except Exception:
|
||||
return "?"
|
||||
return "?"
|
||||
|
||||
|
||||
async def build_auth_observability_report(*, hours: float = 24.0) -> dict[str, Any]:
|
||||
since = datetime.now() - timedelta(hours=hours)
|
||||
db = Tortoise.get_connection("default")
|
||||
placeholder = _query_placeholder()
|
||||
auth_rows = await db.execute_query_dict(
|
||||
"SELECT module, effect, reason, shadow_effect, shadow_reason, latency_ms, "
|
||||
f"overloaded FROM auth_decision_log WHERE create_time >= {placeholder} "
|
||||
"ORDER BY create_time DESC LIMIT 100000",
|
||||
[since],
|
||||
)
|
||||
backpressure_rows = await db.execute_query_dict(
|
||||
"SELECT scope_key, lane, reason, action, queue_size, active_count, duration_ms "
|
||||
f"FROM runtime_backpressure_log WHERE create_time >= {placeholder} "
|
||||
"ORDER BY create_time DESC LIMIT 100000",
|
||||
[since],
|
||||
)
|
||||
|
||||
module_buckets: dict[str, list[dict[str, Any]]] = {}
|
||||
for row in auth_rows:
|
||||
module_buckets.setdefault(str(row.get("module") or "<unknown>"), []).append(row)
|
||||
module_stats: list[dict[str, Any]] = []
|
||||
for module, items in module_buckets.items():
|
||||
latencies = [float(item.get("latency_ms") or 0.0) for item in items]
|
||||
module_stats.append(
|
||||
{
|
||||
"module": module,
|
||||
"total": len(items),
|
||||
"effects": _bucket_counts(items, "effect"),
|
||||
"shadow_effects": _bucket_counts(items, "shadow_effect"),
|
||||
"avg_latency_ms": round(sum(latencies) / len(latencies), 3)
|
||||
if latencies
|
||||
else 0.0,
|
||||
"p95_latency_ms": _percentile(latencies, 0.95),
|
||||
"overloaded": sum(1 for item in items if bool(item.get("overloaded"))),
|
||||
}
|
||||
)
|
||||
|
||||
backpressure_buckets: dict[str, list[dict[str, Any]]] = {}
|
||||
for row in backpressure_rows:
|
||||
key = f"{row.get('lane') or '<unknown>'}:{row.get('reason') or '<none>'}"
|
||||
backpressure_buckets.setdefault(key, []).append(row)
|
||||
backpressure_stats: list[dict[str, Any]] = []
|
||||
for key, items in backpressure_buckets.items():
|
||||
durations = [float(item.get("duration_ms") or 0.0) for item in items]
|
||||
backpressure_stats.append(
|
||||
{
|
||||
"key": key,
|
||||
"total": len(items),
|
||||
"actions": _bucket_counts(items, "action"),
|
||||
"avg_duration_ms": round(sum(durations) / len(durations), 3)
|
||||
if durations
|
||||
else 0.0,
|
||||
"p95_duration_ms": _percentile(durations, 0.95),
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"created_at": datetime.now().isoformat(timespec="seconds"),
|
||||
"window_hours": hours,
|
||||
"auth_decisions": {
|
||||
"total": len(auth_rows),
|
||||
"effects": _bucket_counts(auth_rows, "effect"),
|
||||
"shadow_effects": _bucket_counts(auth_rows, "shadow_effect"),
|
||||
"top_modules_by_p95": sorted(
|
||||
module_stats,
|
||||
key=lambda item: (item["p95_latency_ms"], item["total"]),
|
||||
reverse=True,
|
||||
)[:30],
|
||||
},
|
||||
"backpressure": {
|
||||
"total": len(backpressure_rows),
|
||||
"lane_budget_advice": _lane_budget_advice(backpressure_rows),
|
||||
"top_reasons": sorted(
|
||||
backpressure_stats,
|
||||
key=lambda item: (item["total"], item["p95_duration_ms"]),
|
||||
reverse=True,
|
||||
)[:30],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@PriorityLifecycle.on_shutdown(priority=90)
|
||||
async def _flush_auth_observability_buffer_on_shutdown() -> None:
|
||||
await stop_auth_observability_buffer()
|
||||
+378
-72
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from contextvars import ContextVar
|
||||
from dataclasses import dataclass, field
|
||||
import json
|
||||
import os
|
||||
@@ -35,6 +36,8 @@ def _coerce_int(value, default: int) -> int:
|
||||
|
||||
# RuntimeCache 以模型 save/delete 主动失效为主,周期 refresh 只做兜底。
|
||||
# 这些默认值避免低压力运行时频繁全量扫表。
|
||||
# 权限检查热路径以 RuntimeCache/AuthSnapshot 为唯一数据入口;普通业务的
|
||||
# DataAccess/CacheRoot 缓存不能替代这里的运行态快照。
|
||||
PLUGININFO_MEM_REFRESH_INTERVAL = 3600 # 60分钟
|
||||
BAN_MEM_REFRESH_INTERVAL = 300
|
||||
BAN_MEM_CLEAN_INTERVAL = 60
|
||||
@@ -52,10 +55,16 @@ LIMIT_MEM_REFRESH_INTERVAL = 900 # 15分钟
|
||||
LIMIT_MEM_NEGATIVE_TTL = 30
|
||||
RUNTIME_CACHE_SYNC_ENABLED = True
|
||||
RUNTIME_CACHE_SYNC_CHANNEL = "ZHENXUN_RUNTIME_CACHE_SYNC"
|
||||
RUNTIME_CACHE_LOAD_RETRY_SECONDS = 1.0
|
||||
RUNTIME_CACHE_STARTUP_REFRESH_SKIP_SECONDS = 5.0
|
||||
|
||||
|
||||
INSTANCE_ID = uuid.uuid4().hex
|
||||
_CACHE_READY_EVENT = asyncio.Event()
|
||||
_APPLYING_REMOTE_CACHE_EVENT: ContextVar[bool] = ContextVar(
|
||||
"APPLYING_REMOTE_RUNTIME_CACHE_EVENT",
|
||||
default=False,
|
||||
)
|
||||
|
||||
|
||||
def _env_get(name: str, default: str | None = None) -> str | None:
|
||||
@@ -180,6 +189,65 @@ class PluginInfoSnapshot:
|
||||
plugin._saved_in_db = True
|
||||
return plugin
|
||||
|
||||
def to_payload(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"module": self.module,
|
||||
"module_path": self.module_path,
|
||||
"name": self.name,
|
||||
"status": self.status,
|
||||
"block_type": self.block_type.value if self.block_type else None,
|
||||
"load_status": self.load_status,
|
||||
"author": self.author,
|
||||
"version": self.version,
|
||||
"level": self.level,
|
||||
"default_status": self.default_status,
|
||||
"limit_superuser": self.limit_superuser,
|
||||
"menu_type": self.menu_type,
|
||||
"plugin_type": self.plugin_type.value if self.plugin_type else None,
|
||||
"cost_gold": self.cost_gold,
|
||||
"admin_level": self.admin_level,
|
||||
"ignore_prompt": self.ignore_prompt,
|
||||
"is_delete": self.is_delete,
|
||||
"parent": self.parent,
|
||||
"is_show": self.is_show,
|
||||
"ignore_statistics": self.ignore_statistics,
|
||||
"impression": self.impression,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_payload(cls, payload: dict[str, Any]) -> "PluginInfoSnapshot":
|
||||
block_type = payload.get("block_type")
|
||||
if block_type is not None and not isinstance(block_type, BlockType):
|
||||
block_type = BlockType(block_type)
|
||||
plugin_type = payload.get("plugin_type")
|
||||
if plugin_type is not None and not isinstance(plugin_type, PluginType):
|
||||
plugin_type = PluginType(plugin_type)
|
||||
return cls(
|
||||
id=int(payload.get("id", 0) or 0),
|
||||
module=str(payload.get("module", "") or ""),
|
||||
module_path=str(payload.get("module_path", "") or ""),
|
||||
name=str(payload.get("name", "") or ""),
|
||||
status=bool(payload.get("status", True)),
|
||||
block_type=block_type,
|
||||
load_status=bool(payload.get("load_status", True)),
|
||||
author=payload.get("author"),
|
||||
version=payload.get("version"),
|
||||
level=int(payload.get("level", 0) or 0),
|
||||
default_status=bool(payload.get("default_status", True)),
|
||||
limit_superuser=bool(payload.get("limit_superuser", False)),
|
||||
menu_type=str(payload.get("menu_type", "") or ""),
|
||||
plugin_type=plugin_type,
|
||||
cost_gold=int(payload.get("cost_gold", 0) or 0),
|
||||
admin_level=payload.get("admin_level"),
|
||||
ignore_prompt=bool(payload.get("ignore_prompt", False)),
|
||||
is_delete=bool(payload.get("is_delete", False)),
|
||||
parent=payload.get("parent"),
|
||||
is_show=bool(payload.get("is_show", True)),
|
||||
ignore_statistics=bool(payload.get("ignore_statistics", False)),
|
||||
impression=float(payload.get("impression", 0) or 0),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BanEntry:
|
||||
@@ -648,18 +716,87 @@ class RuntimeCacheSync:
|
||||
cache_type = payload.get("type")
|
||||
action = payload.get("action")
|
||||
data = payload.get("data") or {}
|
||||
if cache_type == "bot":
|
||||
await BotMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "group":
|
||||
await GroupMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "ban":
|
||||
await BanMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "level":
|
||||
await LevelUserMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "task":
|
||||
await TaskInfoMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "plugin_limit":
|
||||
await PluginLimitMemoryCache.apply_sync_event(action, data)
|
||||
token = _APPLYING_REMOTE_CACHE_EVENT.set(True)
|
||||
try:
|
||||
if cache_type == "bot":
|
||||
await BotMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "group":
|
||||
await GroupMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "ban":
|
||||
await BanMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "level":
|
||||
await LevelUserMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "task":
|
||||
await TaskInfoMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "plugin_limit":
|
||||
await PluginLimitMemoryCache.apply_sync_event(action, data)
|
||||
elif cache_type == "plugin":
|
||||
await PluginInfoMemoryCache.apply_sync_event(action, data)
|
||||
finally:
|
||||
_APPLYING_REMOTE_CACHE_EVENT.reset(token)
|
||||
|
||||
|
||||
class RuntimeCacheMutation:
|
||||
"""Small helpers for runtime cache mutation bookkeeping.
|
||||
|
||||
Cache classes still own their storage layout. This helper centralizes the
|
||||
shared mutation side effects: health markers, negative-cache cleanup and
|
||||
cross-process publish.
|
||||
"""
|
||||
|
||||
_load_locks: ClassVar[dict[str, asyncio.Lock]] = {}
|
||||
_retry_after: ClassVar[dict[str, float]] = {}
|
||||
|
||||
@classmethod
|
||||
async def ensure_loaded(cls, cache_cls: type, label: str) -> None:
|
||||
if getattr(cache_cls, "_loaded", False):
|
||||
return
|
||||
now = time.monotonic()
|
||||
if cls._retry_after.get(label, 0.0) > now:
|
||||
return
|
||||
lock = cls._load_locks.setdefault(label, asyncio.Lock())
|
||||
async with lock:
|
||||
if getattr(cache_cls, "_loaded", False):
|
||||
return
|
||||
now = time.monotonic()
|
||||
if cls._retry_after.get(label, 0.0) > now:
|
||||
return
|
||||
try:
|
||||
await cache_cls.refresh()
|
||||
except Exception as exc:
|
||||
cls.mark_error(cache_cls, exc)
|
||||
cls._retry_after[label] = (
|
||||
time.monotonic() + RUNTIME_CACHE_LOAD_RETRY_SECONDS
|
||||
)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def mark_refreshed(cache_cls: type) -> None:
|
||||
setattr(cache_cls, "_loaded", True)
|
||||
setattr(cache_cls, "_last_refresh", time.time())
|
||||
setattr(cache_cls, "_last_error", None)
|
||||
|
||||
@staticmethod
|
||||
def mark_error(cache_cls: type, exc: Exception) -> None:
|
||||
setattr(cache_cls, "_last_error", f"{type(exc).__name__}: {exc}")
|
||||
|
||||
@staticmethod
|
||||
def clear_negative_key(cache_cls: type, key: object) -> None:
|
||||
negative = getattr(cache_cls, "_negative", None)
|
||||
if isinstance(negative, dict):
|
||||
negative.pop(key, None)
|
||||
|
||||
@staticmethod
|
||||
def clear_negative_all(cache_cls: type) -> None:
|
||||
negative = getattr(cache_cls, "_negative", None)
|
||||
if isinstance(negative, dict):
|
||||
negative.clear()
|
||||
|
||||
@staticmethod
|
||||
def publish(cache_type: str, action: str, data: dict[str, Any]) -> None:
|
||||
if _APPLYING_REMOTE_CACHE_EVENT.get():
|
||||
return
|
||||
RuntimeCacheSync.publish_event(cache_type, action, data)
|
||||
|
||||
|
||||
class PluginInfoMemoryCache:
|
||||
@@ -669,6 +806,7 @@ class PluginInfoMemoryCache:
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _to_model(cls, snapshot: PluginInfoSnapshot | None) -> "PluginInfo | None":
|
||||
@@ -687,6 +825,10 @@ class PluginInfoMemoryCache:
|
||||
cls._by_module.pop(old.module, None)
|
||||
cls._by_module_path[snapshot.module_path] = snapshot
|
||||
|
||||
@staticmethod
|
||||
def _module_snapshot_rank(snapshot: PluginInfoSnapshot) -> tuple[int, int]:
|
||||
return (1 if snapshot.load_status else 0, snapshot.id)
|
||||
|
||||
@classmethod
|
||||
async def refresh(cls) -> None:
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
@@ -698,13 +840,16 @@ class PluginInfoMemoryCache:
|
||||
for plugin in plugins:
|
||||
snapshot = PluginInfoSnapshot.from_model(plugin)
|
||||
if snapshot.module:
|
||||
by_module[snapshot.module] = snapshot
|
||||
current = by_module.get(snapshot.module)
|
||||
if current is None or cls._module_snapshot_rank(
|
||||
snapshot
|
||||
) >= cls._module_snapshot_rank(current):
|
||||
by_module[snapshot.module] = snapshot
|
||||
if snapshot.module_path:
|
||||
by_module_path[snapshot.module_path] = snapshot
|
||||
cls._by_module = by_module
|
||||
cls._by_module_path = by_module_path
|
||||
cls._loaded = True
|
||||
cls._last_refresh = time.time()
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(
|
||||
f"plugin cache refreshed: {len(by_module)} entries", LOG_COMMAND
|
||||
)
|
||||
@@ -713,7 +858,7 @@ class PluginInfoMemoryCache:
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "plugin")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -749,8 +894,6 @@ class PluginInfoMemoryCache:
|
||||
return
|
||||
snapshot = PluginInfoSnapshot.from_model(plugin)
|
||||
cls._store_snapshot(snapshot)
|
||||
cls._loaded = True
|
||||
cls._last_refresh = time.time()
|
||||
|
||||
@classmethod
|
||||
def remove_by_module(cls, module: str) -> None:
|
||||
@@ -765,8 +908,16 @@ class PluginInfoMemoryCache:
|
||||
async with cls._lock:
|
||||
snapshot = PluginInfoSnapshot.from_model(plugin)
|
||||
cls._store_snapshot(snapshot)
|
||||
cls._loaded = True
|
||||
cls._last_refresh = time.time()
|
||||
RuntimeCacheMutation.publish("plugin", "upsert", snapshot.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
try:
|
||||
snapshot = PluginInfoSnapshot.from_payload(payload)
|
||||
except Exception:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._store_snapshot(snapshot)
|
||||
|
||||
@classmethod
|
||||
async def remove(
|
||||
@@ -783,6 +934,20 @@ class PluginInfoMemoryCache:
|
||||
snapshot = cls._by_module_path.pop(module_path, None)
|
||||
if snapshot and snapshot.module:
|
||||
cls._by_module.pop(snapshot.module, None)
|
||||
RuntimeCacheMutation.publish(
|
||||
"plugin",
|
||||
"delete",
|
||||
{"module": module, "module_path": module_path},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
if action == "upsert":
|
||||
await cls.upsert_from_payload(data)
|
||||
elif action == "delete":
|
||||
await cls.remove(data.get("module"), data.get("module_path"))
|
||||
elif action == "refresh":
|
||||
await cls.refresh()
|
||||
|
||||
@classmethod
|
||||
async def _refresh_loop(cls, interval: int) -> None:
|
||||
@@ -815,6 +980,8 @@ class BotMemoryCache:
|
||||
_negative: ClassVar[dict[str, float]] = {}
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize(cls, bot_id: str | None) -> str | None:
|
||||
@@ -833,7 +1000,7 @@ class BotMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(bot_id, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, bot_id)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -851,15 +1018,15 @@ class BotMemoryCache:
|
||||
async with cls._lock:
|
||||
records = await BotConsole.all()
|
||||
cls._by_id = {str(r.bot_id): BotSnapshot.from_model(r) for r in records}
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(f"bot cache refreshed: {len(cls._by_id)} entries", LOG_COMMAND)
|
||||
|
||||
@classmethod
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "bot")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -918,15 +1085,16 @@ class BotMemoryCache:
|
||||
available_tasks=entry.available_tasks,
|
||||
)
|
||||
cls._by_id[bot_id] = updated
|
||||
RuntimeCacheSync.publish_event("bot", "upsert", updated.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, bot_id)
|
||||
RuntimeCacheMutation.publish("bot", "upsert", updated.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_model(cls, record) -> None:
|
||||
entry = BotSnapshot.from_model(record)
|
||||
async with cls._lock:
|
||||
cls._by_id[entry.bot_id] = entry
|
||||
cls._negative.pop(entry.bot_id, None)
|
||||
RuntimeCacheSync.publish_event("bot", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.bot_id)
|
||||
RuntimeCacheMutation.publish("bot", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -935,7 +1103,7 @@ class BotMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_id[entry.bot_id] = entry
|
||||
cls._negative.pop(entry.bot_id, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.bot_id)
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, bot_id: str | None) -> None:
|
||||
@@ -944,7 +1112,8 @@ class BotMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_id.pop(bot_id, None)
|
||||
RuntimeCacheSync.publish_event("bot", "delete", {"bot_id": bot_id})
|
||||
RuntimeCacheMutation.clear_negative_key(cls, bot_id)
|
||||
RuntimeCacheMutation.publish("bot", "delete", {"bot_id": bot_id})
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
@@ -986,6 +1155,8 @@ class GroupMemoryCache:
|
||||
_negative: ClassVar[dict[tuple[str, str], float]] = {}
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize(cls, value: str | None) -> str | None:
|
||||
@@ -1014,7 +1185,7 @@ class GroupMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1038,15 +1209,15 @@ class GroupMemoryCache:
|
||||
if key:
|
||||
by_key[key] = entry
|
||||
cls._by_key = by_key
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(f"group cache refreshed: {len(by_key)} entries", LOG_COMMAND)
|
||||
|
||||
@classmethod
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "group")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -1094,8 +1265,8 @@ class GroupMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_key[key] = entry
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheSync.publish_event("group", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
RuntimeCacheMutation.publish("group", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -1105,7 +1276,7 @@ class GroupMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_key[key] = entry
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, group_id: str | None, channel_id: str | None = None) -> None:
|
||||
@@ -1114,7 +1285,8 @@ class GroupMemoryCache:
|
||||
return
|
||||
async with cls._lock:
|
||||
cls._by_key.pop(key, None)
|
||||
RuntimeCacheSync.publish_event(
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
RuntimeCacheMutation.publish(
|
||||
"group", "delete", {"group_id": key[0], "channel_id": key[1] or None}
|
||||
)
|
||||
|
||||
@@ -1186,7 +1358,7 @@ class LevelUserMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1215,16 +1387,15 @@ class LevelUserMemoryCache:
|
||||
by_user_max[entry.user_id] = entry.user_level
|
||||
cls._by_key = by_key
|
||||
cls._by_user_max = by_user_max
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
cls._last_refresh = time.time()
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(f"level cache refreshed: {len(by_key)} entries", LOG_COMMAND)
|
||||
|
||||
@classmethod
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "level")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -1310,13 +1481,13 @@ class LevelUserMemoryCache:
|
||||
async with cls._lock:
|
||||
prev = cls._by_key.get(key)
|
||||
cls._by_key[key] = entry
|
||||
cls._negative.pop(key, None)
|
||||
current = cls._by_user_max.get(entry.user_id, 0)
|
||||
if entry.user_level >= current:
|
||||
cls._by_user_max[entry.user_id] = entry.user_level
|
||||
elif prev and prev.user_level == current and entry.user_level < current:
|
||||
cls._recalc_user_max(entry.user_id)
|
||||
RuntimeCacheSync.publish_event("level", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
RuntimeCacheMutation.publish("level", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -1327,12 +1498,12 @@ class LevelUserMemoryCache:
|
||||
async with cls._lock:
|
||||
prev = cls._by_key.get(key)
|
||||
cls._by_key[key] = entry
|
||||
cls._negative.pop(key, None)
|
||||
current = cls._by_user_max.get(entry.user_id, 0)
|
||||
if entry.user_level >= current:
|
||||
cls._by_user_max[entry.user_id] = entry.user_level
|
||||
elif prev and prev.user_level == current and entry.user_level < current:
|
||||
cls._recalc_user_max(entry.user_id)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, user_id: str | None, group_id: str | None) -> None:
|
||||
@@ -1343,7 +1514,8 @@ class LevelUserMemoryCache:
|
||||
removed = cls._by_key.pop(key, None)
|
||||
if removed and cls._by_user_max.get(removed.user_id) == removed.user_level:
|
||||
cls._recalc_user_max(removed.user_id)
|
||||
RuntimeCacheSync.publish_event(
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
RuntimeCacheMutation.publish(
|
||||
"level", "delete", {"user_id": key[0], "group_id": key[1] or None}
|
||||
)
|
||||
|
||||
@@ -1399,6 +1571,8 @@ class TaskInfoMemoryCache:
|
||||
_negative: ClassVar[dict[str, float]] = {}
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize(cls, module: str | None) -> str | None:
|
||||
@@ -1417,7 +1591,7 @@ class TaskInfoMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(module, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, module)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1443,8 +1617,8 @@ class TaskInfoMemoryCache:
|
||||
by_name[entry.name] = entry
|
||||
cls._by_module = by_module
|
||||
cls._by_name = by_name
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(
|
||||
f"task info cache refreshed: {len(cls._by_module)} entries",
|
||||
LOG_COMMAND,
|
||||
@@ -1454,7 +1628,7 @@ class TaskInfoMemoryCache:
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "task")
|
||||
|
||||
@classmethod
|
||||
async def get(cls, module: str | None) -> TaskInfoSnapshot | None:
|
||||
@@ -1488,10 +1662,21 @@ class TaskInfoMemoryCache:
|
||||
|
||||
@classmethod
|
||||
async def is_disabled(cls, module: str | None) -> bool:
|
||||
"""Backward-compatible runtime disabled check for passive tasks."""
|
||||
return await cls.is_runtime_disabled(module)
|
||||
|
||||
@classmethod
|
||||
async def is_runtime_disabled(cls, module: str | None) -> bool:
|
||||
"""Return whether a passive task is unavailable at runtime.
|
||||
|
||||
Runtime passive availability is defined by TaskInfo.status and
|
||||
TaskInfo.load_status. Bot/group scoped block lists are checked by
|
||||
CommonUtils.task_is_block().
|
||||
"""
|
||||
entry = await cls.get(module)
|
||||
if not entry:
|
||||
return False
|
||||
return not entry.status
|
||||
return not entry.status or not entry.load_status
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_model(cls, record) -> None:
|
||||
@@ -1500,8 +1685,8 @@ class TaskInfoMemoryCache:
|
||||
cls._by_module[entry.module] = entry
|
||||
if entry.name:
|
||||
cls._by_name[entry.name] = entry
|
||||
cls._negative.pop(entry.module, None)
|
||||
RuntimeCacheSync.publish_event("task", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
RuntimeCacheMutation.publish("task", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -1512,7 +1697,7 @@ class TaskInfoMemoryCache:
|
||||
cls._by_module[entry.module] = entry
|
||||
if entry.name:
|
||||
cls._by_name[entry.name] = entry
|
||||
cls._negative.pop(entry.module, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, module: str | None) -> None:
|
||||
@@ -1525,7 +1710,8 @@ class TaskInfoMemoryCache:
|
||||
current = cls._by_name.get(removed.name)
|
||||
if current and current.module == removed.module:
|
||||
cls._by_name.pop(removed.name, None)
|
||||
RuntimeCacheSync.publish_event("task", "delete", {"module": module})
|
||||
RuntimeCacheMutation.clear_negative_key(cls, module)
|
||||
RuntimeCacheMutation.publish("task", "delete", {"module": module})
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
@@ -1568,6 +1754,8 @@ class PluginLimitMemoryCache:
|
||||
_negative: ClassVar[dict[str, float]] = {}
|
||||
_loaded: ClassVar[bool] = False
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize(cls, value: str | None) -> str | None:
|
||||
@@ -1586,7 +1774,7 @@ class PluginLimitMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(module, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, module)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1611,8 +1799,8 @@ class PluginLimitMemoryCache:
|
||||
by_module.setdefault(entry.module, []).append(entry)
|
||||
cls._by_id = by_id
|
||||
cls._by_module = by_module
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(
|
||||
f"plugin limit cache refreshed: {len(by_id)} entries",
|
||||
LOG_COMMAND,
|
||||
@@ -1622,7 +1810,7 @@ class PluginLimitMemoryCache:
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "plugin_limit")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -1668,7 +1856,7 @@ class PluginLimitMemoryCache:
|
||||
async def upsert_from_model(cls, record) -> None:
|
||||
entry = PluginLimitSnapshot.from_model(record)
|
||||
await cls._upsert_entry(entry)
|
||||
RuntimeCacheSync.publish_event("plugin_limit", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.publish("plugin_limit", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def upsert_from_payload(cls, payload: dict[str, Any]) -> None:
|
||||
@@ -1695,6 +1883,7 @@ class PluginLimitMemoryCache:
|
||||
for item in cls._by_module.get(entry.module, [])
|
||||
if item.id != entry.id
|
||||
]
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
return
|
||||
cls._by_id[entry.id] = entry
|
||||
module_limits = [
|
||||
@@ -1704,7 +1893,7 @@ class PluginLimitMemoryCache:
|
||||
]
|
||||
module_limits.append(entry)
|
||||
cls._by_module[entry.module] = module_limits
|
||||
cls._negative.pop(entry.module, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
|
||||
@classmethod
|
||||
async def remove_by_id(cls, limit_id: int | None) -> None:
|
||||
@@ -1718,7 +1907,8 @@ class PluginLimitMemoryCache:
|
||||
for item in cls._by_module.get(entry.module, [])
|
||||
if item.id != entry.id
|
||||
]
|
||||
RuntimeCacheSync.publish_event("plugin_limit", "delete", {"id": int(limit_id)})
|
||||
RuntimeCacheMutation.clear_negative_key(cls, entry.module)
|
||||
RuntimeCacheMutation.publish("plugin_limit", "delete", {"id": int(limit_id)})
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
@@ -1764,6 +1954,8 @@ class BanMemoryCache:
|
||||
_refresh_task: ClassVar[asyncio.Task | None] = None
|
||||
_cleanup_task: ClassVar[asyncio.Task | None] = None
|
||||
_remove_tasks: ClassVar[set[asyncio.Task]] = set()
|
||||
_last_refresh: ClassVar[float] = 0.0
|
||||
_last_error: ClassVar[str | None] = None
|
||||
|
||||
@classmethod
|
||||
def _normalize_id(cls, value: str | None) -> str | None:
|
||||
@@ -1788,7 +1980,7 @@ class BanMemoryCache:
|
||||
if not expire_at:
|
||||
return False
|
||||
if expire_at <= time.time():
|
||||
cls._negative.pop(key, None)
|
||||
RuntimeCacheMutation.clear_negative_key(cls, key)
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -1842,8 +2034,8 @@ class BanMemoryCache:
|
||||
cls._by_user = by_user
|
||||
cls._by_group = by_group
|
||||
cls._by_user_group = by_user_group
|
||||
cls._negative = {}
|
||||
cls._loaded = True
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.mark_refreshed(cls)
|
||||
logger.debug(
|
||||
"ban cache refreshed: "
|
||||
f"user={len(by_user)} group={len(by_group)} "
|
||||
@@ -1855,7 +2047,7 @@ class BanMemoryCache:
|
||||
async def ensure_loaded(cls) -> None:
|
||||
if cls._loaded:
|
||||
return
|
||||
await cls.refresh()
|
||||
await RuntimeCacheMutation.ensure_loaded(cls, "ban")
|
||||
|
||||
@classmethod
|
||||
def is_loaded(cls) -> bool:
|
||||
@@ -1873,13 +2065,13 @@ class BanMemoryCache:
|
||||
cls._by_user[entry.user_id] = entry
|
||||
elif entry.group_id:
|
||||
cls._by_group[entry.group_id] = entry
|
||||
cls._negative = {}
|
||||
RuntimeCacheSync.publish_event("ban", "upsert", entry.to_payload())
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
RuntimeCacheMutation.publish("ban", "upsert", entry.to_payload())
|
||||
|
||||
@classmethod
|
||||
async def remove(cls, user_id: str | None, group_id: str | None) -> None:
|
||||
await cls._remove_local(user_id, group_id)
|
||||
RuntimeCacheSync.publish_event(
|
||||
RuntimeCacheMutation.publish(
|
||||
"ban", "delete", {"user_id": user_id, "group_id": group_id}
|
||||
)
|
||||
|
||||
@@ -1894,7 +2086,7 @@ class BanMemoryCache:
|
||||
cls._by_user.pop(user_id, None)
|
||||
elif group_id:
|
||||
cls._by_group.pop(group_id, None)
|
||||
cls._negative = {}
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
|
||||
@classmethod
|
||||
def _get_entry(cls, user_id: str | None, group_id: str | None) -> BanEntry | None:
|
||||
@@ -1995,7 +2187,7 @@ class BanMemoryCache:
|
||||
elif entry.group_id:
|
||||
cls._by_group.pop(entry.group_id, None)
|
||||
if expired:
|
||||
cls._negative = {}
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
if not delete_db or not expired:
|
||||
return
|
||||
from tortoise.expressions import Q
|
||||
@@ -2064,7 +2256,7 @@ class BanMemoryCache:
|
||||
cls._by_user[entry.user_id] = entry
|
||||
elif entry.group_id:
|
||||
cls._by_group[entry.group_id] = entry
|
||||
cls._negative = {}
|
||||
RuntimeCacheMutation.clear_negative_all(cls)
|
||||
|
||||
@classmethod
|
||||
async def apply_sync_event(cls, action: str, data: dict[str, Any]) -> None:
|
||||
@@ -2078,12 +2270,126 @@ class BanMemoryCache:
|
||||
|
||||
async def _safe_refresh(cache_cls: type, label: str) -> None:
|
||||
"""安全地刷新单个缓存,异常不影响其他缓存。"""
|
||||
if getattr(cache_cls, "_loaded", False):
|
||||
last_refresh = float(getattr(cache_cls, "_last_refresh", 0.0) or 0.0)
|
||||
if time.time() - last_refresh <= RUNTIME_CACHE_STARTUP_REFRESH_SKIP_SECONDS:
|
||||
logger.debug(f"{label} cache startup refresh skipped", LOG_COMMAND)
|
||||
return
|
||||
try:
|
||||
await cache_cls.refresh()
|
||||
except Exception as exc:
|
||||
RuntimeCacheMutation.mark_error(cache_cls, exc)
|
||||
logger.error(f"{label} cache init failed", LOG_COMMAND, e=exc)
|
||||
|
||||
|
||||
def _cache_health(
|
||||
cache_cls: type,
|
||||
*,
|
||||
entry_count: int,
|
||||
negative_count: int = 0,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"loaded": bool(getattr(cache_cls, "_loaded", False)),
|
||||
"entry_count": entry_count,
|
||||
"last_refresh": float(getattr(cache_cls, "_last_refresh", 0.0) or 0.0),
|
||||
"negative_count": negative_count,
|
||||
"last_error": getattr(cache_cls, "_last_error", None),
|
||||
}
|
||||
|
||||
|
||||
def health_snapshot() -> dict[str, dict[str, Any]]:
|
||||
"""Return in-memory runtime cache health without touching the database."""
|
||||
return {
|
||||
"plugin": _cache_health(
|
||||
PluginInfoMemoryCache,
|
||||
entry_count=len(PluginInfoMemoryCache._by_module),
|
||||
),
|
||||
"bot": _cache_health(
|
||||
BotMemoryCache,
|
||||
entry_count=len(BotMemoryCache._by_id),
|
||||
negative_count=len(BotMemoryCache._negative),
|
||||
),
|
||||
"group": _cache_health(
|
||||
GroupMemoryCache,
|
||||
entry_count=len(GroupMemoryCache._by_key),
|
||||
negative_count=len(GroupMemoryCache._negative),
|
||||
),
|
||||
"level": _cache_health(
|
||||
LevelUserMemoryCache,
|
||||
entry_count=len(LevelUserMemoryCache._by_key),
|
||||
negative_count=len(LevelUserMemoryCache._negative),
|
||||
),
|
||||
"task": _cache_health(
|
||||
TaskInfoMemoryCache,
|
||||
entry_count=len(TaskInfoMemoryCache._by_module),
|
||||
negative_count=len(TaskInfoMemoryCache._negative),
|
||||
),
|
||||
"plugin_limit": _cache_health(
|
||||
PluginLimitMemoryCache,
|
||||
entry_count=len(PluginLimitMemoryCache._by_id),
|
||||
negative_count=len(PluginLimitMemoryCache._negative),
|
||||
),
|
||||
"ban": _cache_health(
|
||||
BanMemoryCache,
|
||||
entry_count=(
|
||||
len(BanMemoryCache._by_user)
|
||||
+ len(BanMemoryCache._by_group)
|
||||
+ len(BanMemoryCache._by_user_group)
|
||||
),
|
||||
negative_count=len(BanMemoryCache._negative),
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def passive_status_snapshot(max_modules: int = 50) -> dict[str, Any]:
|
||||
"""Return passive-task state from in-memory caches only.
|
||||
|
||||
This is a local diagnostic helper: it does not query or write the database,
|
||||
and it is not used by runtime decisions.
|
||||
"""
|
||||
tasks = list(TaskInfoMemoryCache._by_module.values())
|
||||
disabled = sorted(task.module for task in tasks if not task.status)
|
||||
unloaded = sorted(task.module for task in tasks if not task.load_status)
|
||||
runtime_enabled = [
|
||||
task.module for task in tasks if task.status and task.load_status
|
||||
]
|
||||
bot_block_total = sum(
|
||||
len(_parse_block_modules(bot.block_tasks))
|
||||
for bot in BotMemoryCache._by_id.values()
|
||||
)
|
||||
group_block_total = sum(
|
||||
len(group.block_task_set) + len(group.superuser_block_task_set)
|
||||
for group in GroupMemoryCache._by_key.values()
|
||||
)
|
||||
return {
|
||||
"cache": health_snapshot(),
|
||||
"passive_tasks": {
|
||||
"total": len(tasks),
|
||||
"status_enabled": sum(1 for task in tasks if task.status),
|
||||
"load_status_enabled": sum(1 for task in tasks if task.load_status),
|
||||
"runtime_enabled": len(runtime_enabled),
|
||||
"disabled_modules": disabled[:max_modules],
|
||||
"disabled_modules_total": len(disabled),
|
||||
"unloaded_modules": unloaded[:max_modules],
|
||||
"unloaded_modules_total": len(unloaded),
|
||||
},
|
||||
"scoped_blocks": {
|
||||
"bot_block_tasks_total": bot_block_total,
|
||||
"group_block_tasks_total": group_block_total,
|
||||
},
|
||||
"semantics": {
|
||||
"available_tasks": "management_display_mirror_not_runtime_whitelist",
|
||||
"runtime_truth": [
|
||||
"TaskInfo.status",
|
||||
"TaskInfo.load_status",
|
||||
"BotConsole.block_tasks",
|
||||
"GroupConsole.block_task",
|
||||
"GroupConsole.superuser_block_task",
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=6)
|
||||
async def _init_runtime_cache():
|
||||
await RuntimeCacheSync.start()
|
||||
|
||||
@@ -12,6 +12,10 @@ T = TypeVar("T", bound=Model)
|
||||
class DataAccess(Generic[T]):
|
||||
"""数据访问兼容层,根据配置保留单点缓存读取和清理能力
|
||||
|
||||
边界说明:DataAccess 面向普通业务查询和低频管理链路。权限检查热路径
|
||||
必须优先使用 RuntimeCache/AuthSnapshot,避免高并发消息处理时触发
|
||||
DB/cache-aside 读放大。
|
||||
|
||||
新的高频运行态路径应优先使用 RuntimeCache 或 BoundedTTLCache。
|
||||
这里不再把 filter/all/create/update_or_create 结果写入通用缓存,
|
||||
create/update_or_create 只负责清理旧缓存,避免旧值残留。
|
||||
|
||||
@@ -5,6 +5,9 @@ from pydantic import BaseModel
|
||||
# 数据库操作超时设置(秒)
|
||||
DB_TIMEOUT_SECONDS = 3.0
|
||||
|
||||
# 启动期自动补齐字段/索引可能需要等待数据库锁或扫描较大的表,单独放宽超时
|
||||
DB_SCHEMA_GUARD_TIMEOUT_SECONDS = 30.0
|
||||
|
||||
# 性能监控阈值(秒)
|
||||
SLOW_QUERY_THRESHOLD = 0.5
|
||||
|
||||
|
||||
@@ -13,7 +13,7 @@ from tortoise.exceptions import OperationalError
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import DB_TIMEOUT_SECONDS, LOG_COMMAND
|
||||
from .config import DB_SCHEMA_GUARD_TIMEOUT_SECONDS, LOG_COMMAND
|
||||
|
||||
Dialect = Literal["sqlite", "postgres", "mysql", "unknown"]
|
||||
|
||||
@@ -466,7 +466,7 @@ async def repair_safe_schema_drift() -> SchemaGuardResult:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
connection.execute_query_dict(sql),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
timeout=DB_SCHEMA_GUARD_TIMEOUT_SECONDS,
|
||||
)
|
||||
columns[source] = ColumnInfo(name=source, data_type="")
|
||||
result.repaired_columns += 1
|
||||
@@ -521,7 +521,7 @@ async def repair_safe_schema_drift() -> SchemaGuardResult:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
connection.execute_query_dict(sql),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
timeout=DB_SCHEMA_GUARD_TIMEOUT_SECONDS,
|
||||
)
|
||||
existing_indexes.add(index_columns)
|
||||
result.repaired_indexes += 1
|
||||
@@ -529,6 +529,16 @@ async def repair_safe_schema_drift() -> SchemaGuardResult:
|
||||
f"SchemaGuard 已补齐索引: {table}.{index_columns}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
except TimeoutError as exc:
|
||||
result.warnings += 1
|
||||
result.skipped_indexes += 1
|
||||
logger.warning(
|
||||
"SchemaGuard 补齐索引超时,已跳过: "
|
||||
f"{table}.{index_columns} "
|
||||
f"({DB_SCHEMA_GUARD_TIMEOUT_SECONDS}s)",
|
||||
LOG_COMMAND,
|
||||
e=exc,
|
||||
)
|
||||
except OperationalError as exc:
|
||||
err = str(exc).lower()
|
||||
if any(
|
||||
|
||||
@@ -140,9 +140,6 @@ def register_runtime_bootstrap(_driver) -> None:
|
||||
global _thread_executor
|
||||
await _stop_launcher_watchdog()
|
||||
await stop_send_queue()
|
||||
from zhenxun.models._bot_message_buffer import stop_bot_message_store_buffer
|
||||
|
||||
await stop_bot_message_store_buffer()
|
||||
await stop_memory_governor()
|
||||
executor = _thread_executor
|
||||
_thread_executor = None
|
||||
|
||||
@@ -27,6 +27,7 @@ from zhenxun.services.log import logger
|
||||
from zhenxun.services.message_load import should_pause_tasks
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.decorator.retry import Retry
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
from zhenxun.utils.pydantic_compat import parse_as
|
||||
|
||||
from .repository import ScheduleRepository
|
||||
@@ -37,6 +38,37 @@ SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit"
|
||||
_LAST_PRESSURE_SKIP = 0.0
|
||||
|
||||
|
||||
def _resolve_scheduler_bot(bot_id: str | None, log_target: str) -> Bot | None:
|
||||
if bot_id:
|
||||
try:
|
||||
return nonebot.get_bot(bot_id)
|
||||
except KeyError:
|
||||
logger.warning(f"{log_target} 需要的 Bot {bot_id} 不在线,本次执行跳过。")
|
||||
return None
|
||||
|
||||
bots = list(nonebot.get_bots().values())
|
||||
if not bots:
|
||||
logger.warning(f"{log_target} 当前没有可用 Bot,本次执行跳过。")
|
||||
return None
|
||||
if len(bots) == 1:
|
||||
return bots[0]
|
||||
|
||||
qq_client_bots = [
|
||||
bot for bot in bots if PlatformUtils.get_platform_scope(bot) == "qq_client"
|
||||
]
|
||||
if len(qq_client_bots) == 1:
|
||||
bot = qq_client_bots[0]
|
||||
logger.warning(
|
||||
f"{log_target} 未指定 Bot,多 Bot 在线,自动选择 OneBot {bot.self_id}。"
|
||||
)
|
||||
return bot
|
||||
|
||||
logger.warning(
|
||||
f"{log_target} 未指定 Bot 且多 Bot 在线,无法安全选择," "本次执行跳过。"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
class APSchedulerAdapter:
|
||||
"""封装对 APScheduler 的操作"""
|
||||
|
||||
@@ -343,7 +375,12 @@ async def _execute_job(
|
||||
return
|
||||
|
||||
try:
|
||||
bot = nonebot.get_bot()
|
||||
bot = _resolve_scheduler_bot(
|
||||
context_override.bot_id,
|
||||
f"临时任务 {plugin_name}",
|
||||
)
|
||||
if bot is None:
|
||||
return
|
||||
logger.info(f"开始执行临时任务: {plugin_name}")
|
||||
injected_params = {"context": context_override}
|
||||
state: T_State = {ScheduleContext: context_override}
|
||||
@@ -380,18 +417,9 @@ async def _execute_job(
|
||||
logger.warning(f"定时任务 {schedule_id} 不存在或已禁用,跳过执行。")
|
||||
return
|
||||
|
||||
try:
|
||||
bot = (
|
||||
nonebot.get_bot(schedule.bot_id)
|
||||
if schedule.bot_id
|
||||
else nonebot.get_bot()
|
||||
)
|
||||
except (KeyError, ValueError):
|
||||
logger.warning(
|
||||
f"任务 {schedule_id} 需要的 Bot {schedule.bot_id} "
|
||||
f"不在线,本次执行跳过。"
|
||||
)
|
||||
raise
|
||||
bot = _resolve_scheduler_bot(schedule.bot_id, f"任务 {schedule_id}")
|
||||
if bot is None:
|
||||
return
|
||||
|
||||
resolver = scheduler_manager._target_resolvers.get(schedule.target_type)
|
||||
if not resolver:
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable
|
||||
import contextlib
|
||||
import importlib
|
||||
from typing import Any, cast
|
||||
|
||||
from nonebot.adapters import Bot, Event
|
||||
@@ -10,6 +11,9 @@ from nonebot.log import logger
|
||||
_PATCHED = False
|
||||
_ORIGINAL_FETCH: Callable[..., Awaitable[Any]] | None = None
|
||||
_ORIGINAL_ONEBOT11_GROUP_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
|
||||
_ORIGINAL_QQ_C2C_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
|
||||
_ORIGINAL_QQ_GROUP_AT_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
|
||||
_ORIGINAL_QQ_GUILD_MESSAGE: Callable[..., Awaitable[dict[str, Any]]] | None = None
|
||||
|
||||
|
||||
def _sender_value(sender: Any, key: str, default: Any = None) -> Any:
|
||||
@@ -81,6 +85,88 @@ async def _fast_onebot11_group_message(bot: Bot, event: Event) -> dict[str, Any]
|
||||
}
|
||||
|
||||
|
||||
def _qq_bot_app_id(bot: Bot) -> str:
|
||||
bot_info = getattr(bot, "bot_info", None)
|
||||
app_id = getattr(bot_info, "id", None)
|
||||
return str(app_id or getattr(bot, "self_id", ""))
|
||||
|
||||
|
||||
async def _fast_qq_c2c_message(bot: Bot, event: Event) -> dict[str, Any]:
|
||||
"""Build Uninfo session for QQ official C2C messages from event fields."""
|
||||
|
||||
author = _event_value(event, "author")
|
||||
user_id = str(
|
||||
_sender_value(author, "user_openid")
|
||||
or _sender_value(author, "id")
|
||||
or _event_value(event, "user_id", "")
|
||||
)
|
||||
username = str(_sender_value(author, "username", "") or "")
|
||||
return {
|
||||
"user_id": user_id,
|
||||
"name": username,
|
||||
"nickname": username,
|
||||
"avatar": f"https://q.qlogo.cn/qqapp/{_qq_bot_app_id(bot)}/{user_id}/100",
|
||||
}
|
||||
|
||||
|
||||
async def _fast_qq_group_at_message(bot: Bot, event: Event) -> dict[str, Any]:
|
||||
"""Build Uninfo session for QQ official group-at messages from event fields."""
|
||||
|
||||
author = _event_value(event, "author")
|
||||
user_id = str(
|
||||
_sender_value(author, "member_openid")
|
||||
or _sender_value(author, "id")
|
||||
or _event_value(event, "user_id", "")
|
||||
)
|
||||
username = str(_sender_value(author, "username", "") or "")
|
||||
group_id = str(
|
||||
_event_value(event, "group_openid") or _event_value(event, "group_id") or ""
|
||||
)
|
||||
return {
|
||||
"user_id": user_id,
|
||||
"name": username,
|
||||
"nickname": username,
|
||||
"avatar": f"https://q.qlogo.cn/qqapp/{_qq_bot_app_id(bot)}/{user_id}/100",
|
||||
"group_id": group_id,
|
||||
}
|
||||
|
||||
|
||||
async def _fast_qq_guild_message(bot: Bot, event: Event) -> dict[str, Any]:
|
||||
"""Build Uninfo session for QQ official guild/channel messages locally.
|
||||
|
||||
nonebot-plugin-uninfo enriches guild messages through remote guild/channel
|
||||
APIs. Runtime auth only needs stable scene/user ids, so avoid remote calls
|
||||
during matcher fanout.
|
||||
"""
|
||||
|
||||
author = _event_value(event, "author")
|
||||
member = _event_value(event, "member")
|
||||
guild_id = str(_event_value(event, "guild_id", "") or "")
|
||||
channel_id = str(_event_value(event, "channel_id", "") or "")
|
||||
user_id = str(_sender_value(author, "id", "") or "")
|
||||
nickname = str(_sender_value(member, "nick", "") or "")
|
||||
username = str(_sender_value(author, "username", "") or "")
|
||||
base: dict[str, Any] = {
|
||||
"user_id": user_id,
|
||||
"name": username,
|
||||
"nickname": nickname or username,
|
||||
"avatar": _sender_value(author, "avatar"),
|
||||
"guild_id": guild_id,
|
||||
"channel_id": channel_id,
|
||||
"guild_name": "",
|
||||
"guild_avatar": None,
|
||||
"channel_name": "",
|
||||
"channel_type": -1,
|
||||
}
|
||||
roles = _sender_value(member, "roles")
|
||||
if roles is not None:
|
||||
base["roles"] = roles
|
||||
joined_at = _sender_value(member, "joined_at")
|
||||
if joined_at is not None:
|
||||
base["joined_at"] = joined_at
|
||||
return base
|
||||
|
||||
|
||||
async def _singleflight_fetch(self: Any, bot: Bot, event: Event) -> Any:
|
||||
original = _ORIGINAL_FETCH
|
||||
if original is None:
|
||||
@@ -114,6 +200,8 @@ async def _singleflight_fetch(self: Any, bot: Bot, event: Event) -> Any:
|
||||
|
||||
def apply_uninfo_onebot11_patch() -> None:
|
||||
global _ORIGINAL_FETCH, _ORIGINAL_ONEBOT11_GROUP_MESSAGE, _PATCHED
|
||||
global _ORIGINAL_QQ_C2C_MESSAGE, _ORIGINAL_QQ_GROUP_AT_MESSAGE
|
||||
global _ORIGINAL_QQ_GUILD_MESSAGE
|
||||
if _PATCHED:
|
||||
return
|
||||
|
||||
@@ -129,6 +217,59 @@ def apply_uninfo_onebot11_patch() -> None:
|
||||
setattr(_fast_onebot11_group_message, "__zhenxun_fast_onebot11__", True)
|
||||
fetcher.endpoint[GroupMessageEvent] = _fast_onebot11_group_message
|
||||
|
||||
with contextlib.suppress(Exception):
|
||||
qq_event_module = importlib.import_module("nonebot.adapters.qq.event")
|
||||
AtMessageCreateEvent = getattr(qq_event_module, "AtMessageCreateEvent")
|
||||
C2CMessageCreateEvent = getattr(qq_event_module, "C2CMessageCreateEvent")
|
||||
DirectMessageCreateEvent = getattr(qq_event_module, "DirectMessageCreateEvent")
|
||||
GroupAtMessageCreateEvent = getattr(
|
||||
qq_event_module,
|
||||
"GroupAtMessageCreateEvent",
|
||||
)
|
||||
GroupMessageCreateEvent = getattr(
|
||||
qq_event_module,
|
||||
"GroupMessageCreateEvent",
|
||||
)
|
||||
MessageCreateEvent = getattr(qq_event_module, "MessageCreateEvent")
|
||||
from nonebot_plugin_uninfo.adapters.qq.main import fetcher as qq_fetcher
|
||||
|
||||
original_c2c = qq_fetcher.endpoint.get(C2CMessageCreateEvent)
|
||||
if not getattr(original_c2c, "__zhenxun_fast_qq__", False):
|
||||
_ORIGINAL_QQ_C2C_MESSAGE = cast(
|
||||
Callable[..., Awaitable[dict[str, Any]]] | None,
|
||||
original_c2c,
|
||||
)
|
||||
setattr(_fast_qq_c2c_message, "__zhenxun_fast_qq__", True)
|
||||
qq_fetcher.endpoint[C2CMessageCreateEvent] = _fast_qq_c2c_message
|
||||
|
||||
for event_type in (GroupMessageCreateEvent, GroupAtMessageCreateEvent):
|
||||
original_group_at = qq_fetcher.endpoint.get(event_type)
|
||||
if getattr(original_group_at, "__zhenxun_fast_qq__", False):
|
||||
continue
|
||||
if _ORIGINAL_QQ_GROUP_AT_MESSAGE is None and original_group_at is not None:
|
||||
_ORIGINAL_QQ_GROUP_AT_MESSAGE = cast(
|
||||
Callable[..., Awaitable[dict[str, Any]]],
|
||||
original_group_at,
|
||||
)
|
||||
setattr(_fast_qq_group_at_message, "__zhenxun_fast_qq__", True)
|
||||
qq_fetcher.endpoint[event_type] = _fast_qq_group_at_message
|
||||
|
||||
for event_type in (
|
||||
MessageCreateEvent,
|
||||
AtMessageCreateEvent,
|
||||
DirectMessageCreateEvent,
|
||||
):
|
||||
original_guild = qq_fetcher.endpoint.get(event_type)
|
||||
if getattr(original_guild, "__zhenxun_fast_qq__", False):
|
||||
continue
|
||||
if _ORIGINAL_QQ_GUILD_MESSAGE is None and original_guild is not None:
|
||||
_ORIGINAL_QQ_GUILD_MESSAGE = cast(
|
||||
Callable[..., Awaitable[dict[str, Any]]],
|
||||
original_guild,
|
||||
)
|
||||
setattr(_fast_qq_guild_message, "__zhenxun_fast_qq__", True)
|
||||
qq_fetcher.endpoint[event_type] = _fast_qq_guild_message
|
||||
|
||||
try:
|
||||
from nonebot_plugin_uninfo.fetch import InfoFetcher
|
||||
except Exception as e:
|
||||
@@ -146,4 +287,4 @@ def apply_uninfo_onebot11_patch() -> None:
|
||||
setattr(_singleflight_fetch, "__zhenxun_singleflight__", True)
|
||||
setattr(InfoFetcher, "fetch", _singleflight_fetch)
|
||||
_PATCHED = True
|
||||
logger.debug("Uninfo OneBot11 fast fetch and singleflight patch applied")
|
||||
logger.debug("Uninfo fast fetch and singleflight patch applied")
|
||||
|
||||
Reference in New Issue
Block a user