ADK-gateway/backend/app/repository/task_repo.py
2026-08-04 17:20:08 +08:00

166 lines
5.8 KiB
Python

"""任务池 Redis 存储层。"""
import time
from typing import Any
import redis.asyncio as aioredis
from app import constants as C
from app.config import settings
from app.constants import TaskStatus
from app.models.schemas import TaskInfo
class TaskRepo:
def __init__(self, redis: aioredis.Redis):
self.redis = redis
@staticmethod
def _info_key(request_id: str) -> str:
return C.TASK_INFO_KEY.format(request_id=request_id)
# ---------- 写入 ----------
async def create(self, task: TaskInfo) -> bool:
"""录入任务(幂等),返回是否新建。"""
key = self._info_key(task.request_id)
existed = await self.redis.exists(key)
if existed:
return False
await self.redis.hset(
key,
mapping={
"request_id": task.request_id,
"task_type": task.task_type,
"task_tags": ",".join(task.task_tags),
"description": task.description or "",
"payload": _json(task.payload),
"status": task.status.value,
"agent_id": task.agent_id or "",
"cli_session_id": task.cli_session_id or "",
"create_time": str(task.create_time),
"timeout": str(task.timeout),
"progress": str(task.progress),
"result": _json(task.result),
"error_info": task.error_info or "",
},
)
await self.redis.expire(key, settings.task_ttl)
await self.redis.sadd(C.TASK_PENDING_SET, task.request_id)
return True
async def update(
self,
request_id: str,
*,
status: TaskStatus | None = None,
agent_id: str | None = None,
progress: int | None = None,
result: dict | None = None,
error_info: str | None = None,
) -> bool:
key = self._info_key(request_id)
if not await self.redis.exists(key):
return False
mapping: dict[str, Any] = {}
if status is not None:
mapping["status"] = status.value
if agent_id is not None:
mapping["agent_id"] = agent_id
if progress is not None:
mapping["progress"] = str(progress)
if result is not None:
mapping["result"] = _json(result)
if error_info is not None:
mapping["error_info"] = error_info
if mapping:
await self.redis.hset(key, mapping=mapping)
return True
# ---------- 状态集合维护 ----------
async def mark_pending(self, request_id: str) -> None:
await self.redis.sadd(C.TASK_PENDING_SET, request_id)
await self.redis.srem(C.TASK_RUNNING_SET, request_id)
async def mark_running(self, request_id: str) -> None:
await self.redis.sadd(C.TASK_RUNNING_SET, request_id)
await self.redis.srem(C.TASK_PENDING_SET, request_id)
async def mark_finished(self, request_id: str) -> None:
await self.redis.srem(C.TASK_PENDING_SET, request_id)
await self.redis.srem(C.TASK_RUNNING_SET, request_id)
async def pending_tasks(self) -> list[str]:
return list(await self.redis.smembers(C.TASK_PENDING_SET))
async def running_tasks(self) -> list[str]:
return list(await self.redis.smembers(C.TASK_RUNNING_SET))
async def remove_pending(self, request_id: str) -> None:
await self.redis.srem(C.TASK_PENDING_SET, request_id)
# ---------- 读取 ----------
async def get(self, request_id: str) -> TaskInfo | None:
raw = await self.redis.hgetall(self._info_key(request_id))
if not raw:
return None
return await self._to_task(raw)
async def _to_task(self, raw: dict) -> TaskInfo:
return TaskInfo(
request_id=raw.get("request_id", ""),
task_type=raw.get("task_type", ""),
task_tags=[t for t in raw.get("task_tags", "").split(",") if t],
description=raw.get("description") or None,
payload=_unjson(raw.get("payload")),
status=TaskStatus(raw.get("status", TaskStatus.PENDING.value)),
agent_id=raw.get("agent_id") or None,
cli_session_id=raw.get("cli_session_id") or None,
create_time=float(raw.get("create_time", 0) or 0),
timeout=int(raw.get("timeout", 0) or 0),
progress=int(raw.get("progress", 0) or 0),
result=_unjson(raw.get("result")),
error_info=raw.get("error_info") or None,
)
async def list(self, status: TaskStatus | None = None, limit: int = 100) -> list[TaskInfo]:
"""按状态过滤返回任务列表(含终态任务,通过 scan 全量扫描 task:info:*)。"""
ids = await self.scan_ids()
out = []
for rid in ids:
t = await self.get(rid)
if not t:
continue
if status is not None and t.status != status:
continue
out.append(t)
# 按创建时间倒序
out.sort(key=lambda x: x.create_time, reverse=True)
return out[:limit]
async def scan_ids(self) -> list[str]:
"""扫描所有 task:info:* 键,返回 RequestID 列表。"""
prefix = C.TASK_INFO_KEY.replace("{request_id}", "")
ids = []
async for key in self.redis.scan_iter(match=f"{prefix}*", count=1000):
ids.append(key[len(prefix):])
return ids
async def delete(self, request_id: str) -> None:
await self.redis.delete(self._info_key(request_id))
await self.mark_finished(request_id)
def _json(obj: Any) -> str:
import json
return json.dumps(obj, ensure_ascii=False) if obj is not None else ""
def _unjson(s: str | None) -> Any:
import json
if not s:
return None
try:
return json.loads(s)
except json.JSONDecodeError:
return None