106 lines
4.2 KiB
Python
106 lines
4.2 KiB
Python
"""任务池服务:受理、状态流转、幂等、超时失败。"""
|
||
import logging
|
||
import time
|
||
|
||
import redis.asyncio as aioredis
|
||
|
||
from app.config import settings
|
||
from app.constants import 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 可取消)。"""
|
||
task = await self.repo.get(request_id)
|
||
if not task:
|
||
return False
|
||
if task.status in (TaskStatus.SUCCESS, TaskStatus.FAILED):
|
||
return False
|
||
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 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:
|
||
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 |