"""任务池服务:受理、状态流转、幂等、超时失败。""" import logging import time import redis.asyncio as aioredis from app.config import settings from app.constants import AgentStatus, TaskStatus from app.models.schemas import TaskInfo, TaskSubmit from app.repository.agent_repo import AgentRepo from app.repository.task_repo import TaskRepo from app.services.notifier import TaskNotifier logger = logging.getLogger(__name__) class TaskService: def __init__(self, redis: aioredis.Redis, repo: TaskRepo | None = None): self.redis = redis self.repo = repo or TaskRepo(redis) self.agent_repo = AgentRepo(redis) async def _release_if_running(self, task: TaskInfo) -> None: """若任务处于 running 且绑定 Agent,释放其算力。""" if task.status == TaskStatus.RUNNING and task.agent_id: await self.agent_repo.adjust_load(task.agent_id, -1) async def submit(self, submit: TaskSubmit) -> TaskInfo: """受理任务:生成/复用 RequestID,幂等写入任务池。""" task = TaskInfo( request_id=submit.request_id or submit.request_id or _new_id(), task_type=submit.task_type, task_tags=submit.task_tags, description=submit.description, payload=submit.payload, cli_session_id=submit.cli_session_id, timeout=submit.timeout or settings.default_task_timeout, status=TaskStatus.PENDING, ) created = await self.repo.create(task) if not created: # 幂等:返回已有任务 existing = await self.repo.get(task.request_id) if existing: return existing logger.info("task submitted request_id=%s type=%s", task.request_id, task.task_type) return task async def get(self, request_id: str) -> TaskInfo | None: return await self.repo.get(request_id) async def list(self, status: TaskStatus | None = None, limit: int = 100) -> list[TaskInfo]: return await self.repo.list(status=status, limit=limit) async def cancel(self, request_id: str) -> bool: """取消任务(仅 pending/running 可取消)。 - running 且绑定 agent:置 agent 为 stopping,并通过同一 cli_session_id 的 SSE 通道下发 task_stop 停止指令;agent 停止后回传结果,由 on_result 终态 保护将其置回 ready。 - pending:无需下发,直接标记失败。 """ task = await self.repo.get(request_id) if not task: return False if task.status in (TaskStatus.SUCCESS, TaskStatus.FAILED): return False was_running = task.status == TaskStatus.RUNNING await self._release_if_running(task) await self.repo.update(request_id, status=TaskStatus.FAILED, error_info="cancelled by user") await self.repo.mark_finished(request_id) if was_running and task.agent_id: # 通知 agent 停止当前任务 await self.agent_repo.update_status(task.agent_id, AgentStatus.STOPPING) await TaskNotifier(self.redis).publish_task_stop(task.cli_session_id, request_id) if task.cli_session_id: finished = await self.repo.get(request_id) if finished: await TaskNotifier(self.redis).publish_task_done(finished) logger.info("task cancelled request_id=%s", request_id) return True async def reset(self, request_id: str) -> bool: """重置任务状态为 pending(重新调度)。""" task = await self.repo.get(request_id) if not task: return False await self.repo.update(request_id, status=TaskStatus.PENDING, agent_id=None, error_info=None, progress=0) await self.repo.mark_pending(request_id) logger.info("task reset request_id=%s", request_id) return True async def check_timeouts(self) -> int: """扫描 running 任务,超时自动标记失败。""" now = time.time() failed = 0 for rid in await self.repo.running_tasks(): task = await self.repo.get(rid) if not task: continue if task.timeout and (now - task.create_time) > task.timeout: if task.status == TaskStatus.RUNNING and task.agent_id: await self.agent_repo.mark_idle(task.agent_id) await self._release_if_running(task) await self.repo.update(rid, status=TaskStatus.FAILED, error_info="timeout") await self.repo.mark_finished(rid) if task.cli_session_id: finished = await self.repo.get(rid) if finished: await TaskNotifier(self.redis).publish_task_done(finished) failed += 1 logger.info("task timeout request_id=%s", rid) return failed def _new_id() -> str: import uuid return uuid.uuid4().hex