166 lines
5.8 KiB
Python
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 |