"""通用任务调度器 —— CAS 锁、并发控制、超时检测、后台执行 重构要点(v2): - 不再继承 TaskHandler(scheduler 不是 handler,是编排层) - 接收 TaskRegistry 实例,消除全局 _TASK_HANDLER_REGISTRY - 使用结构化 LogRecord 替代自由 dict """ import asyncio import json import time from datetime import datetime, timedelta from typing import Any, Callable, Dict, List, Optional from supply.domain.task.config import TaskConfig, TaskStatus, get_default_task_config from supply.domain.task.base import TaskRegistry from supply.domain.task.exceptions import ( TaskConcurrencyError, TaskValidationError, ) from supply.infra import TaskScheduleResponse AlertCallback = Callable[[str, dict], Any] def _validate_table_name(table_name: str) -> str: cleaned = table_name.strip().replace(" ", "_") if not cleaned or "--" in cleaned or ";" in cleaned.lower(): raise TaskValidationError(f"非法的表名: {table_name}") return cleaned def _validate_task_name(task_name: Any) -> str: task_name = str(task_name).strip() if not task_name: raise TaskValidationError("task_name 不能为空") return task_name def _format_error(e: Exception) -> str: import traceback return "".join( traceback.format_exception(type(e), e, e.__traceback__) )[-1000:] class TaskScheduler: """通用任务调度器 —— CAS 锁、并发控制、超时检测、后台执行 使用方式: registry = TaskRegistry() # ... 注册 handler ... scheduler = TaskScheduler( data={"task_name": "my_crawl", "date_string": "2026-06-30"}, registry=registry, log_service=log, db_client=db, trace_id="trace-123", task_configs={"my_crawl": TaskConfig(timeout=600, max_concurrent=3)}, table_name="my_task_table", alert_callback=my_alert_func, ) result = await scheduler.deal() """ def __init__( self, data: dict, registry: TaskRegistry, *, log_service: Any = None, db_client: Any = None, trace_id: str = "", task_configs: Optional[Dict[str, TaskConfig]] = None, table_name: str = "task_manager", alert_callback: Optional[AlertCallback] = None, ): self.data = data self.registry = registry self.log_service = log_service self.db_client = db_client self.trace_id = trace_id self.table = _validate_table_name(table_name) self.task_configs = task_configs or {} self._alert_callback = alert_callback # ==================== 日志 ==================== async def _log(self, event: str, **kwargs) -> None: """结构化日志推送""" if self.log_service is None: return try: from supply.infra.observability.log_schema import ( LogRecord, LogLevel, LogCategory, ) record = LogRecord( level=LogLevel.INFO, category=LogCategory.TASK, event=event, trace_id=self.trace_id, task_name=self.data.get("task_name", ""), extra=kwargs, ) await self.log_service.log(record) except Exception: pass async def _alert(self, title: str, detail: dict) -> None: if self._alert_callback: try: await self._alert_callback(title, detail) except Exception: pass def get_task_config(self, task_name: str) -> TaskConfig: return self.task_configs.get(task_name, get_default_task_config()) # ==================== 数据库操作 ==================== async def _insert_or_ignore_task(self, task_name: str, date_str: str) -> None: query = f""" INSERT IGNORE INTO {self.table} (date_string, task_name, start_timestamp, task_status, trace_id, data) VALUES (%s, %s, %s, %s, %s, %s) """ await self.db_client.async_save( query=query, params=( date_str, task_name, int(time.time()), TaskStatus.INIT, self.trace_id, json.dumps(self.data, ensure_ascii=False), ), ) async def _try_lock_task(self) -> bool: query = f""" UPDATE {self.table} SET task_status = %s WHERE trace_id = %s AND task_status = %s """ result = await self.db_client.async_save( query=query, params=(TaskStatus.PROCESSING, self.trace_id, TaskStatus.INIT), ) return bool(result) async def _release_task(self, status: int) -> None: query = f""" UPDATE {self.table} SET task_status = %s, finish_timestamp = %s WHERE trace_id = %s AND task_status = %s """ await self.db_client.async_save( query=query, params=(status, int(time.time()), self.trace_id, TaskStatus.PROCESSING), ) async def _get_processing_tasks(self, task_name: str) -> List[Dict[str, Any]]: query = f""" SELECT trace_id, start_timestamp, data FROM {self.table} WHERE task_status = %s AND task_name = %s """ rows = await self.db_client.async_fetch( query=query, params=(TaskStatus.PROCESSING, task_name), ) return rows or [] # ==================== 并发/超时检查 ==================== async def _check_task_concurrency_and_timeout(self, task_name: str) -> None: processing_tasks = await self._get_processing_tasks(task_name) if not processing_tasks: return config = self.get_task_config(task_name) current_time = int(time.time()) # 超时检测:自动释放超时任务 timeout_tasks = [ t for t in processing_tasks if current_time - t["start_timestamp"] > config.timeout ] if timeout_tasks: await self._log( "task_timeout_detected", timeout_count=len(timeout_tasks), timeout_traces=[t["trace_id"] for t in timeout_tasks], ) await self._alert( title=f"Task Timeout Alert: {task_name}", detail={ "task_name": task_name, "timeout_count": len(timeout_tasks), "timeout_threshold": config.timeout, "timeout_tasks": [ { "trace_id": t["trace_id"], "running_time": current_time - t["start_timestamp"], } for t in timeout_tasks ], }, ) for t in timeout_tasks: await self._force_release_task(t["trace_id"], TaskStatus.FAILED) # 并发限制 active_tasks = [ t for t in processing_tasks if current_time - t["start_timestamp"] <= config.timeout ] if len(active_tasks) >= config.max_concurrent: await self._log( "task_concurrency_limit", current_count=len(active_tasks), max_concurrent=config.max_concurrent, ) await self._alert( title=f"Task Concurrency Limit: {task_name}", detail={ "task_name": task_name, "current_count": len(active_tasks), "max_concurrent": config.max_concurrent, "active_tasks": [t["trace_id"] for t in active_tasks], }, ) raise TaskConcurrencyError( f"Task {task_name} has reached max concurrency limit " f"({len(active_tasks)}/{config.max_concurrent})", task_name=task_name, ) # ==================== 任务执行 ==================== async def _run_with_guard( self, task_name: str, date_str: str, handler_callable: Callable, ) -> dict: """带保护的任务执行: 并发检查 → CAS 锁 → 后台执行 → 释放""" # 1. 并发/超时检查 try: await self._check_task_concurrency_and_timeout(task_name) except TaskConcurrencyError as e: return TaskScheduleResponse.fail("5005", str(e)) # 2. 创建任务记录并 CAS 获取锁 await self._insert_or_ignore_task(task_name, date_str) if not await self._try_lock_task(): return TaskScheduleResponse.fail("5001", "Task is already processing") # 3. 后台执行 async def _task_wrapper(): status = TaskStatus.FAILED config = self.get_task_config(task_name) start_time = time.time() try: await self._log("task_started") result = await handler_callable() status = result if isinstance(result, int) else TaskStatus.SUCCESS duration = time.time() - start_time await self._log( "task_completed", status=status, duration_ms=int(duration * 1000), ) except Exception as e: duration = time.time() - start_time error_detail = _format_error(e) await self._log( "task_failed", error=error_detail, duration_ms=int(duration * 1000), ) if config.alert_on_failure: await self._alert( title=f"Task Failed: {task_name}", detail={ "task_name": task_name, "trace_id": self.trace_id, "error": error_detail, "duration": duration, }, ) finally: await self._release_task(status) asyncio.create_task(_task_wrapper(), name=f"{task_name}_{self.trace_id}") return TaskScheduleResponse.success( task_name=task_name, data={ "code": 0, "message": "Task started successfully", "trace_id": self.trace_id, }, ) # ==================== 任务管理接口 ==================== async def get_task_status(self, trace_id: Optional[str] = None) -> Optional[Dict[str, Any]]: tid = trace_id or self.trace_id query = f"SELECT * FROM {self.table} WHERE trace_id = %s" return await self.db_client.async_fetch_one(query, params=(tid,)) async def cancel_task(self, trace_id: Optional[str] = None) -> bool: tid = trace_id or self.trace_id query = f""" UPDATE {self.table} SET task_status = %s, finish_timestamp = %s WHERE trace_id = %s AND task_status IN (%s, %s) """ result = await self.db_client.async_save( query, params=( TaskStatus.FAILED, int(time.time()), tid, TaskStatus.INIT, TaskStatus.PROCESSING, ), ) if result: await self._log("task_cancelled", trace_id=tid) return bool(result) async def retry_task(self, trace_id: Optional[str] = None) -> bool: tid = trace_id or self.trace_id query = f""" UPDATE {self.table} SET task_status = %s, start_timestamp = %s, finish_timestamp = NULL WHERE trace_id = %s """ result = await self.db_client.async_save( query, params=(TaskStatus.INIT, int(time.time()), tid) ) if result: await self._log("task_retried", trace_id=tid) return bool(result) async def _force_release_task(self, trace_id: str, status: int) -> None: query = f""" UPDATE {self.table} SET task_status = %s, finish_timestamp = %s WHERE trace_id = %s """ await self.db_client.async_save( query, params=(status, int(time.time()), trace_id) ) await self._log("task_force_released", trace_id=trace_id, status=status) # ==================== 主入口 ==================== async def deal(self) -> dict: """任务调度主入口""" task_name = self.data.get("task_name") if not task_name: return TaskScheduleResponse.fail("4003", "task_name is required") try: task_name = _validate_task_name(task_name) except TaskValidationError as e: return TaskScheduleResponse.fail("4003", str(e)) date_str = self.data.get("date_string") or ( datetime.utcnow() + timedelta(hours=8) ).strftime("%Y-%m-%d") # 从 TaskRegistry 查找 handler handler_cls = self.registry.get(task_name) if handler_cls is None: return TaskScheduleResponse.fail( "4001", f"Unknown task: {task_name}. " f"Available tasks: {', '.join(self.registry.list_names())}", ) # 实例化 handler handler = handler_cls(log_service=self.log_service, db_client=self.db_client) return await self._run_with_guard( task_name, date_str, lambda: handler.execute( type("Task", (), {"task_name": task_name, "params": self.data, "trace_id": self.trace_id})() ), )