117 lines
5.7 KiB
Python
117 lines
5.7 KiB
Python
"""通信中转服务:正向下发任务指令、反向回传结果/日志,RequestID+AgentID 关联。"""
|
||
import logging
|
||
|
||
import httpx
|
||
import redis.asyncio as aioredis
|
||
|
||
from app.config import settings
|
||
from app.constants import AgentStatus, TaskStatus
|
||
from app.models.schemas import TaskResult
|
||
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 RelayService:
|
||
"""网关作为唯一通信中枢:正向把任务指令推送到 Agent 端点,反向接收结果回传。"""
|
||
|
||
def __init__(self, redis: aioredis.Redis, task_repo: TaskRepo | None = None):
|
||
self.redis = redis
|
||
self.task_repo = task_repo or TaskRepo(redis)
|
||
self.agent_repo = AgentRepo(redis)
|
||
|
||
async def dispatch_command(self, agent_id: str, request_id: str, payload: dict) -> bool:
|
||
"""正向:向 Agent 真实 HTTP 推送任务指令(POST {agent.endpoint}/tasks/{request_id})。
|
||
|
||
成功返回 True;Agent 不存在 / 端点缺失 / 网络失败 / 非 202 均返回 False,
|
||
由调度器回退任务状态并释放负载。
|
||
"""
|
||
agent = await self.agent_repo.get(agent_id)
|
||
if not agent:
|
||
logger.warning("relay dispatch failed: agent not found agent=%s request=%s", agent_id, request_id)
|
||
return False
|
||
if not agent.endpoint:
|
||
logger.warning("relay dispatch failed: endpoint empty agent=%s request=%s", agent_id, request_id)
|
||
return False
|
||
task = await self.task_repo.get(request_id)
|
||
url = f"{agent.endpoint.rstrip('/')}/tasks/{request_id}"
|
||
body = {
|
||
"auth": settings.gateway_auth,
|
||
"request_id": request_id,
|
||
"payload": payload,
|
||
"cli_session_id": getattr(task, "cli_session_id", None) if task else None,
|
||
"task_type": getattr(task, "task_type", None) if task else None,
|
||
"task_tags": getattr(task, "task_tags", None) if task else None,
|
||
"timeout": getattr(task, "timeout", None) if task else None,
|
||
}
|
||
try:
|
||
async with httpx.AsyncClient(timeout=5.0) as client:
|
||
resp = await client.post(url, json=body)
|
||
except httpx.HTTPError as e:
|
||
logger.error("relay dispatch network error agent=%s request=%s err=%s", agent_id, request_id, e)
|
||
await self._log(agent_id, request_id, "dispatch", f"push failed: {e}")
|
||
return False
|
||
if resp.status_code != 202:
|
||
logger.warning("relay dispatch http %s agent=%s request=%s body=%s",
|
||
resp.status_code, agent_id, request_id, resp.text[:200])
|
||
await self._log(agent_id, request_id, "dispatch", f"push failed http {resp.status_code}")
|
||
return False
|
||
await self._log(agent_id, request_id, "dispatch", f"command pushed to {agent.endpoint}")
|
||
logger.info("relay dispatch ok agent=%s request=%s url=%s", agent_id, request_id, url)
|
||
return True
|
||
|
||
_TERMINAL = {TaskStatus.SUCCESS, TaskStatus.FAILED}
|
||
|
||
async def on_result(self, result: TaskResult) -> bool:
|
||
"""反向:Agent 回传结果,更新任务状态并记录日志。
|
||
|
||
终态保护:任务已处于终态(如已取消置 failed)时,忽略本次任务状态覆盖,
|
||
但匹配的 Agent 仍会释放负载并置回 ready。
|
||
"""
|
||
task = await self.task_repo.get(result.request_id)
|
||
if not task:
|
||
logger.warning("result for unknown task request=%s", result.request_id)
|
||
return False
|
||
if task.agent_id and task.agent_id != result.agent_id:
|
||
logger.warning("result agent mismatch request=%s expected=%s got=%s",
|
||
result.request_id, task.agent_id, result.agent_id)
|
||
return False
|
||
already_terminal = task.status in self._TERMINAL
|
||
if not already_terminal:
|
||
await self.task_repo.update(
|
||
result.request_id,
|
||
status=result.status,
|
||
progress=result.progress,
|
||
result=result.result,
|
||
error_info=result.error_info,
|
||
)
|
||
await self.task_repo.mark_finished(result.request_id)
|
||
# 释放 Agent 算力(任务确实绑定在该 Agent 时才释放)
|
||
if task.agent_id == result.agent_id:
|
||
await self.agent_repo.adjust_load(result.agent_id, -1)
|
||
# Agent 空闲,置回就绪(若已被手动置为不可用则保持不可用,不会复活)
|
||
if task.agent_id:
|
||
await self.agent_repo.mark_idle(task.agent_id)
|
||
await self._log(result.agent_id, result.request_id, "result", f"result received status={result.status.value}")
|
||
# 任务到达终态,通知对应 CLI 会话(若携带了 cli_session_id)
|
||
if not already_terminal and task.cli_session_id:
|
||
finished = await self.task_repo.get(result.request_id)
|
||
if finished:
|
||
await TaskNotifier(self.redis).publish_task_done(finished)
|
||
logger.info("task result request=%s agent=%s status=%s", result.request_id, result.agent_id, result.status.value)
|
||
return True
|
||
|
||
async def on_progress(self, request_id: str, agent_id: str, progress: int) -> bool:
|
||
await self.task_repo.update(request_id, progress=progress)
|
||
await self._log(agent_id, request_id, "progress", f"progress {progress}%")
|
||
return True
|
||
|
||
async def _log(self, agent_id: str, request_id: str, action: str, message: str) -> None:
|
||
from app.repository.log_repo import LogRepo
|
||
|
||
await LogRepo(self.redis).append(
|
||
source="gateway", scope="task", message=message,
|
||
request_id=request_id, agent_id=agent_id,
|
||
) |