"""任务池 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