ADK-gateway/backend/app/services/task_service.py
2026-08-05 16:33:27 +08:00

119 lines
4.9 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""任务池服务:受理、状态流转、幂等、超时失败。"""
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