mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-05 03:39:59 +08:00
✨ feat(core): 增强定时任务与群组标签管理,重构调度核心
✨ 新功能 * **标签 (tags)**: 引入群组标签服务。 * 支持静态标签和动态标签 (基于 Alconna 规则自动匹配群信息)。 * 支持黑名单模式及 `@all` 特殊标签。 * 提供 `tag_manage` 超级用户插件 (list, create, edit, delete 等)。 * 群成员变动时自动失效动态标签缓存。 * **调度 (scheduler)**: 增强定时任务。 * 重构 `ScheduledJob` 模型,支持 `TAG`, `ALL_GROUPS` 等多种目标类型。 * 新增任务别名 (`name`)、创建者、权限、来源等字段。 * 支持一次性任务 (`schedule_once`) 和 Alconna 命令行参数 (`--params-cli`)。 * 新增执行选项 (`jitter`, `spread`) 和并发策略 (`ALLOW`, `SKIP`, `QUEUE`)。 * 支持批量获取任务状态。 ♻️ 重构优化 * **调度器核心**: * 拆分 `service.py` 为 `manager.py` (API) 和 `types.py` (模型)。 * 合并 `adapter.py` / `job.py` 至 `engine.py` (统一调度引擎)。 * 引入 `targeting.py` 模块管理任务目标解析。 * **调度器插件 (scheduler_admin)**: * 迁移命令参数校验逻辑至 `ArparmaBehavior`。 * 引入 `dependencies.py` 和 `data_source.py` 解耦业务逻辑与依赖注入。 * 适配新的任务目标类型展示。
This commit is contained in:
@@ -45,12 +45,18 @@ from .llm import (
|
||||
from .log import logger
|
||||
from .plugin_init import PluginInit, PluginInitManager
|
||||
from .renderer import renderer_service
|
||||
from .scheduler import scheduler_manager
|
||||
from .scheduler import (
|
||||
ExecutionPolicy,
|
||||
ScheduleContext,
|
||||
Trigger,
|
||||
scheduler_manager,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AI",
|
||||
"AIConfig",
|
||||
"CommonOverrides",
|
||||
"ExecutionPolicy",
|
||||
"LLMContentPart",
|
||||
"LLMException",
|
||||
"LLMGenerationConfig",
|
||||
@@ -58,6 +64,8 @@ __all__ = [
|
||||
"Model",
|
||||
"PluginInit",
|
||||
"PluginInitManager",
|
||||
"ScheduleContext",
|
||||
"Trigger",
|
||||
"avatar_service",
|
||||
"chat",
|
||||
"clear_model_cache",
|
||||
|
||||
@@ -4,11 +4,10 @@
|
||||
提供一个统一的、持久化的定时任务管理器,供所有插件使用。
|
||||
"""
|
||||
|
||||
from .job import ScheduleContext
|
||||
from .lifecycle import _load_schedules_from_db
|
||||
from .service import ExecutionPolicy, scheduler_manager
|
||||
from .triggers import Trigger
|
||||
from . import lifecycle
|
||||
from .manager import scheduler_manager
|
||||
from .types import ExecutionPolicy, ScheduleContext, Trigger
|
||||
|
||||
_ = _load_schedules_from_db
|
||||
_ = lifecycle
|
||||
|
||||
__all__ = ["ExecutionPolicy", "ScheduleContext", "Trigger", "scheduler_manager"]
|
||||
|
||||
@@ -1,174 +0,0 @@
|
||||
"""
|
||||
引擎适配层 (Adapter)
|
||||
|
||||
封装所有对具体调度器引擎 (APScheduler) 的操作,
|
||||
使上层服务与调度器实现解耦。
|
||||
"""
|
||||
|
||||
from collections.abc import Callable
|
||||
|
||||
from nonebot_plugin_apscheduler import scheduler
|
||||
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .job import ScheduleContext, _execute_job
|
||||
|
||||
JOB_PREFIX = "zhenxun_schedule_"
|
||||
|
||||
|
||||
class APSchedulerAdapter:
|
||||
"""封装对 APScheduler 的操作"""
|
||||
|
||||
@staticmethod
|
||||
def _get_job_id(schedule_id: int) -> str:
|
||||
"""
|
||||
生成 APScheduler 的 Job ID
|
||||
|
||||
参数:
|
||||
schedule_id: 定时任务的ID。
|
||||
|
||||
返回:
|
||||
str: APScheduler 使用的 Job ID。
|
||||
"""
|
||||
return f"{JOB_PREFIX}{schedule_id}"
|
||||
|
||||
@staticmethod
|
||||
def add_or_reschedule_job(schedule: ScheduledJob):
|
||||
"""
|
||||
根据 ScheduledJob 添加或重新调度一个 APScheduler 任务
|
||||
|
||||
参数:
|
||||
schedule: 定时任务对象,包含任务的所有配置信息。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule.id)
|
||||
|
||||
if not isinstance(schedule.trigger_config, dict):
|
||||
logger.error(
|
||||
f"任务 {schedule.id} 的 trigger_config 不是字典类型: "
|
||||
f"{type(schedule.trigger_config)}"
|
||||
)
|
||||
return
|
||||
|
||||
job = scheduler.get_job(job_id)
|
||||
if job:
|
||||
scheduler.reschedule_job(
|
||||
job_id, trigger=schedule.trigger_type, **schedule.trigger_config
|
||||
)
|
||||
logger.debug(f"已更新APScheduler任务: {job_id}")
|
||||
else:
|
||||
scheduler.add_job(
|
||||
_execute_job,
|
||||
trigger=schedule.trigger_type,
|
||||
id=job_id,
|
||||
misfire_grace_time=300,
|
||||
args=[schedule.id],
|
||||
**schedule.trigger_config,
|
||||
)
|
||||
logger.debug(f"已添加新的APScheduler任务: {job_id}")
|
||||
|
||||
@staticmethod
|
||||
def remove_job(schedule_id: int):
|
||||
"""
|
||||
移除一个 APScheduler 任务
|
||||
|
||||
参数:
|
||||
schedule_id: 要移除的定时任务ID。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule_id)
|
||||
try:
|
||||
scheduler.remove_job(job_id)
|
||||
logger.debug(f"已从APScheduler中移除任务: {job_id}")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def pause_job(schedule_id: int):
|
||||
"""
|
||||
暂停一个 APScheduler 任务
|
||||
|
||||
参数:
|
||||
schedule_id: 要暂停的定时任务ID。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule_id)
|
||||
try:
|
||||
scheduler.pause_job(job_id)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def resume_job(schedule_id: int):
|
||||
"""
|
||||
恢复一个 APScheduler 任务
|
||||
|
||||
参数:
|
||||
schedule_id: 要恢复的定时任务ID。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule_id)
|
||||
try:
|
||||
scheduler.resume_job(job_id)
|
||||
except Exception:
|
||||
import asyncio
|
||||
|
||||
from .repository import ScheduleRepository
|
||||
|
||||
async def _re_add_job():
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
if schedule:
|
||||
APSchedulerAdapter.add_or_reschedule_job(schedule)
|
||||
|
||||
asyncio.create_task(_re_add_job()) # noqa: RUF006
|
||||
|
||||
@staticmethod
|
||||
def get_job_status(schedule_id: int) -> dict:
|
||||
"""
|
||||
获取 APScheduler Job 的状态
|
||||
|
||||
参数:
|
||||
schedule_id: 定时任务的ID。
|
||||
|
||||
返回:
|
||||
dict: 包含任务状态信息的字典,包含next_run_time等字段。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule_id)
|
||||
job = scheduler.get_job(job_id)
|
||||
return {
|
||||
"next_run_time": job.next_run_time.strftime("%Y-%m-%d %H:%M:%S")
|
||||
if job and job.next_run_time
|
||||
else "N/A",
|
||||
"is_paused_in_scheduler": not bool(job.next_run_time) if job else "N/A",
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def add_ephemeral_job(
|
||||
job_id: str,
|
||||
func: Callable,
|
||||
trigger_type: str,
|
||||
trigger_config: dict,
|
||||
context: ScheduleContext,
|
||||
):
|
||||
"""
|
||||
直接向 APScheduler 添加一个临时的、非持久化的任务
|
||||
|
||||
参数:
|
||||
job_id: 临时任务的唯一ID。
|
||||
func: 要执行的函数。
|
||||
trigger_type: 触发器类型。
|
||||
trigger_config: 触发器配置字典。
|
||||
context: 任务执行上下文。
|
||||
"""
|
||||
job = scheduler.get_job(job_id)
|
||||
if job:
|
||||
logger.warning(f"尝试添加一个已存在的临时任务ID: {job_id},操作被忽略。")
|
||||
return
|
||||
|
||||
scheduler.add_job(
|
||||
_execute_job,
|
||||
trigger=trigger_type,
|
||||
id=job_id,
|
||||
misfire_grace_time=60,
|
||||
args=[None],
|
||||
kwargs={"context_override": context},
|
||||
**trigger_config,
|
||||
)
|
||||
logger.debug(f"已添加新的临时APScheduler任务: {job_id}")
|
||||
@@ -0,0 +1,454 @@
|
||||
"""
|
||||
引擎适配层 (Adapter) 与 任务执行逻辑 (Job)
|
||||
|
||||
封装所有对具体调度器引擎 (APScheduler) 的操作,
|
||||
以及被 APScheduler 实际调度的函数。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from functools import partial
|
||||
import random
|
||||
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.dependencies import Dependent
|
||||
from nonebot.exception import FinishedException, PausedException, SkippedException
|
||||
from nonebot.matcher import Matcher
|
||||
from nonebot.typing import T_State
|
||||
from nonebot_plugin_apscheduler import scheduler
|
||||
from pydantic import BaseModel
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.decorator.retry import Retry
|
||||
from zhenxun.utils.pydantic_compat import parse_as
|
||||
|
||||
from .repository import ScheduleRepository
|
||||
from .types import ExecutionPolicy, ScheduleContext
|
||||
|
||||
JOB_PREFIX = "zhenxun_schedule_"
|
||||
SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit"
|
||||
|
||||
|
||||
class APSchedulerAdapter:
|
||||
"""封装对 APScheduler 的操作"""
|
||||
|
||||
@staticmethod
|
||||
def _get_job_id(schedule_id: int) -> str:
|
||||
"""
|
||||
生成 APScheduler 的 Job ID
|
||||
|
||||
参数:
|
||||
schedule_id: 定时任务的ID。
|
||||
|
||||
返回:
|
||||
str: APScheduler 使用的 Job ID。
|
||||
"""
|
||||
return f"{JOB_PREFIX}{schedule_id}"
|
||||
|
||||
@staticmethod
|
||||
def add_or_reschedule_job(schedule: ScheduledJob):
|
||||
"""
|
||||
根据 ScheduledJob 添加或重新调度一个 APScheduler 任务
|
||||
|
||||
参数:
|
||||
schedule: 定时任务对象,包含任务的所有配置信息。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule.id)
|
||||
|
||||
try:
|
||||
scheduler.remove_job(job_id)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if not isinstance(schedule.trigger_config, dict):
|
||||
logger.error(
|
||||
f"任务 {schedule.id} 的 trigger_config 不是字典类型: "
|
||||
f"{type(schedule.trigger_config)}"
|
||||
)
|
||||
return
|
||||
|
||||
trigger_params = schedule.trigger_config.copy()
|
||||
execution_options = (
|
||||
schedule.execution_options
|
||||
if isinstance(schedule.execution_options, dict)
|
||||
else {}
|
||||
)
|
||||
if jitter := execution_options.get("jitter"):
|
||||
if isinstance(jitter, int) and jitter > 0:
|
||||
trigger_params["jitter"] = jitter
|
||||
|
||||
concurrency_policy = execution_options.get("concurrency_policy", "ALLOW")
|
||||
job_params = {
|
||||
"id": job_id,
|
||||
"misfire_grace_time": 300,
|
||||
"args": [schedule.id],
|
||||
}
|
||||
|
||||
if concurrency_policy == "SKIP":
|
||||
job_params["max_instances"] = 1
|
||||
job_params["coalesce"] = True
|
||||
elif concurrency_policy == "QUEUE":
|
||||
job_params["max_instances"] = 1
|
||||
job_params["coalesce"] = False
|
||||
|
||||
scheduler.add_job(
|
||||
_execute_job,
|
||||
trigger=schedule.trigger_type,
|
||||
**job_params,
|
||||
**trigger_params,
|
||||
)
|
||||
logger.debug(
|
||||
f"已添加或更新APScheduler任务: {job_id} | 并发策略: {concurrency_policy}, "
|
||||
f"抖动: {trigger_params.get('jitter', '无')}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def remove_job(schedule_id: int):
|
||||
"""
|
||||
移除一个 APScheduler 任务
|
||||
|
||||
参数:
|
||||
schedule_id: 要移除的定时任务ID。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule_id)
|
||||
try:
|
||||
scheduler.remove_job(job_id)
|
||||
logger.debug(f"已从APScheduler中移除任务: {job_id}")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def pause_job(schedule_id: int):
|
||||
"""
|
||||
暂停一个 APScheduler 任务
|
||||
|
||||
参数:
|
||||
schedule_id: 要暂停的定时任务ID。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule_id)
|
||||
try:
|
||||
scheduler.pause_job(job_id)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def resume_job(schedule_id: int):
|
||||
"""
|
||||
恢复一个 APScheduler 任务
|
||||
|
||||
参数:
|
||||
schedule_id: 要恢复的定时任务ID。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule_id)
|
||||
try:
|
||||
scheduler.resume_job(job_id)
|
||||
except Exception:
|
||||
import asyncio
|
||||
|
||||
async def _re_add_job():
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
if schedule:
|
||||
APSchedulerAdapter.add_or_reschedule_job(schedule)
|
||||
|
||||
asyncio.create_task(_re_add_job()) # noqa: RUF006
|
||||
|
||||
@staticmethod
|
||||
def get_job_status(schedule_id: int) -> dict:
|
||||
"""
|
||||
获取 APScheduler Job 的状态
|
||||
|
||||
参数:
|
||||
schedule_id: 定时任务的ID。
|
||||
|
||||
返回:
|
||||
dict: 包含任务状态信息的字典,包含next_run_time等字段。
|
||||
"""
|
||||
job_id = APSchedulerAdapter._get_job_id(schedule_id)
|
||||
job = scheduler.get_job(job_id)
|
||||
return {
|
||||
"next_run_time": job.next_run_time.strftime("%Y-%m-%d %H:%M:%S")
|
||||
if job and job.next_run_time
|
||||
else "N/A",
|
||||
"is_paused_in_scheduler": not bool(job.next_run_time) if job else "N/A",
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def add_ephemeral_job(
|
||||
job_id: str,
|
||||
func: Callable,
|
||||
trigger_type: str,
|
||||
trigger_config: dict,
|
||||
context: ScheduleContext,
|
||||
):
|
||||
"""
|
||||
直接向 APScheduler 添加一个临时的、非持久化的任务
|
||||
|
||||
参数:
|
||||
job_id: 临时任务的唯一ID。
|
||||
func: 要执行的函数。
|
||||
trigger_type: 触发器类型。
|
||||
trigger_config: 触发器配置字典。
|
||||
context: 任务执行上下文。
|
||||
"""
|
||||
job = scheduler.get_job(job_id)
|
||||
if job:
|
||||
logger.warning(f"尝试添加一个已存在的临时任务ID: {job_id},操作被忽略。")
|
||||
return
|
||||
|
||||
scheduler.add_job(
|
||||
_execute_job,
|
||||
trigger=trigger_type,
|
||||
id=job_id,
|
||||
misfire_grace_time=60,
|
||||
args=[None],
|
||||
kwargs={"context_override": context},
|
||||
**trigger_config,
|
||||
)
|
||||
logger.debug(f"已添加新的临时APScheduler任务: {job_id}")
|
||||
|
||||
|
||||
async def _execute_single_job_instance(
|
||||
schedule: ScheduledJob, bot, group_id: str | None = None
|
||||
):
|
||||
"""
|
||||
负责执行一个具体目标的任务实例。
|
||||
"""
|
||||
from .manager import scheduler_manager
|
||||
|
||||
plugin_name = schedule.plugin_name
|
||||
if group_id is None and schedule.target_type == "GROUP":
|
||||
group_id = schedule.target_identifier
|
||||
|
||||
task_meta = scheduler_manager._registered_tasks.get(plugin_name)
|
||||
|
||||
if not task_meta:
|
||||
logger.error(f"无法执行任务:插件 '{plugin_name}' 在执行期间变得不可用。")
|
||||
return
|
||||
|
||||
is_blocked = await CommonUtils.task_is_block(bot, plugin_name, group_id)
|
||||
if is_blocked:
|
||||
target_desc = f"群 {group_id}" if group_id else "全局"
|
||||
logger.info(
|
||||
f"插件 '{plugin_name}' 的定时任务在目标 [{target_desc}] "
|
||||
f"因功能被禁用而跳过执行。"
|
||||
)
|
||||
return
|
||||
|
||||
context = ScheduleContext(
|
||||
schedule_id=schedule.id,
|
||||
plugin_name=plugin_name,
|
||||
bot_id=bot.self_id,
|
||||
group_id=group_id,
|
||||
job_kwargs=schedule.job_kwargs if isinstance(schedule.job_kwargs, dict) else {},
|
||||
)
|
||||
state: T_State = {ScheduleContext: context}
|
||||
|
||||
policy_data = context.job_kwargs.pop("execution_policy", {})
|
||||
policy = ExecutionPolicy(**policy_data)
|
||||
|
||||
async def task_execution_coro():
|
||||
injected_params = {"context": context}
|
||||
|
||||
params_model = task_meta.get("model")
|
||||
if params_model and isinstance(context.job_kwargs, dict):
|
||||
try:
|
||||
if isinstance(params_model, type) and issubclass(
|
||||
params_model, BaseModel
|
||||
):
|
||||
params_instance = parse_as(params_model, context.job_kwargs)
|
||||
injected_params["params"] = params_instance # type: ignore
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"任务 {schedule.id} (目标: {group_id}) 参数验证失败: {e}", e=e
|
||||
)
|
||||
raise
|
||||
|
||||
async def wrapper(bot: Bot):
|
||||
return await task_meta["func"](bot=bot, **injected_params) # type: ignore
|
||||
|
||||
dependent = Dependent.parse(
|
||||
call=wrapper,
|
||||
allow_types=Matcher.HANDLER_PARAM_TYPES,
|
||||
)
|
||||
return await dependent(bot=bot, state=state)
|
||||
|
||||
try:
|
||||
if policy.retries > 0:
|
||||
on_success_handler = None
|
||||
if policy.on_success_callback:
|
||||
on_success_handler = partial(policy.on_success_callback, context)
|
||||
|
||||
on_failure_handler = None
|
||||
if policy.on_failure_callback:
|
||||
on_failure_handler = partial(policy.on_failure_callback, context)
|
||||
|
||||
retry_exceptions = tuple(policy.retry_on_exceptions or [])
|
||||
|
||||
retry_decorator = Retry.api(
|
||||
stop_max_attempt=policy.retries + 1,
|
||||
strategy="exponential" if policy.retry_backoff else "fixed",
|
||||
wait_fixed_seconds=policy.retry_delay_seconds,
|
||||
exception=retry_exceptions,
|
||||
on_success=on_success_handler,
|
||||
on_failure=on_failure_handler,
|
||||
log_name=f"ScheduledJob-{schedule.id}-{group_id or 'global'}",
|
||||
)
|
||||
|
||||
decorated_executor = retry_decorator(task_execution_coro)
|
||||
await decorated_executor()
|
||||
else:
|
||||
logger.info(
|
||||
f"插件 '{plugin_name}' 开始为目标 [{group_id or '全局'}] "
|
||||
f"执行定时任务 (ID: {schedule.id})。"
|
||||
)
|
||||
await task_execution_coro()
|
||||
|
||||
except (PausedException, FinishedException, SkippedException) as e:
|
||||
logger.warning(
|
||||
f"定时任务 {schedule.id} (目标: {group_id}) 被中断: {type(e).__name__}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"执行定时任务 {schedule.id} (目标: {group_id}) "
|
||||
f"时发生未被策略处理的最终错误",
|
||||
e=e,
|
||||
)
|
||||
|
||||
|
||||
async def _execute_job(
|
||||
schedule_id: int | None,
|
||||
force: bool = False,
|
||||
context_override: ScheduleContext | None = None,
|
||||
):
|
||||
"""
|
||||
APScheduler 调度的入口函数,现在作为分发器。
|
||||
"""
|
||||
from .manager import scheduler_manager
|
||||
|
||||
schedule = None
|
||||
|
||||
if context_override:
|
||||
plugin_name = context_override.plugin_name
|
||||
task_meta = scheduler_manager._registered_tasks.get(plugin_name)
|
||||
if not task_meta or not task_meta["func"]:
|
||||
logger.error(f"无法执行临时任务:函数 '{plugin_name}' 未注册。")
|
||||
return
|
||||
|
||||
try:
|
||||
bot = nonebot.get_bot()
|
||||
logger.info(f"开始执行临时任务: {plugin_name}")
|
||||
injected_params = {"context": context_override}
|
||||
state: T_State = {ScheduleContext: context_override}
|
||||
|
||||
async def wrapper(bot: Bot):
|
||||
return await task_meta["func"](bot=bot, **injected_params) # type: ignore
|
||||
|
||||
dependent = Dependent.parse(
|
||||
call=wrapper,
|
||||
allow_types=Matcher.HANDLER_PARAM_TYPES,
|
||||
)
|
||||
await dependent(bot=bot, state=state)
|
||||
logger.info(f"临时任务 '{plugin_name}' 执行完成。")
|
||||
except Exception as e:
|
||||
logger.error(f"执行临时任务 '{plugin_name}' 时发生错误", e=e)
|
||||
return
|
||||
|
||||
if schedule_id is None:
|
||||
logger.error("执行持久化任务时 schedule_id 不能为空。")
|
||||
return
|
||||
|
||||
scheduler_manager._running_tasks.add(schedule_id)
|
||||
try:
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
if not schedule or (not schedule.is_enabled and not force):
|
||||
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
|
||||
|
||||
resolver = scheduler_manager._target_resolvers.get(schedule.target_type)
|
||||
if not resolver:
|
||||
logger.error(
|
||||
f"任务 {schedule.id} 的目标类型 '{schedule.target_type}' "
|
||||
f"没有注册解析器,执行跳过。"
|
||||
)
|
||||
raise ValueError(f"未知的目标类型: {schedule.target_type}")
|
||||
|
||||
try:
|
||||
resolved_targets = await resolver(schedule.target_identifier, bot)
|
||||
except Exception as e:
|
||||
logger.error(f"为任务 {schedule.id} 解析目标失败", e=e)
|
||||
raise
|
||||
|
||||
logger.info(
|
||||
f"任务 {schedule.id} ({schedule.name or schedule.plugin_name}) 开始执行, "
|
||||
f"目标类型: {schedule.target_type}, "
|
||||
f"解析出 {len(resolved_targets)} 个目标"
|
||||
)
|
||||
|
||||
concurrency_limit = Config.get_config(
|
||||
"SchedulerManager", SCHEDULE_CONCURRENCY_KEY, 5
|
||||
)
|
||||
semaphore = asyncio.Semaphore(concurrency_limit if concurrency_limit > 0 else 5)
|
||||
|
||||
spread_config = (
|
||||
schedule.execution_options
|
||||
if isinstance(schedule.execution_options, dict)
|
||||
else {}
|
||||
)
|
||||
spread_seconds = spread_config.get("spread", 1.0)
|
||||
|
||||
async def worker(target_id: str | None):
|
||||
await asyncio.sleep(random.uniform(0.1, spread_seconds))
|
||||
async with semaphore:
|
||||
await _execute_single_job_instance(schedule, bot, group_id=target_id)
|
||||
|
||||
tasks_to_run = [worker(target_id) for target_id in resolved_targets]
|
||||
if tasks_to_run:
|
||||
await asyncio.gather(*tasks_to_run, return_exceptions=True)
|
||||
|
||||
schedule.last_run_at = datetime.now()
|
||||
schedule.last_run_status = "SUCCESS"
|
||||
schedule.consecutive_failures = 0
|
||||
await schedule.save(
|
||||
update_fields=["last_run_at", "last_run_status", "consecutive_failures"]
|
||||
)
|
||||
|
||||
if schedule.is_one_off:
|
||||
logger.info(f"一次性任务 {schedule.id} 执行成功,将被删除。")
|
||||
await ScheduledJob.filter(id=schedule.id).delete()
|
||||
APSchedulerAdapter.remove_job(schedule.id)
|
||||
if schedule.plugin_name.startswith("runtime_one_off__"):
|
||||
scheduler_manager._registered_tasks.pop(schedule.plugin_name, None)
|
||||
logger.debug(f"已注销一次性运行时任务: {schedule.plugin_name}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"执行任务 {schedule_id} 期间发生严重错误", e=e)
|
||||
|
||||
if schedule:
|
||||
schedule.last_run_at = datetime.now()
|
||||
schedule.last_run_status = "FAILURE"
|
||||
schedule.consecutive_failures = (schedule.consecutive_failures or 0) + 1
|
||||
await schedule.save(
|
||||
update_fields=["last_run_at", "last_run_status", "consecutive_failures"]
|
||||
)
|
||||
|
||||
finally:
|
||||
if schedule_id is not None:
|
||||
scheduler_manager._running_tasks.discard(schedule_id)
|
||||
@@ -1,239 +0,0 @@
|
||||
"""
|
||||
定时任务的执行逻辑
|
||||
|
||||
包含被 APScheduler 实际调度的函数,以及处理不同目标(单个、所有群组)的执行策略。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
from functools import partial
|
||||
import random
|
||||
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.dependencies import Dependent
|
||||
from nonebot.exception import FinishedException, PausedException, SkippedException
|
||||
from nonebot.matcher import Matcher
|
||||
from nonebot.typing import T_State
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services.log import logger
|
||||
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
|
||||
|
||||
SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit"
|
||||
|
||||
|
||||
class ScheduleContext(BaseModel):
|
||||
"""
|
||||
定时任务执行上下文,可通过依赖注入获取。
|
||||
"""
|
||||
|
||||
schedule_id: int = Field(..., description="数据库中的任务ID")
|
||||
plugin_name: str = Field(..., description="任务所属的插件名称")
|
||||
bot_id: str | None = Field(None, description="执行任务的Bot ID")
|
||||
group_id: str | None = Field(None, description="任务目标群组ID")
|
||||
job_kwargs: dict = Field(default_factory=dict, description="任务配置的参数")
|
||||
|
||||
|
||||
async def _execute_single_job_instance(schedule: ScheduledJob, bot):
|
||||
"""
|
||||
负责执行一个具体目标的任务实例。
|
||||
"""
|
||||
plugin_name = schedule.plugin_name
|
||||
group_id = schedule.group_id
|
||||
|
||||
from .service import ExecutionPolicy, scheduler_manager
|
||||
|
||||
task_meta = scheduler_manager._registered_tasks.get(plugin_name)
|
||||
|
||||
if not task_meta:
|
||||
logger.error(f"无法执行任务:插件 '{plugin_name}' 在执行期间变得不可用。")
|
||||
return
|
||||
|
||||
is_blocked = await CommonUtils.task_is_block(bot, plugin_name, group_id)
|
||||
if is_blocked:
|
||||
target_desc = f"群 {group_id}" if group_id else "全局"
|
||||
logger.info(
|
||||
f"插件 '{plugin_name}' 的定时任务在目标 [{target_desc}] "
|
||||
f"因功能被禁用而跳过执行。"
|
||||
)
|
||||
return
|
||||
|
||||
context = ScheduleContext(
|
||||
schedule_id=schedule.id,
|
||||
plugin_name=schedule.plugin_name,
|
||||
bot_id=bot.self_id,
|
||||
group_id=schedule.group_id,
|
||||
job_kwargs=schedule.job_kwargs if isinstance(schedule.job_kwargs, dict) else {},
|
||||
)
|
||||
state: T_State = {ScheduleContext: context}
|
||||
|
||||
policy_data = context.job_kwargs.pop("execution_policy", {})
|
||||
policy = ExecutionPolicy(**policy_data)
|
||||
|
||||
async def task_execution_coro():
|
||||
injected_params = {"context": context}
|
||||
|
||||
params_model = task_meta.get("model")
|
||||
if params_model and isinstance(context.job_kwargs, dict):
|
||||
try:
|
||||
if isinstance(params_model, type) and issubclass(
|
||||
params_model, BaseModel
|
||||
):
|
||||
params_instance = parse_as(params_model, context.job_kwargs)
|
||||
injected_params["params"] = params_instance # type: ignore
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"任务 {schedule.id} (目标: {group_id}) 参数验证失败: {e}", e=e
|
||||
)
|
||||
raise
|
||||
|
||||
async def wrapper(bot: Bot):
|
||||
return await task_meta["func"](bot=bot, **injected_params) # type: ignore
|
||||
|
||||
dependent = Dependent.parse(
|
||||
call=wrapper,
|
||||
allow_types=Matcher.HANDLER_PARAM_TYPES,
|
||||
)
|
||||
return await dependent(bot=bot, state=state)
|
||||
|
||||
try:
|
||||
if policy.retries > 0:
|
||||
on_success_handler = None
|
||||
if policy.on_success_callback:
|
||||
on_success_handler = partial(policy.on_success_callback, context)
|
||||
|
||||
on_failure_handler = None
|
||||
if policy.on_failure_callback:
|
||||
on_failure_handler = partial(policy.on_failure_callback, context)
|
||||
|
||||
retry_exceptions = tuple(policy.retry_on_exceptions or [])
|
||||
|
||||
retry_decorator = Retry.api(
|
||||
stop_max_attempt=policy.retries + 1,
|
||||
strategy="exponential" if policy.retry_backoff else "fixed",
|
||||
wait_fixed_seconds=policy.retry_delay_seconds,
|
||||
exception=retry_exceptions,
|
||||
on_success=on_success_handler,
|
||||
on_failure=on_failure_handler,
|
||||
log_name=f"ScheduledJob-{schedule.id}-{schedule.group_id or 'global'}",
|
||||
)
|
||||
|
||||
decorated_executor = retry_decorator(task_execution_coro)
|
||||
await decorated_executor()
|
||||
else:
|
||||
logger.info(
|
||||
f"插件 '{plugin_name}' 开始为目标 [{group_id or '全局'}] "
|
||||
f"执行定时任务 (ID: {schedule.id})。"
|
||||
)
|
||||
await task_execution_coro()
|
||||
|
||||
except (PausedException, FinishedException, SkippedException) as e:
|
||||
logger.warning(
|
||||
f"定时任务 {schedule.id} (目标: {group_id}) 被中断: {type(e).__name__}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"执行定时任务 {schedule.id} (目标: {group_id}) "
|
||||
f"时发生未被策略处理的最终错误",
|
||||
e=e,
|
||||
)
|
||||
|
||||
|
||||
async def _execute_job(schedule_id: int):
|
||||
"""
|
||||
APScheduler 调度的入口函数,现在作为分发器。
|
||||
"""
|
||||
from .repository import ScheduleRepository
|
||||
from .service import scheduler_manager
|
||||
|
||||
scheduler_manager._running_tasks.add(schedule_id)
|
||||
try:
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
if not schedule or not schedule.is_enabled:
|
||||
logger.warning(f"定时任务 {schedule_id} 不存在或已禁用,跳过执行。")
|
||||
return
|
||||
|
||||
if schedule.plugin_name not in scheduler_manager._registered_tasks:
|
||||
logger.error(
|
||||
f"无法执行定时任务:插件 '{schedule.plugin_name}' "
|
||||
f"未注册或已卸载。将禁用该任务。"
|
||||
)
|
||||
schedule.is_enabled = False
|
||||
await ScheduleRepository.save(schedule, update_fields=["is_enabled"])
|
||||
from .adapter import APSchedulerAdapter
|
||||
|
||||
APSchedulerAdapter.remove_job(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"不在线,本次执行跳过。"
|
||||
)
|
||||
return
|
||||
|
||||
if schedule.group_id == scheduler_manager.ALL_GROUPS:
|
||||
concurrency_limit = Config.get_config(
|
||||
"SchedulerManager", SCHEDULE_CONCURRENCY_KEY, 5
|
||||
)
|
||||
if not isinstance(concurrency_limit, int) or concurrency_limit <= 0:
|
||||
concurrency_limit = 5
|
||||
|
||||
logger.info(
|
||||
f"开始执行针对 [所有群组] 的任务 (ID: {schedule.id}, "
|
||||
f"插件: {schedule.plugin_name}, Bot: {bot.self_id}),"
|
||||
f"并发限制: {concurrency_limit}"
|
||||
)
|
||||
|
||||
try:
|
||||
group_list, _ = await PlatformUtils.get_group_list(bot)
|
||||
all_gids = {
|
||||
g.group_id for g in group_list if g.group_id and not g.channel_id
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"为 'all' 任务获取 Bot {bot.self_id} 的群列表失败", e=e)
|
||||
return
|
||||
|
||||
specific_tasks_gids = set(
|
||||
await ScheduledJob.filter(
|
||||
plugin_name=schedule.plugin_name, group_id__in=list(all_gids)
|
||||
).values_list("group_id", flat=True)
|
||||
)
|
||||
|
||||
semaphore = asyncio.Semaphore(concurrency_limit)
|
||||
|
||||
async def worker(gid: str):
|
||||
await asyncio.sleep(random.uniform(0.1, 1.0))
|
||||
async with semaphore:
|
||||
temp_schedule = copy.deepcopy(schedule)
|
||||
temp_schedule.group_id = gid
|
||||
await _execute_single_job_instance(temp_schedule, bot)
|
||||
|
||||
tasks_to_run = [
|
||||
worker(gid) for gid in all_gids if gid not in specific_tasks_gids
|
||||
]
|
||||
|
||||
if tasks_to_run:
|
||||
await asyncio.gather(*tasks_to_run)
|
||||
logger.info(
|
||||
f"针对 [所有群组] 的任务 (ID: {schedule.id}) 执行完毕,"
|
||||
f"共处理 {len(tasks_to_run)} 个群组。"
|
||||
)
|
||||
|
||||
else:
|
||||
await _execute_single_job_instance(schedule, bot)
|
||||
|
||||
finally:
|
||||
scheduler_manager._running_tasks.discard(schedule_id)
|
||||
@@ -8,10 +8,10 @@ from zhenxun.services.log import logger
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
from .adapter import APSchedulerAdapter
|
||||
from .job import ScheduleContext
|
||||
from .engine import APSchedulerAdapter
|
||||
from .manager import scheduler_manager
|
||||
from .repository import ScheduleRepository
|
||||
from .service import scheduler_manager
|
||||
from .types import ScheduleContext
|
||||
|
||||
|
||||
@PriorityLifecycle.on_startup(priority=90)
|
||||
@@ -37,7 +37,7 @@ async def _load_schedules_from_db():
|
||||
|
||||
query_kwargs = {
|
||||
"plugin_name": plugin_name,
|
||||
"group_id": group_id,
|
||||
"target_identifier": group_id or "",
|
||||
"bot_id": bot_id,
|
||||
}
|
||||
exists = await ScheduleRepository.exists(**query_kwargs)
|
||||
@@ -49,9 +49,13 @@ async def _load_schedules_from_db():
|
||||
task_info.trigger, exclude={"trigger_type"}
|
||||
)
|
||||
|
||||
target_type = "GROUP" if group_id else "GLOBAL"
|
||||
target_identifier = group_id or ""
|
||||
|
||||
schedule = await scheduler_manager.add_schedule(
|
||||
plugin_name=plugin_name,
|
||||
group_id=group_id,
|
||||
target_type=target_type,
|
||||
target_identifier=target_identifier,
|
||||
trigger_type=task_info.trigger.trigger_type,
|
||||
trigger_config=trigger_config_dict,
|
||||
job_kwargs=task_info.job_kwargs,
|
||||
|
||||
@@ -1,81 +1,92 @@
|
||||
"""
|
||||
服务层 (Service)
|
||||
服务层 (Service Manager)
|
||||
|
||||
定义 SchedulerManager 类作为定时任务服务的公共 API 入口。
|
||||
它负责编排业务逻辑,并调用 Repository 和 Adapter 层来完成具体工作。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable, Coroutine
|
||||
from datetime import datetime
|
||||
import inspect
|
||||
from typing import Any, ClassVar
|
||||
import uuid
|
||||
|
||||
from arclet.alconna import Alconna, Option
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot
|
||||
from pydantic import BaseModel
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
from zhenxun.utils.pydantic_compat import model_dump, model_validate
|
||||
|
||||
from .adapter import APSchedulerAdapter
|
||||
from .job import ScheduleContext, _execute_job
|
||||
from .engine import APSchedulerAdapter
|
||||
from .repository import ScheduleRepository
|
||||
from .targeter import ScheduleTargeter
|
||||
from .triggers import BaseTrigger
|
||||
|
||||
|
||||
class ExecutionPolicy(BaseModel):
|
||||
"""
|
||||
封装定时任务的执行策略,包括重试和回调。
|
||||
"""
|
||||
|
||||
retries: int = 0
|
||||
retry_delay_seconds: int = 30
|
||||
retry_backoff: bool = False
|
||||
retry_on_exceptions: list[type[Exception]] | None = None
|
||||
on_success_callback: Callable[[ScheduleContext, Any], Awaitable[None]] | None = None
|
||||
on_failure_callback: (
|
||||
Callable[[ScheduleContext, Exception], Awaitable[None]] | None
|
||||
) = None
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
|
||||
class ScheduledJobDeclaration(BaseModel):
|
||||
"""用于在启动时声明默认定时任务的内部数据模型"""
|
||||
|
||||
plugin_name: str
|
||||
group_id: str | None
|
||||
bot_id: str | None
|
||||
trigger: BaseTrigger
|
||||
job_kwargs: dict[str, Any]
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
|
||||
class EphemeralJobDeclaration(BaseModel):
|
||||
"""用于在启动时声明临时任务的内部数据模型"""
|
||||
|
||||
plugin_name: str
|
||||
func: Callable[..., Coroutine]
|
||||
trigger: BaseTrigger
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
from .targeting import (
|
||||
ScheduleTargeter,
|
||||
)
|
||||
from .types import (
|
||||
BaseTrigger,
|
||||
EphemeralJobDeclaration,
|
||||
ExecutionOptions,
|
||||
ExecutionPolicy,
|
||||
ScheduleContext,
|
||||
ScheduledJobDeclaration,
|
||||
)
|
||||
|
||||
|
||||
class SchedulerManager:
|
||||
ALL_GROUPS: ClassVar[str] = "__ALL_GROUPS__"
|
||||
_registered_tasks: ClassVar[
|
||||
dict[str, dict[str, Callable | type[BaseModel] | None]]
|
||||
dict[
|
||||
str,
|
||||
dict[str, Callable | type[BaseModel] | int | list[Option] | Alconna | None],
|
||||
]
|
||||
] = {}
|
||||
_declared_tasks: ClassVar[list[ScheduledJobDeclaration]] = []
|
||||
_ephemeral_declared_tasks: ClassVar[list[EphemeralJobDeclaration]] = []
|
||||
_running_tasks: ClassVar[set] = set()
|
||||
_target_resolvers: ClassVar[
|
||||
dict[str, Callable[[str, Bot], Awaitable[list[str | None]]]]
|
||||
] = {}
|
||||
|
||||
def __init__(self):
|
||||
self._register_builtin_resolvers()
|
||||
|
||||
def _register_builtin_resolvers(self):
|
||||
"""在管理器初始化时注册所有内置的目标解析器。"""
|
||||
from .targeting import (
|
||||
_resolve_all_groups,
|
||||
_resolve_global_or_user,
|
||||
_resolve_group,
|
||||
_resolve_tag,
|
||||
_resolve_user,
|
||||
)
|
||||
|
||||
if "GROUP" in self._target_resolvers:
|
||||
return
|
||||
self.register_target_resolver("GROUP", _resolve_group)
|
||||
self.register_target_resolver("TAG", _resolve_tag)
|
||||
self.register_target_resolver("ALL_GROUPS", _resolve_all_groups)
|
||||
self.register_target_resolver("GLOBAL", _resolve_global_or_user)
|
||||
self.register_target_resolver("USER", _resolve_user)
|
||||
logger.debug("已注册所有内置的定时任务目标解析器。")
|
||||
|
||||
def register_target_resolver(
|
||||
self,
|
||||
target_type: str,
|
||||
resolver_func: Callable[[str, Bot], Awaitable[list[str | None]]],
|
||||
):
|
||||
"""
|
||||
注册一个新的目标类型解析器。
|
||||
"""
|
||||
if target_type in self._target_resolvers:
|
||||
logger.warning(f"目标解析器 '{target_type}' 已存在,将被覆盖。")
|
||||
self._target_resolvers[target_type.upper()] = resolver_func
|
||||
logger.info(f"已注册新的定时任务目标解析器: '{target_type}'")
|
||||
|
||||
def target(self, **filters: Any) -> ScheduleTargeter:
|
||||
"""
|
||||
@@ -96,22 +107,11 @@ class SchedulerManager:
|
||||
bot_id: str | None = None,
|
||||
default_params: BaseModel | None = None,
|
||||
policy: ExecutionPolicy | None = None,
|
||||
default_jitter: int | None = None,
|
||||
default_spread: int | None = None,
|
||||
):
|
||||
"""
|
||||
声明式定时任务的统一装饰器。
|
||||
|
||||
此装饰器用于将一个异步函数注册为一个可调度的定时任务,
|
||||
并为其创建一个默认的调度计划。
|
||||
|
||||
参数:
|
||||
trigger: 一个由 `Trigger` 工厂类创建的触发器配置对象
|
||||
(例如 `Trigger.cron(hour=8)`)。
|
||||
group_id: 默认的目标群组ID。`None` 表示全局任务,
|
||||
`SchedulerManager.ALL_GROUPS` 表示所有群组。
|
||||
bot_id: 默认的目标Bot ID,`None` 表示使用任意可用Bot。
|
||||
default_params: (可选) 一个Pydantic模型实例,为任务提供默认参数。
|
||||
任务函数需要有对应的Pydantic模型类型注解。
|
||||
policy: (可选) 一个ExecutionPolicy实例,定义任务的执行策略。
|
||||
"""
|
||||
|
||||
def decorator(func: Callable[..., Coroutine]) -> Callable[..., Coroutine]:
|
||||
@@ -122,7 +122,7 @@ class SchedulerManager:
|
||||
plugin_name = plugin.name
|
||||
|
||||
params_model = None
|
||||
from .job import ScheduleContext
|
||||
from .types import ScheduleContext
|
||||
|
||||
for param in inspect.signature(func).parameters.values():
|
||||
if (
|
||||
@@ -138,6 +138,8 @@ class SchedulerManager:
|
||||
self._registered_tasks[plugin_name] = {
|
||||
"func": func,
|
||||
"model": params_model,
|
||||
"default_jitter": default_jitter,
|
||||
"default_spread": default_spread,
|
||||
}
|
||||
|
||||
job_kwargs = model_dump(default_params) if default_params else {}
|
||||
@@ -165,13 +167,6 @@ class SchedulerManager:
|
||||
def runtime_job(self, trigger: BaseTrigger):
|
||||
"""
|
||||
声明一个临时的、非持久化的定时任务。
|
||||
|
||||
这个任务只存在于内存中,随程序重启而消失。
|
||||
它非常适合用于插件内部的、固定的、无需用户配置的系统级定时任务。
|
||||
被此装饰器修饰的函数依然可以享受完整的依赖注入功能。
|
||||
|
||||
参数:
|
||||
trigger: 一个由 `Trigger` 工厂类创建的触发器配置对象。
|
||||
"""
|
||||
|
||||
def decorator(func: Callable[..., Coroutine]) -> Callable[..., Coroutine]:
|
||||
@@ -203,17 +198,16 @@ class SchedulerManager:
|
||||
return decorator
|
||||
|
||||
def register(
|
||||
self, plugin_name: str, params_model: type[BaseModel] | None = None
|
||||
self,
|
||||
plugin_name: str,
|
||||
params_model: type[BaseModel] | None = None,
|
||||
cli_parser: Alconna | None = None,
|
||||
default_permission: int = 5,
|
||||
default_jitter: int | None = None,
|
||||
default_spread: int | None = None,
|
||||
) -> Callable:
|
||||
"""
|
||||
注册可调度的任务函数
|
||||
|
||||
参数:
|
||||
plugin_name: 插件名称,用于标识任务。
|
||||
params_model: 参数验证模型,继承自BaseModel的类。
|
||||
|
||||
返回:
|
||||
Callable: 装饰器函数。
|
||||
"""
|
||||
|
||||
def decorator(func: Callable[..., Coroutine]) -> Callable[..., Coroutine]:
|
||||
@@ -222,6 +216,10 @@ class SchedulerManager:
|
||||
self._registered_tasks[plugin_name] = {
|
||||
"func": func,
|
||||
"model": params_model,
|
||||
"cli_parser": cli_parser,
|
||||
"default_permission": default_permission,
|
||||
"default_jitter": default_jitter,
|
||||
"default_spread": default_spread,
|
||||
}
|
||||
model_name = params_model.__name__ if params_model else "无"
|
||||
logger.debug(
|
||||
@@ -234,25 +232,14 @@ class SchedulerManager:
|
||||
def get_registered_plugins(self) -> list[str]:
|
||||
"""
|
||||
获取已注册插件列表
|
||||
|
||||
返回:
|
||||
list[str]: 已注册的插件名称列表。
|
||||
"""
|
||||
return list(self._registered_tasks.keys())
|
||||
|
||||
async def run_at(self, func: Callable[..., Coroutine], trigger: BaseTrigger) -> str:
|
||||
"""
|
||||
【新增】在未来的某个时间点,运行一个一次性的临时任务。
|
||||
|
||||
这是一个编程式API,用于动态调度一个非持久化的任务。
|
||||
|
||||
参数:
|
||||
func: 要执行的异步函数。
|
||||
trigger: 一个由 `Trigger` 工廠類創建的觸發器配置對象。
|
||||
|
||||
返回:
|
||||
str: 临时任务的唯一ID,可用于未来的管理(如取消)。
|
||||
在未来的某个时间点,运行一个一次性的临时任务。
|
||||
"""
|
||||
|
||||
job_id = f"ephemeral_runtime_{uuid.uuid4()}"
|
||||
|
||||
context = ScheduleContext(
|
||||
@@ -273,6 +260,47 @@ class SchedulerManager:
|
||||
logger.info(f"已动态调度一个临时任务 (ID: {job_id}),将在 {trigger} 触发。")
|
||||
return job_id
|
||||
|
||||
async def schedule_once(
|
||||
self,
|
||||
func: Callable[..., Coroutine],
|
||||
trigger: BaseTrigger,
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
group_id: str | None = None,
|
||||
bot_id: str | None = None,
|
||||
job_kwargs: dict | None = None,
|
||||
name: str | None = None,
|
||||
created_by: str | None = None,
|
||||
required_permission: int = 5,
|
||||
) -> "ScheduledJob | None":
|
||||
"""
|
||||
编程式API,用于动态调度一个持久化的、一次性的任务。
|
||||
"""
|
||||
if user_id and group_id:
|
||||
raise ValueError("user_id 和 group_id 不能同时提供。")
|
||||
|
||||
temp_plugin_name = f"runtime_one_off__{func.__module__}.{func.__name__}__{uuid.uuid4().hex[:8]}" # noqa: E501
|
||||
|
||||
self._registered_tasks[temp_plugin_name] = {"func": func, "model": None}
|
||||
logger.debug(f"为一次性任务动态注册临时插件: '{temp_plugin_name}'")
|
||||
|
||||
target_type = "USER" if user_id else ("GROUP" if group_id else "GLOBAL")
|
||||
target_identifier = user_id or group_id or ""
|
||||
|
||||
return await self.add_schedule(
|
||||
plugin_name=temp_plugin_name,
|
||||
target_type=target_type,
|
||||
target_identifier=target_identifier,
|
||||
trigger_type=trigger.trigger_type,
|
||||
trigger_config=model_dump(trigger, exclude={"trigger_type"}),
|
||||
job_kwargs=job_kwargs,
|
||||
bot_id=bot_id,
|
||||
name=name,
|
||||
created_by=created_by,
|
||||
required_permission=required_permission,
|
||||
is_one_off=True,
|
||||
)
|
||||
|
||||
async def add_daily_task(
|
||||
self,
|
||||
plugin_name: str,
|
||||
@@ -285,18 +313,6 @@ class SchedulerManager:
|
||||
) -> "ScheduledJob | None":
|
||||
"""
|
||||
添加每日定时任务
|
||||
|
||||
参数:
|
||||
plugin_name: 插件名称。
|
||||
group_id: 目标群组ID,None表示全局任务。
|
||||
hour: 执行小时(0-23)。
|
||||
minute: 执行分钟(0-59)。
|
||||
second: 执行秒数(0-59),默认为0。
|
||||
job_kwargs: 任务参数字典。
|
||||
bot_id: 目标Bot ID,None表示使用默认Bot。
|
||||
|
||||
返回:
|
||||
ScheduledJob | None: 创建的任务信息,失败时返回None。
|
||||
"""
|
||||
trigger_config = {
|
||||
"hour": hour,
|
||||
@@ -306,9 +322,10 @@ class SchedulerManager:
|
||||
}
|
||||
return await self.add_schedule(
|
||||
plugin_name,
|
||||
group_id,
|
||||
"cron",
|
||||
trigger_config,
|
||||
target_type="GROUP" if group_id else "GLOBAL",
|
||||
target_identifier=group_id or "",
|
||||
trigger_type="cron",
|
||||
trigger_config=trigger_config,
|
||||
job_kwargs=job_kwargs,
|
||||
bot_id=bot_id,
|
||||
)
|
||||
@@ -329,17 +346,6 @@ class SchedulerManager:
|
||||
) -> "ScheduledJob | None":
|
||||
"""
|
||||
添加间隔性定时任务
|
||||
|
||||
参数:
|
||||
plugin_name: 插件名称。
|
||||
group_id: 目标群组ID,None表示全局任务。
|
||||
weeks/days/hours/minutes/seconds: 间隔时间,至少指定一个。
|
||||
start_date: 开始时间,None表示立即开始。
|
||||
job_kwargs: 任务参数字典。
|
||||
bot_id: 目标Bot ID。
|
||||
|
||||
返回:
|
||||
ScheduledJob | None: 创建的任务信息,失败时返回None。
|
||||
"""
|
||||
trigger_config = {
|
||||
"weeks": weeks,
|
||||
@@ -352,9 +358,10 @@ class SchedulerManager:
|
||||
trigger_config = {k: v for k, v in trigger_config.items() if v}
|
||||
return await self.add_schedule(
|
||||
plugin_name,
|
||||
group_id,
|
||||
"interval",
|
||||
trigger_config,
|
||||
target_type="GROUP" if group_id else "GLOBAL",
|
||||
target_identifier=group_id or "",
|
||||
trigger_type="interval",
|
||||
trigger_config=trigger_config,
|
||||
job_kwargs=job_kwargs,
|
||||
bot_id=bot_id,
|
||||
)
|
||||
@@ -384,11 +391,7 @@ class SchedulerManager:
|
||||
return False, f"插件 '{plugin_name}' 的参数模型配置错误"
|
||||
|
||||
try:
|
||||
model_validate = getattr(params_model, "model_validate", None)
|
||||
if not model_validate:
|
||||
return False, f"插件 '{plugin_name}' 的参数模型不支持验证"
|
||||
|
||||
validated_model = model_validate(job_kwargs)
|
||||
validated_model = model_validate(params_model, job_kwargs)
|
||||
|
||||
return True, model_dump(validated_model)
|
||||
except ValidationError as e:
|
||||
@@ -400,22 +403,37 @@ class SchedulerManager:
|
||||
async def add_schedule(
|
||||
self,
|
||||
plugin_name: str,
|
||||
group_id: str | None,
|
||||
target_type: str,
|
||||
target_identifier: str,
|
||||
trigger_type: str,
|
||||
trigger_config: dict,
|
||||
job_kwargs: dict | None = None,
|
||||
bot_id: str | None = None,
|
||||
*,
|
||||
name: str | None = None,
|
||||
created_by: str | None = None,
|
||||
required_permission: int = 5,
|
||||
source: str = "USER",
|
||||
is_one_off: bool = False,
|
||||
execution_options: dict | None = None,
|
||||
) -> "ScheduledJob | None":
|
||||
"""
|
||||
添加定时任务(通用方法)
|
||||
|
||||
参数:
|
||||
plugin_name: 插件名称。
|
||||
group_id: 目标群组ID,None表示全局任务。
|
||||
trigger_type: 触发器类型,如'cron'、'interval'等。
|
||||
target_type: 目标类型 (GROUP, USER, TAG, ALL_GROUPS, GLOBAL)。
|
||||
target_identifier: 目标标识符。
|
||||
trigger_type: 触发器类型 (cron, interval, date)。
|
||||
trigger_config: 触发器配置字典。
|
||||
job_kwargs: 任务参数字典。
|
||||
bot_id: 目标Bot ID,None表示使用默认Bot。
|
||||
job_kwargs: 传递给任务函数的额外参数。
|
||||
bot_id: Bot ID约束。
|
||||
name: 任务别名。
|
||||
created_by: 创建者ID。
|
||||
required_permission: 管理此任务所需的权限。
|
||||
source: 任务来源 (USER, PLUGIN_DEFAULT)。
|
||||
is_one_off: 是否为一次性任务。
|
||||
execution_options: 任务执行的额外选项 (例如: jitter, spread)。
|
||||
|
||||
返回:
|
||||
ScheduledJob | None: 创建的任务信息,失败时返回None。
|
||||
@@ -429,51 +447,86 @@ class SchedulerManager:
|
||||
logger.error(f"任务参数校验失败: {result}")
|
||||
return None
|
||||
|
||||
search_kwargs = {"plugin_name": plugin_name, "group_id": group_id}
|
||||
if bot_id and group_id == self.ALL_GROUPS:
|
||||
options_dict = execution_options or {}
|
||||
validated_options = ExecutionOptions(**options_dict)
|
||||
|
||||
search_kwargs = {
|
||||
"plugin_name": plugin_name,
|
||||
"target_type": target_type,
|
||||
"target_identifier": target_identifier,
|
||||
}
|
||||
if bot_id:
|
||||
search_kwargs["bot_id"] = bot_id
|
||||
else:
|
||||
search_kwargs["bot_id__isnull"] = True
|
||||
|
||||
defaults = {
|
||||
"name": name,
|
||||
"trigger_type": trigger_type,
|
||||
"trigger_config": trigger_config,
|
||||
"job_kwargs": result,
|
||||
"is_enabled": True,
|
||||
"created_by": created_by,
|
||||
"required_permission": required_permission,
|
||||
"source": source,
|
||||
"is_one_off": is_one_off,
|
||||
"execution_options": model_dump(
|
||||
validated_options, exclude_none=True
|
||||
),
|
||||
}
|
||||
|
||||
defaults = {k: v for k, v in defaults.items() if v is not None}
|
||||
|
||||
schedule, created = await ScheduleRepository.update_or_create(
|
||||
defaults, **search_kwargs
|
||||
)
|
||||
APSchedulerAdapter.add_or_reschedule_job(schedule)
|
||||
|
||||
action = "设置" if created else "更新"
|
||||
action_str = "创建" if created else "更新"
|
||||
logger.info(
|
||||
f"已成功{action}插件 '{plugin_name}' 的定时任务 (ID: {schedule.id})。"
|
||||
f"已成功{action_str}任务 '{name or plugin_name}' (ID: {schedule.id})"
|
||||
)
|
||||
return schedule
|
||||
|
||||
async def get_schedules(
|
||||
self,
|
||||
plugin_name: str | None = None,
|
||||
group_id: str | None = None,
|
||||
bot_id: str | None = None,
|
||||
) -> list[ScheduledJob]:
|
||||
self, page: int | None = None, page_size: int | None = None, **filters: Any
|
||||
) -> tuple[list[ScheduledJob], int]:
|
||||
"""
|
||||
根据条件获取定时任务列表
|
||||
|
||||
参数:
|
||||
plugin_name: 插件名称,None表示不限制。
|
||||
group_id: 群组ID,None表示不限制。
|
||||
bot_id: Bot ID,None表示不限制。
|
||||
|
||||
返回:
|
||||
list[ScheduledJob]: 符合条件的任务信息列表。
|
||||
"""
|
||||
cleaned_filters = {k: v for k, v in filters.items() if v is not None}
|
||||
return await ScheduleRepository.query_schedules(
|
||||
plugin_name=plugin_name, group_id=group_id, bot_id=bot_id
|
||||
page=page, page_size=page_size, **cleaned_filters
|
||||
)
|
||||
|
||||
async def get_schedules_status_bulk(
|
||||
self, schedule_ids: list[int]
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
批量获取多个定时任务的详细状态信息
|
||||
"""
|
||||
if not schedule_ids:
|
||||
return []
|
||||
|
||||
schedules = await ScheduleRepository.filter(id__in=schedule_ids).all()
|
||||
schedule_map = {s.id: s for s in schedules}
|
||||
|
||||
statuses = []
|
||||
for schedule_id in schedule_ids:
|
||||
if schedule := schedule_map.get(schedule_id):
|
||||
status_from_scheduler = APSchedulerAdapter.get_job_status(schedule.id)
|
||||
status_dict = {
|
||||
field: getattr(schedule, field)
|
||||
for field in schedule._meta.fields_map
|
||||
}
|
||||
status_dict.update(status_from_scheduler)
|
||||
status_dict["is_enabled"] = (
|
||||
"运行中"
|
||||
if schedule_id in self._running_tasks
|
||||
else ("启用" if schedule.is_enabled else "暂停")
|
||||
)
|
||||
statuses.append(status_dict)
|
||||
|
||||
return statuses
|
||||
|
||||
async def update_schedule(
|
||||
self,
|
||||
schedule_id: int,
|
||||
@@ -483,15 +536,6 @@ class SchedulerManager:
|
||||
) -> tuple[bool, str]:
|
||||
"""
|
||||
更新定时任务配置
|
||||
|
||||
参数:
|
||||
schedule_id: 任务ID。
|
||||
trigger_type: 新的触发器类型,None表示不更新。
|
||||
trigger_config: 新的触发器配置,None表示不更新。
|
||||
job_kwargs: 新的任务参数,None表示不更新。
|
||||
|
||||
返回:
|
||||
tuple[bool, str]: (是否成功, 结果消息)。
|
||||
"""
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
if not schedule:
|
||||
@@ -533,12 +577,6 @@ class SchedulerManager:
|
||||
async def get_schedule_status(self, schedule_id: int) -> dict | None:
|
||||
"""
|
||||
获取定时任务的详细状态信息
|
||||
|
||||
参数:
|
||||
schedule_id: 定时任务的ID。
|
||||
|
||||
返回:
|
||||
dict | None: 任务详细信息字典,不存在时返回None。
|
||||
"""
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
if not schedule:
|
||||
@@ -556,7 +594,8 @@ class SchedulerManager:
|
||||
"id": schedule.id,
|
||||
"bot_id": schedule.bot_id,
|
||||
"plugin_name": schedule.plugin_name,
|
||||
"group_id": schedule.group_id,
|
||||
"target_type": schedule.target_type,
|
||||
"target_identifier": schedule.target_identifier,
|
||||
"is_enabled": status_text,
|
||||
"trigger_type": schedule.trigger_type,
|
||||
"trigger_config": schedule.trigger_config,
|
||||
@@ -567,12 +606,6 @@ class SchedulerManager:
|
||||
async def pause_schedule(self, schedule_id: int) -> tuple[bool, str]:
|
||||
"""
|
||||
暂停指定的定时任务
|
||||
|
||||
参数:
|
||||
schedule_id: 要暂停的定时任务ID。
|
||||
|
||||
返回:
|
||||
tuple[bool, str]: (是否成功, 操作结果消息)。
|
||||
"""
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
if not schedule or not schedule.is_enabled:
|
||||
@@ -586,12 +619,6 @@ class SchedulerManager:
|
||||
async def resume_schedule(self, schedule_id: int) -> tuple[bool, str]:
|
||||
"""
|
||||
恢复指定的定时任务
|
||||
|
||||
参数:
|
||||
schedule_id: 要恢复的定时任务ID。
|
||||
|
||||
返回:
|
||||
tuple[bool, str]: (是否成功, 操作结果消息)。
|
||||
"""
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
if not schedule or schedule.is_enabled:
|
||||
@@ -605,13 +632,9 @@ class SchedulerManager:
|
||||
async def trigger_now(self, schedule_id: int) -> tuple[bool, str]:
|
||||
"""
|
||||
立即手动触发指定的定时任务
|
||||
|
||||
参数:
|
||||
schedule_id: 要触发的定时任务ID。
|
||||
|
||||
返回:
|
||||
tuple[bool, str]: (是否成功, 操作结果消息)。
|
||||
"""
|
||||
from .engine import _execute_job
|
||||
|
||||
schedule = await ScheduleRepository.get_by_id(schedule_id)
|
||||
if not schedule:
|
||||
return False, f"未找到 ID 为 {schedule_id} 的定时任务。"
|
||||
@@ -619,12 +642,23 @@ class SchedulerManager:
|
||||
return False, f"插件 '{schedule.plugin_name}' 没有注册可用的定时任务。"
|
||||
|
||||
try:
|
||||
await _execute_job(schedule.id)
|
||||
await _execute_job(schedule.id, force=True)
|
||||
return True, f"已手动触发任务 (ID: {schedule.id})。"
|
||||
except Exception as e:
|
||||
logger.error(f"手动触发任务失败: {e}")
|
||||
return False, f"手动触发任务失败: {e}"
|
||||
|
||||
async def get_schedule_by_id(self, schedule_id: int) -> "ScheduledJob | None":
|
||||
"""
|
||||
通过ID获取任务对象的公共方法。
|
||||
|
||||
参数:
|
||||
schedule_id: 任务ID。
|
||||
|
||||
返回:
|
||||
ScheduledJob | None: 任务对象,不存在时返回None。
|
||||
"""
|
||||
return await ScheduleRepository.get_by_id(schedule_id)
|
||||
|
||||
|
||||
scheduler_manager = SchedulerManager()
|
||||
scheduler = scheduler_manager
|
||||
@@ -64,9 +64,9 @@ class ScheduleRepository:
|
||||
async def get_by_plugin_and_group(
|
||||
plugin_name: str, group_ids: list[str]
|
||||
) -> list[ScheduledJob]:
|
||||
"""根据插件和群组ID列表获取任务"""
|
||||
"""[DEPRECATED] 根据插件和群组ID列表获取任务"""
|
||||
return await ScheduledJob.filter(
|
||||
plugin_name=plugin_name, group_id__in=group_ids
|
||||
plugin_name=plugin_name, target_descriptor__in=group_ids
|
||||
).all()
|
||||
|
||||
@staticmethod
|
||||
@@ -77,20 +77,30 @@ class ScheduleRepository:
|
||||
return await ScheduledJob.update_or_create(defaults=defaults, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
async def query_schedules(**filters: Any) -> list[ScheduledJob]:
|
||||
async def query_schedules(
|
||||
page: int | None = None, page_size: int | None = None, **filters: Any
|
||||
) -> tuple[list[ScheduledJob], int]:
|
||||
"""
|
||||
根据任意条件查询任务列表
|
||||
|
||||
参数:
|
||||
page: 页码(从1开始)
|
||||
page_size: 每页数量
|
||||
**filters: 过滤条件,如 group_id="123", plugin_name="abc"
|
||||
|
||||
返回:
|
||||
list[ScheduledJob]: 任务列表
|
||||
tuple[list[ScheduledJob], int]: (任务列表, 总数)
|
||||
"""
|
||||
cleaned_filters = {k: v for k, v in filters.items() if v is not None}
|
||||
if not cleaned_filters:
|
||||
return await ScheduledJob.all()
|
||||
return await ScheduledJob.filter(**cleaned_filters).all()
|
||||
query = ScheduledJob.filter(**cleaned_filters)
|
||||
|
||||
total_count = await query.count()
|
||||
|
||||
if page is not None and page_size is not None:
|
||||
offset = (page - 1) * page_size
|
||||
query = query.offset(offset).limit(page_size)
|
||||
|
||||
return await query.all(), total_count
|
||||
|
||||
@staticmethod
|
||||
def filter(**kwargs: Any) -> QuerySet[ScheduledJob]:
|
||||
|
||||
@@ -1,14 +1,46 @@
|
||||
"""
|
||||
目标选择器 (Targeter)
|
||||
目标解析与选择器 (Targeting)
|
||||
|
||||
提供链式API,用于构建和执行对多个定时任务的批量操作。
|
||||
提供用于解析任务目标和批量操作目标的 ScheduleTargeter 类。
|
||||
"""
|
||||
|
||||
from collections.abc import Callable, Coroutine
|
||||
from typing import Any
|
||||
|
||||
from .adapter import APSchedulerAdapter
|
||||
from .repository import ScheduleRepository
|
||||
from nonebot.adapters import Bot
|
||||
|
||||
from zhenxun.services.tags import tag_manager
|
||||
|
||||
__all__ = [
|
||||
"ScheduleTargeter",
|
||||
"_resolve_all_groups",
|
||||
"_resolve_global_or_user",
|
||||
"_resolve_group",
|
||||
"_resolve_tag",
|
||||
"_resolve_user",
|
||||
]
|
||||
|
||||
|
||||
async def _resolve_group(target_identifier: str, bot: Bot) -> list[str | None]:
|
||||
return [target_identifier]
|
||||
|
||||
|
||||
async def _resolve_tag(target_identifier: str, bot: Bot) -> list[str | None]:
|
||||
result = await tag_manager.resolve_tag_to_group_ids(target_identifier)
|
||||
return result # type: ignore
|
||||
|
||||
|
||||
async def _resolve_user(target_identifier: str, bot: Bot) -> list[str | None]:
|
||||
return [target_identifier]
|
||||
|
||||
|
||||
async def _resolve_all_groups(target_identifier: str, bot: Bot) -> list[str | None]:
|
||||
result = await tag_manager.resolve_tag_to_group_ids("@all", bot=bot)
|
||||
return result
|
||||
|
||||
|
||||
async def _resolve_global_or_user(target_identifier: str, bot: Bot) -> list[str | None]:
|
||||
return [None]
|
||||
|
||||
|
||||
class ScheduleTargeter:
|
||||
@@ -34,6 +66,8 @@ class ScheduleTargeter:
|
||||
返回:
|
||||
list[ScheduledJob]: 符合过滤条件的任务列表。
|
||||
"""
|
||||
from .repository import ScheduleRepository
|
||||
|
||||
query = ScheduleRepository.filter(**self._filters)
|
||||
return await query.all()
|
||||
|
||||
@@ -48,12 +82,14 @@ class ScheduleTargeter:
|
||||
return f"任务 ID {self._filters['id']} 的"
|
||||
|
||||
parts = []
|
||||
if "group_id" in self._filters:
|
||||
group_id = self._filters["group_id"]
|
||||
if group_id == self._manager.ALL_GROUPS:
|
||||
if "target_descriptor" in self._filters:
|
||||
descriptor = self._filters["target_descriptor"]
|
||||
if descriptor == self._manager.ALL_GROUPS:
|
||||
parts.append("所有群组中")
|
||||
elif descriptor.startswith("tag:"):
|
||||
parts.append(f"标签 '{descriptor[4:]}' 的")
|
||||
else:
|
||||
parts.append(f"群 {group_id} 中")
|
||||
parts.append(f"群 {descriptor} 中")
|
||||
|
||||
if "plugin_name" in self._filters:
|
||||
parts.append(f"插件 '{self._filters['plugin_name']}' 的")
|
||||
@@ -111,6 +147,9 @@ class ScheduleTargeter:
|
||||
返回:
|
||||
tuple[int, str]: (成功移除的任务数量, 操作结果消息)。
|
||||
"""
|
||||
from .engine import APSchedulerAdapter
|
||||
from .repository import ScheduleRepository
|
||||
|
||||
schedules = await self._get_schedules()
|
||||
if not schedules:
|
||||
target_desc = self._generate_target_description()
|
||||
@@ -0,0 +1,145 @@
|
||||
"""
|
||||
定时任务服务的数据模型与类型定义
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from datetime import datetime
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class BaseTrigger(BaseModel):
|
||||
"""触发器配置的基类"""
|
||||
|
||||
trigger_type: str = Field(..., exclude=True)
|
||||
|
||||
|
||||
class CronTrigger(BaseTrigger):
|
||||
"""Cron 触发器配置"""
|
||||
|
||||
trigger_type: Literal["cron"] = "cron" # type: ignore
|
||||
year: int | str | None = None
|
||||
month: int | str | None = None
|
||||
day: int | str | None = None
|
||||
week: int | str | None = None
|
||||
day_of_week: int | str | None = None
|
||||
hour: int | str | None = None
|
||||
minute: int | str | None = None
|
||||
second: int | str | None = None
|
||||
start_date: datetime | str | None = None
|
||||
end_date: datetime | str | None = None
|
||||
timezone: str | None = None
|
||||
jitter: int | None = None
|
||||
|
||||
|
||||
class IntervalTrigger(BaseTrigger):
|
||||
"""Interval 触发器配置"""
|
||||
|
||||
trigger_type: Literal["interval"] = "interval" # type: ignore
|
||||
weeks: int = 0
|
||||
days: int = 0
|
||||
hours: int = 0
|
||||
minutes: int = 0
|
||||
seconds: int = 0
|
||||
start_date: datetime | str | None = None
|
||||
end_date: datetime | str | None = None
|
||||
timezone: str | None = None
|
||||
jitter: int | None = None
|
||||
|
||||
|
||||
class DateTrigger(BaseTrigger):
|
||||
"""Date 触发器配置"""
|
||||
|
||||
trigger_type: Literal["date"] = "date" # type: ignore
|
||||
run_date: datetime | str
|
||||
timezone: str | None = None
|
||||
|
||||
|
||||
class Trigger:
|
||||
"""
|
||||
一个用于创建类型安全触发器配置的工厂类。
|
||||
提供了流畅的、具备IDE自动补全功能的API。
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def cron(**kwargs) -> CronTrigger:
|
||||
"""创建一个 Cron 触发器配置。"""
|
||||
return CronTrigger(**kwargs)
|
||||
|
||||
@staticmethod
|
||||
def interval(**kwargs) -> IntervalTrigger:
|
||||
"""创建一个 Interval 触发器配置。"""
|
||||
return IntervalTrigger(**kwargs)
|
||||
|
||||
@staticmethod
|
||||
def date(**kwargs) -> DateTrigger:
|
||||
"""创建一个 Date 触发器配置。"""
|
||||
return DateTrigger(**kwargs)
|
||||
|
||||
class ExecutionOptions(BaseModel):
|
||||
"""
|
||||
封装定时任务的执行策略,包括重试和回调。
|
||||
"""
|
||||
|
||||
jitter: int | None = Field(None, description="触发时间抖动(秒)")
|
||||
spread: int | None = Field(None, description="多目标执行的分散延迟(秒)")
|
||||
concurrency_policy: Literal["ALLOW", "SKIP", "QUEUE"] = Field(
|
||||
"ALLOW", description="并发策略"
|
||||
)
|
||||
retries: int = 0
|
||||
retry_delay_seconds: int = 30
|
||||
|
||||
|
||||
class ScheduleContext(BaseModel):
|
||||
"""
|
||||
定时任务执行上下文,可通过依赖注入获取。
|
||||
"""
|
||||
|
||||
schedule_id: int = Field(..., description="数据库中的任务ID")
|
||||
plugin_name: str = Field(..., description="任务所属的插件名称")
|
||||
bot_id: str | None = Field(None, description="执行任务的Bot ID")
|
||||
group_id: str | None = Field(None, description="当前执行实例的目标群组ID")
|
||||
job_kwargs: dict = Field(default_factory=dict, description="任务配置的参数")
|
||||
|
||||
|
||||
class ExecutionPolicy(BaseModel):
|
||||
"""
|
||||
封装定时任务的执行策略,包括重试和回调。
|
||||
"""
|
||||
|
||||
retries: int = 0
|
||||
retry_delay_seconds: int = 30
|
||||
retry_backoff: bool = False
|
||||
retry_on_exceptions: list[type[Exception]] | None = None
|
||||
on_success_callback: Callable[[ScheduleContext, Any], Awaitable[None]] | None = None
|
||||
on_failure_callback: (
|
||||
Callable[[ScheduleContext, Exception], Awaitable[None]] | None
|
||||
) = None
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
|
||||
class ScheduledJobDeclaration(BaseModel):
|
||||
"""用于在启动时声明默认定时任务的内部数据模型"""
|
||||
|
||||
plugin_name: str
|
||||
group_id: str | None
|
||||
bot_id: str | None
|
||||
trigger: BaseTrigger
|
||||
job_kwargs: dict[str, Any]
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
|
||||
class EphemeralJobDeclaration(BaseModel):
|
||||
"""用于在启动时声明临时任务的内部数据模型"""
|
||||
|
||||
plugin_name: str
|
||||
func: Callable[..., Awaitable[Any]]
|
||||
trigger: BaseTrigger
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
@@ -0,0 +1,11 @@
|
||||
"""
|
||||
标签服务入口,提供 ``TagManager`` 实例并加载内置规则。
|
||||
"""
|
||||
|
||||
from .manager import TagManager
|
||||
|
||||
tag_manager = TagManager()
|
||||
|
||||
from . import filters # noqa: F401
|
||||
|
||||
__all__ = ["tag_manager"]
|
||||
@@ -0,0 +1,11 @@
|
||||
"""
|
||||
动态标签的内置过滤器集合,可通过装饰器注册到标签管理器。
|
||||
"""
|
||||
|
||||
from . import tag_manager
|
||||
|
||||
tag_manager.add_field_rule("member_count", db_field="member_count", value_type=int)
|
||||
tag_manager.add_field_rule("level", db_field="level", value_type=int)
|
||||
tag_manager.add_field_rule("status", db_field="status", value_type=bool)
|
||||
tag_manager.add_field_rule("is_super", db_field="is_super", value_type=bool)
|
||||
tag_manager.add_field_rule("group_name", db_field="group_name", value_type=str)
|
||||
@@ -0,0 +1,511 @@
|
||||
"""
|
||||
标签服务的核心实现,负责标签的增删改查与动态规则解析。
|
||||
"""
|
||||
|
||||
from collections.abc import Callable, Coroutine
|
||||
from dataclasses import dataclass
|
||||
from functools import partial
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from aiocache import Cache, cached
|
||||
from arclet.alconna import Alconna, Args
|
||||
from nonebot.adapters import Bot
|
||||
from tortoise.exceptions import IntegrityError
|
||||
from tortoise.expressions import Q
|
||||
from tortoise.transactions import in_transaction
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.group_tag import GroupTag, GroupTagLink
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
from .models import (
|
||||
ErrorResult,
|
||||
IDSetResult,
|
||||
QueryResult,
|
||||
RuleExecutionError,
|
||||
RuleExecutionResult,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class HandlerInfo:
|
||||
"""存储已注册处理器的元信息。"""
|
||||
|
||||
func: Callable[..., Coroutine[Any, Any, RuleExecutionResult]]
|
||||
alconna: Alconna
|
||||
|
||||
|
||||
def invalidate_on_change(func: Callable) -> Callable:
|
||||
"""装饰器: 在方法成功执行后自动使标签缓存失效。"""
|
||||
|
||||
async def wrapper(self: "TagManager", *args, **kwargs):
|
||||
result = await func(self, *args, **kwargs)
|
||||
await self._invalidate_cache()
|
||||
return result
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
class TagManager:
|
||||
"""群组标签管理服务。提供对群组标签的注册、解析与维护等操作。"""
|
||||
|
||||
_dynamic_handlers: ClassVar[dict[str, HandlerInfo]] = {}
|
||||
|
||||
def add_field_rule(self, name: str, db_field: str, value_type: type):
|
||||
"""
|
||||
一个便捷的快捷方式,用于快速创建一个基于 `GroupConsole` 模型字段的规则。
|
||||
它在内部使用 `register_rule`。
|
||||
"""
|
||||
from arclet.alconna import CommandMeta
|
||||
|
||||
alc = Alconna(
|
||||
name,
|
||||
Args["op", str]["value", value_type],
|
||||
meta=CommandMeta(
|
||||
fuzzy_match=True,
|
||||
compact=False,
|
||||
),
|
||||
)
|
||||
|
||||
handler = partial(self._generic_field_handler, db_field=db_field)
|
||||
|
||||
self.register_rule(alc)(handler)
|
||||
|
||||
logger.debug(f"已添加字段规则: '{name}' -> {db_field} ({value_type.__name__})")
|
||||
|
||||
async def _generic_field_handler(
|
||||
self, db_field: str, op: str, value: Any
|
||||
) -> QueryResult:
|
||||
"""所有通过 add_field_rule 添加的规则共享的处理器。"""
|
||||
op_map = {">": "__gt", ">=": "__gte", "<": "__lt", "<=": "__lte", "=": ""}
|
||||
op_lower = op.lower()
|
||||
|
||||
if op_lower == "contains":
|
||||
op_suffix = "__iposix_regex"
|
||||
elif op_lower == "in":
|
||||
op_suffix = "__in"
|
||||
value = [v.strip() for v in str(value).split(",")]
|
||||
elif op == "!=":
|
||||
return QueryResult(q_object=~Q(**{db_field: value}))
|
||||
elif op in op_map:
|
||||
op_suffix = op_map[op]
|
||||
else:
|
||||
raise RuleExecutionError(f"字段 '{db_field}' 不支持操作符: {op}")
|
||||
|
||||
q_kwargs: dict[str, Any] = {
|
||||
f"{db_field}{op_suffix}" if op_suffix else db_field: value
|
||||
}
|
||||
return QueryResult(q_object=Q(**q_kwargs))
|
||||
|
||||
def register_rule(self, alconna: Alconna):
|
||||
"""
|
||||
装饰器:注册一个完全自定义的规则处理器及其语法定义(Alconna)。
|
||||
"""
|
||||
|
||||
def decorator(handler: Callable[..., Coroutine[Any, Any, RuleExecutionResult]]):
|
||||
name = alconna.command
|
||||
if name in self._dynamic_handlers:
|
||||
logger.warning(f"动态标签规则 '{name}' 已被注册,将被覆盖。")
|
||||
self._dynamic_handlers[name] = HandlerInfo(func=handler, alconna=alconna)
|
||||
logger.debug(f"已注册动态标签规则: '{name}'")
|
||||
return handler
|
||||
|
||||
return decorator
|
||||
|
||||
async def _invalidate_cache(self):
|
||||
"""辅助函数,用于清除标签相关的缓存,确保数据一致性。"""
|
||||
cache = Cache(Cache.MEMORY, namespace="tag_service")
|
||||
await cache.clear()
|
||||
logger.debug("已清除所有群组标签缓存。")
|
||||
|
||||
@invalidate_on_change
|
||||
async def create_tag(
|
||||
self,
|
||||
name: str,
|
||||
is_blacklist: bool = False,
|
||||
description: str | None = None,
|
||||
group_ids: list[str] | None = None,
|
||||
tag_type: str = "STATIC",
|
||||
dynamic_rule: dict | str | None = None,
|
||||
) -> GroupTag:
|
||||
"""
|
||||
创建新的群组标签。
|
||||
|
||||
参数:
|
||||
name: 标签名称。
|
||||
is_blacklist: 是否为黑名单标签,黑名单标签会在最终结果中剔除关联群组。
|
||||
description: 标签描述信息。
|
||||
group_ids: 需要关联的静态群组 ID 列表,动态标签必须留空。
|
||||
tag_type: 标签类型,支持 ``STATIC`` 或 ``DYNAMIC``。
|
||||
dynamic_rule: 动态标签所使用的规则配置。
|
||||
|
||||
返回:
|
||||
新创建的 ``GroupTag`` 实例。
|
||||
"""
|
||||
if tag_type == "DYNAMIC" and group_ids:
|
||||
raise ValueError("动态标签不能在创建时关联静态群组。")
|
||||
if tag_type == "STATIC" and dynamic_rule:
|
||||
raise ValueError("静态标签不能设置动态规则。")
|
||||
async with in_transaction():
|
||||
tag = await GroupTag.create(
|
||||
name=name,
|
||||
is_blacklist=is_blacklist,
|
||||
description=description,
|
||||
tag_type=tag_type,
|
||||
dynamic_rule=dynamic_rule,
|
||||
)
|
||||
if group_ids:
|
||||
await GroupTagLink.bulk_create(
|
||||
[GroupTagLink(tag=tag, group_id=gid) for gid in group_ids]
|
||||
)
|
||||
return tag
|
||||
|
||||
@invalidate_on_change
|
||||
async def delete_tag(self, name: str) -> bool:
|
||||
"""
|
||||
删除指定标签。
|
||||
|
||||
参数:
|
||||
name: 标签名称。
|
||||
|
||||
返回:
|
||||
``True`` 表示删除成功,``False`` 表示标签不存在。
|
||||
"""
|
||||
deleted_count = await GroupTag.filter(name=name).delete()
|
||||
return deleted_count > 0
|
||||
|
||||
@invalidate_on_change
|
||||
async def add_groups_to_tag(self, name: str, group_ids: list[str]) -> int: # type: ignore
|
||||
"""
|
||||
向静态标签追加群组关联。
|
||||
"""
|
||||
tag = await GroupTag.get_or_none(name=name)
|
||||
if not tag:
|
||||
raise ValueError(f"标签 '{name}' 不存在。")
|
||||
if tag.tag_type == "DYNAMIC":
|
||||
raise ValueError("不能向动态标签手动添加群组。")
|
||||
|
||||
await GroupTagLink.bulk_create(
|
||||
[GroupTagLink(tag=tag, group_id=gid) for gid in group_ids],
|
||||
ignore_conflicts=True,
|
||||
)
|
||||
return len(group_ids)
|
||||
|
||||
@invalidate_on_change
|
||||
async def remove_groups_from_tag(self, name: str, group_ids: list[str]) -> int:
|
||||
"""从静态标签移除指定群组。"""
|
||||
tag = await GroupTag.get_or_none(name=name)
|
||||
if not tag:
|
||||
return 0
|
||||
if tag.tag_type == "DYNAMIC":
|
||||
raise ValueError("不能从动态标签手动移除群组。")
|
||||
deleted_count = await GroupTagLink.filter(
|
||||
tag=tag, group_id__in=group_ids
|
||||
).delete()
|
||||
return deleted_count
|
||||
|
||||
async def list_tags_with_counts(self) -> list[dict]:
|
||||
"""列出所有标签及其关联的群组数量。"""
|
||||
tags = await GroupTag.all().prefetch_related("groups")
|
||||
return [
|
||||
{
|
||||
"name": tag.name,
|
||||
"description": tag.description,
|
||||
"is_blacklist": tag.is_blacklist,
|
||||
"tag_type": tag.tag_type,
|
||||
"group_count": len(tag.groups),
|
||||
}
|
||||
for tag in tags
|
||||
]
|
||||
|
||||
async def get_tag_details(self, name: str, bot: Bot | None = None) -> dict | None:
|
||||
"""
|
||||
获取标签的完整信息,包括基础属性、静态群组与动态解析结果。
|
||||
|
||||
参数:
|
||||
name: 标签名称。
|
||||
bot: 可选的 ``Bot`` 实例,用于在动态标签下获取实时群组信息。
|
||||
|
||||
返回:
|
||||
包含标签详情的字典;若标签不存在则返回 ``None``。
|
||||
"""
|
||||
tag = await GroupTag.get_or_none(name=name).prefetch_related("groups")
|
||||
if not tag:
|
||||
return None
|
||||
|
||||
resolved_groups = None
|
||||
if tag.tag_type == "DYNAMIC" and bot:
|
||||
resolved_group_ids = await self.resolve_tag_to_group_ids(name, bot=bot)
|
||||
if resolved_group_ids:
|
||||
groups_from_db = await GroupConsole.filter(
|
||||
group_id__in=resolved_group_ids
|
||||
).all()
|
||||
resolved_groups = [(g.group_id, g.group_name) for g in groups_from_db]
|
||||
else:
|
||||
resolved_groups = []
|
||||
|
||||
return {
|
||||
"name": tag.name,
|
||||
"description": tag.description,
|
||||
"is_blacklist": tag.is_blacklist,
|
||||
"tag_type": tag.tag_type,
|
||||
"dynamic_rule": tag.dynamic_rule,
|
||||
"groups": [link.group_id for link in tag.groups],
|
||||
"resolved_groups": resolved_groups,
|
||||
}
|
||||
|
||||
async def _execute_rule(
|
||||
self, rule_str: str, bot: Bot | None
|
||||
) -> RuleExecutionResult:
|
||||
"""使用Alconna解析并执行单个规则。"""
|
||||
rule_str = " ".join(rule_str.split())
|
||||
|
||||
parts = rule_str.strip().split(maxsplit=1)
|
||||
if not parts:
|
||||
raise RuleExecutionError("规则字符串不能为空")
|
||||
|
||||
rule_name = parts[0]
|
||||
|
||||
handler_info = self._dynamic_handlers.get(rule_name)
|
||||
if not handler_info:
|
||||
available_rules = ", ".join(sorted(self._dynamic_handlers.keys()))
|
||||
raise RuleExecutionError(
|
||||
f"未知的规则名称: '{rule_name}'\n可用规则: {available_rules}"
|
||||
)
|
||||
|
||||
try:
|
||||
arparma = handler_info.alconna.parse(rule_str)
|
||||
if not arparma.matched:
|
||||
error_msg = (
|
||||
str(arparma.error_info) if arparma.error_info else "未知语法错误"
|
||||
)
|
||||
|
||||
args_info = []
|
||||
if handler_info.alconna.args:
|
||||
for arg in handler_info.alconna.args.argument:
|
||||
arg_name = arg.name
|
||||
arg_type = getattr(arg.value, "origin", arg.value)
|
||||
type_name = getattr(arg_type, "__name__", str(arg_type))
|
||||
args_info.append(f"<{arg_name}:{type_name}>")
|
||||
|
||||
expected_format = (
|
||||
f"{rule_name} {' '.join(args_info)}" if args_info else rule_name
|
||||
)
|
||||
|
||||
example = ""
|
||||
if rule_name in ["member_count", "level"]:
|
||||
example = f"\n示例: {rule_name} > 100"
|
||||
elif rule_name in ["status", "is_super"]:
|
||||
example = f"\n示例: {rule_name} = true"
|
||||
elif rule_name == "group_name":
|
||||
example = f"\n示例: {rule_name} contains 测试"
|
||||
|
||||
raise RuleExecutionError(
|
||||
f"规则 '{rule_name}' 参数错误: {error_msg}\n"
|
||||
f"期望格式: {expected_format}{example}"
|
||||
)
|
||||
|
||||
func_to_check = (
|
||||
handler_info.func.func
|
||||
if isinstance(handler_info.func, partial)
|
||||
else handler_info.func
|
||||
)
|
||||
|
||||
extra_kwargs = {}
|
||||
if "bot" in getattr(func_to_check, "__annotations__", {}):
|
||||
extra_kwargs["bot"] = bot
|
||||
|
||||
result = await arparma.call(handler_info.func, **extra_kwargs)
|
||||
|
||||
if not isinstance(result, RuleExecutionResult):
|
||||
raise TypeError(
|
||||
f"处理器 '{rule_name}' 返回了不支持的类型 '{type(result)}'。 "
|
||||
"必须返回 QueryResult, IDSetResult 或 ErrorResult。"
|
||||
)
|
||||
return result
|
||||
|
||||
except RuleExecutionError:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise RuleExecutionError(f"执行规则 '{rule_name}' 时发生内部错误: {e}")
|
||||
|
||||
async def _resolve_dynamic_tag(
|
||||
self, rule: dict | str, bot: Bot | None = None
|
||||
) -> set[str]:
|
||||
"""根据动态规则解析符合条件的群组 ID 集合。"""
|
||||
if isinstance(rule, dict):
|
||||
raise RuleExecutionError("动态规则必须是字符串格式。")
|
||||
|
||||
final_ids: set[str] = set()
|
||||
or_clauses = [part.strip() for part in rule.split(" or ")]
|
||||
|
||||
for or_clause in or_clauses:
|
||||
current_and_q = Q()
|
||||
current_and_ids: set[str] | None = None
|
||||
|
||||
and_rules = [part.strip() for part in or_clause.split(" and ")]
|
||||
for simple_rule in and_rules:
|
||||
try:
|
||||
result = await self._execute_rule(simple_rule, bot)
|
||||
if isinstance(result, QueryResult):
|
||||
current_and_q &= result.q_object
|
||||
elif isinstance(result, IDSetResult):
|
||||
if current_and_ids is None:
|
||||
current_and_ids = result.group_ids
|
||||
else:
|
||||
current_and_ids.intersection_update(result.group_ids)
|
||||
elif isinstance(result, ErrorResult):
|
||||
raise RuleExecutionError(result.message)
|
||||
|
||||
except Exception as e:
|
||||
raise RuleExecutionError(
|
||||
f"解析规则 '{simple_rule}' 时失败: {e}"
|
||||
) from e
|
||||
|
||||
ids_from_q: set[str] | None = None
|
||||
if current_and_q.children:
|
||||
q_filtered_groups = await GroupConsole.filter(
|
||||
current_and_q
|
||||
).values_list("group_id", flat=True)
|
||||
ids_from_q = {str(gid) for gid in q_filtered_groups}
|
||||
|
||||
if ids_from_q is not None:
|
||||
if current_and_ids is None:
|
||||
clause_result_ids = ids_from_q
|
||||
else:
|
||||
clause_result_ids = current_and_ids.intersection(ids_from_q)
|
||||
else:
|
||||
if current_and_ids is None:
|
||||
clause_result_ids = set()
|
||||
else:
|
||||
clause_result_ids = current_and_ids
|
||||
|
||||
final_ids.update(clause_result_ids)
|
||||
|
||||
if bot:
|
||||
bot_groups, _ = await PlatformUtils.get_group_list(bot)
|
||||
bot_group_ids = {g.group_id for g in bot_groups if g.group_id}
|
||||
final_ids.intersection_update(bot_group_ids)
|
||||
|
||||
return final_ids
|
||||
|
||||
@cached(ttl=300, namespace="tag_service")
|
||||
async def resolve_tag_to_group_ids(
|
||||
self, name: str, bot: Bot | None = None
|
||||
) -> list[str]:
|
||||
"""
|
||||
核心解析方法:根据标签名解析出最终的群组ID列表
|
||||
|
||||
参数:
|
||||
name: 需要解析的标签名称,特殊值 ``@all`` 表示所有群。
|
||||
bot: 可选的 ``Bot`` 实例,用于拉取最新的群信息。
|
||||
|
||||
返回:
|
||||
标签对应的群组 ID 列表。当标签不存在或无法解析时返回空列表。
|
||||
"""
|
||||
if name == "@all":
|
||||
if bot:
|
||||
all_groups, _ = await PlatformUtils.get_group_list(bot)
|
||||
return [str(g.group_id) for g in all_groups if g.group_id]
|
||||
else:
|
||||
all_group_ids = await GroupConsole.all().values_list(
|
||||
"group_id", flat=True
|
||||
)
|
||||
return [str(gid) for gid in all_group_ids]
|
||||
|
||||
tag = await GroupTag.get_or_none(name=name).prefetch_related("groups")
|
||||
if not tag:
|
||||
return []
|
||||
|
||||
if tag.tag_type == "DYNAMIC":
|
||||
if not tag.dynamic_rule or not isinstance(tag.dynamic_rule, dict | str):
|
||||
return []
|
||||
associated_groups = await self._resolve_dynamic_tag(tag.dynamic_rule, bot)
|
||||
else:
|
||||
associated_groups = {link.group_id for link in tag.groups}
|
||||
|
||||
if tag.is_blacklist:
|
||||
all_group_ids_from_db = await GroupConsole.all().values_list(
|
||||
"group_id", flat=True
|
||||
)
|
||||
return list({str(gid) for gid in all_group_ids_from_db} - associated_groups)
|
||||
else:
|
||||
return list(associated_groups)
|
||||
|
||||
@invalidate_on_change
|
||||
async def rename_tag(self, old_name: str, new_name: str) -> GroupTag:
|
||||
"""重命名已有标签"""
|
||||
if await GroupTag.exists(name=new_name):
|
||||
raise IntegrityError(f"标签 '{new_name}' 已存在。")
|
||||
tag = await GroupTag.get(name=old_name)
|
||||
tag.name = new_name
|
||||
await tag.save(update_fields=["name"])
|
||||
return tag
|
||||
|
||||
@invalidate_on_change
|
||||
async def update_tag_attributes(
|
||||
self,
|
||||
name: str,
|
||||
description: str | None = None,
|
||||
is_blacklist: bool | None = None,
|
||||
dynamic_rule: dict | str | None = None,
|
||||
) -> GroupTag:
|
||||
"""
|
||||
局部更新标签属性。
|
||||
|
||||
参数:
|
||||
name: 标签名称。
|
||||
description: 可选的新描述。
|
||||
is_blacklist: 可选的新黑名单标记。
|
||||
dynamic_rule: 可选的新动态规则配置。
|
||||
|
||||
返回:
|
||||
更新后的 ``GroupTag`` 实例。
|
||||
"""
|
||||
tag = await GroupTag.get(name=name)
|
||||
update_fields = []
|
||||
if dynamic_rule is not None:
|
||||
if tag.tag_type != "DYNAMIC":
|
||||
raise ValueError("只能为动态标签更新规则。")
|
||||
tag.dynamic_rule = dynamic_rule # type: ignore
|
||||
update_fields.append("dynamic_rule")
|
||||
if description is not None:
|
||||
tag.description = description
|
||||
update_fields.append("description")
|
||||
if is_blacklist is not None:
|
||||
tag.is_blacklist = is_blacklist
|
||||
update_fields.append("is_blacklist")
|
||||
|
||||
if update_fields:
|
||||
await tag.save(update_fields=update_fields)
|
||||
return tag
|
||||
|
||||
@invalidate_on_change
|
||||
async def set_groups_for_tag(self, name: str, group_ids: list[str]) -> int:
|
||||
"""
|
||||
覆盖设置静态标签的群组列表。
|
||||
|
||||
参数:
|
||||
name: 标签名称。
|
||||
group_ids: 需要绑定的群组 ID 列表。
|
||||
|
||||
返回:
|
||||
设置成功后的群组数量。
|
||||
"""
|
||||
tag = await GroupTag.get(name=name)
|
||||
if tag.tag_type == "DYNAMIC":
|
||||
raise ValueError("不能为动态标签设置静态群组列表。")
|
||||
async with in_transaction():
|
||||
await GroupTagLink.filter(tag=tag).delete()
|
||||
await GroupTagLink.bulk_create(
|
||||
[GroupTagLink(tag=tag, group_id=gid) for gid in group_ids],
|
||||
ignore_conflicts=True,
|
||||
)
|
||||
return len(group_ids)
|
||||
|
||||
@invalidate_on_change
|
||||
async def clear_all_tags(self) -> int:
|
||||
"""删除所有标签,并清空缓存。"""
|
||||
deleted_count = await GroupTag.all().delete()
|
||||
return deleted_count
|
||||
@@ -0,0 +1,41 @@
|
||||
"""
|
||||
动态标签的规则执行结果模型。
|
||||
"""
|
||||
|
||||
from abc import ABC
|
||||
|
||||
from pydantic import BaseModel
|
||||
from tortoise.expressions import Q
|
||||
|
||||
|
||||
class RuleExecutionError(ValueError):
|
||||
"""在规则执行期间,由处理器返回的、可向用户展示的错误。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class RuleExecutionResult(BaseModel, ABC):
|
||||
"""规则执行结果的抽象基类。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class QueryResult(RuleExecutionResult):
|
||||
"""表示数据库查询条件的结果。"""
|
||||
|
||||
q_object: Q
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
|
||||
class IDSetResult(RuleExecutionResult):
|
||||
"""表示一组群组ID的结果。"""
|
||||
|
||||
group_ids: set[str]
|
||||
|
||||
|
||||
class ErrorResult(RuleExecutionResult):
|
||||
"""表示一个可向用户显示的错误。"""
|
||||
|
||||
message: str
|
||||
Reference in New Issue
Block a user