# services/redis_task_manager.py """Redis 任务管理器 - 任务状态热缓存(D7:不再有进程内存回退)。 存储格式:Redis Hash(field -> JSON 字符串)。 - update_task 走 HSET 字段级原子更新,消除旧 get->merge->set 三步竞态 (后台处理流程与导出端点并发写同一任务时丢更新); - 进度 tick 只重写变化字段,不再全量重写整个任务 blob; - 兼容读旧 string 格式(升级前写入的在途任务),新写入一律 Hash; - **PG 是任务状态单一事实源**:Redis 不可用时本管理器不再降级进程内 dict (多副本下各进程内存互相不可见,造成同一任务不同副本读到不同状态), 而是 no-op / 返回 None——状态查询路径(TaskQueryService)自然落到 PG。 """ import json from typing import Dict, Any, Optional from datetime import datetime import redis.asyncio as aioredis from shared.config.settings import settings from shared.utils.logger import get_logger logger = get_logger(__name__) class RedisTaskManager: """基于 Redis Hash 的任务状态管理""" _instance: Optional["RedisTaskManager"] = None def __init__(self): self._redis: Optional[aioredis.Redis] = None self._prefix = "moldinsight:task:" self._ttl = 86400 * 7 # 任务默认保留 7 天 self._connected = False @classmethod def get_instance(cls) -> "RedisTaskManager": if cls._instance is None: cls._instance = RedisTaskManager() return cls._instance async def connect(self): """连接 Redis(配置统一来自 shared.config.settings,不再硬编码主机名)""" if self._connected and self._redis: return host = settings.REDIS_HOST port = settings.REDIS_PORT password = settings.REDIS_PASSWORD db = settings.REDIS_DB try: self._redis = aioredis.Redis( host=host, port=port, password=password if password else None, db=db, decode_responses=True, socket_connect_timeout=5, socket_timeout=5, retry_on_timeout=True, ) # 测试连接 await self._redis.ping() self._connected = True logger.info(f"Redis 连接成功: {host}:{port}") except Exception as e: logger.error(f"Redis 连接失败: {e},任务状态将使用内存回退") self._redis = None self._connected = False async def reconnect(self): """强制重新连接 用于事件循环会变更的场景(如 Celery 每个任务经 asyncio.run 创建新循环): redis.asyncio 客户端绑定到创建它的循环,旧循环关闭后客户端失效, 必须在新循环中重建客户端才能继续使用。 """ # 丢弃绑定在旧(已关闭)循环上的客户端,connect() 会重建 self._redis = None self._connected = False await self.connect() async def disconnect(self): """断开 Redis 连接""" if self._redis: await self._redis.aclose() self._redis = None self._connected = False logger.info("Redis 连接已断开") @property def is_connected(self) -> bool: return self._connected and self._redis is not None @property def redis_client(self) -> aioredis.Redis: """暴露底层客户端(batch 元数据等非任务结构数据使用)。 未连接时抛出明确错误,而不是让调用方踩 AttributeError。 """ if not self.is_connected or self._redis is None: raise RuntimeError("Redis 未连接,无法直接访问 redis_client") return self._redis # ---- 内部工具 ---- def _key(self, task_id: str) -> str: return f"{self._prefix}{task_id}" @staticmethod def _dump_mapping(data: Dict[str, Any]) -> Dict[str, str]: """把任务 dict 序列化为 Hash mapping(field -> JSON 字符串)""" serializable = RedisTaskManager._make_serializable(data) return {k: json.dumps(v, ensure_ascii=False) for k, v in serializable.items()} async def _load_hash(self, key: str) -> Optional[Dict[str, Any]]: raw = await self._redis.hgetall(key) if not raw: return None result = {} for field, value in raw.items(): try: result[field] = json.loads(value) except (json.JSONDecodeError, TypeError): result[field] = value return result async def _load_any(self, key: str) -> Optional[Dict[str, Any]]: """读取任务数据,自动识别 Hash(新)与 string(旧)格式。""" key_type = await self._redis.type(key) if key_type == "hash": return await self._load_hash(key) if key_type == "string": legacy = await self._redis.get(key) if not legacy: return None try: return json.loads(legacy) except json.JSONDecodeError: logger.warning(f"任务数据解析失败(旧 string 格式): {key}") return None return None # ---- 公共接口 ---- async def set_task(self, task_id: str, data: Dict[str, Any], ttl: Optional[int] = None): """整包写入任务数据(Hash,覆盖旧值,含旧 string 格式清理)。 Redis 不可用时 no-op:任务状态事实源在 PG,缓存缺失不影响正确性。 """ if not self.is_connected: return effective_ttl = ttl or self._ttl mapping = self._dump_mapping(data) try: key = self._key(task_id) # DEL 先清掉可能存在的旧 string/Hash,保证覆盖语义 pipe = self._redis.pipeline() pipe.delete(key) pipe.hset(key, mapping=mapping) pipe.expire(key, effective_ttl) await pipe.execute() except Exception as e: logger.warning(f"Redis 写入失败(任务状态以 PG 为准): task={task_id}, {e}") async def get_task(self, task_id: str) -> Optional[Dict[str, Any]]: """获取任务数据(Hash / 旧 string 兼容)。 Redis 不可用 / 未命中返回 None,调用方落到 PG 路径。 """ if not self.is_connected: return None try: return await self._load_any(self._key(task_id)) except Exception as e: logger.warning(f"Redis 读取失败(任务状态以 PG 为准): task={task_id}, {e}") return None async def update_task(self, task_id: str, updates: Dict[str, Any]): """字段级原子更新(HSET),无读改写竞态。 兼容旧 string 格式:先迁移为 Hash 再更新。 Redis 不可用时 no-op(状态事实源在 PG)。 """ if not self.is_connected: return mapping = self._dump_mapping(updates) try: key = self._key(task_id) key_type = await self._redis.type(key) if key_type == "none": logger.warning(f"任务 {task_id} 不存在,无法更新") return if key_type == "string": # 旧格式迁移:string -> Hash legacy = await self._redis.get(key) try: base = json.loads(legacy) if legacy else {} except json.JSONDecodeError: base = {} base.update(mapping) pipe = self._redis.pipeline() pipe.delete(key) pipe.hset(key, mapping=self._dump_mapping(base)) pipe.expire(key, self._ttl) await pipe.execute() return await self._redis.hset(key, mapping=mapping) await self._redis.expire(key, self._ttl) except Exception as e: logger.warning(f"Redis 更新失败(任务状态以 PG 为准): task={task_id}, {e}") async def delete_task(self, task_id: str): """删除任务(DEL 对 Hash/string 均有效);Redis 不可用时 no-op""" if not self.is_connected: return try: await self._redis.delete(self._key(task_id)) except Exception as e: logger.warning(f"Redis 删除失败: task={task_id}, {e}") async def get_all_tasks(self) -> Dict[str, Dict[str, Any]]: """获取所有任务;Redis 不可用时返回空 dict(调用方需容忍)""" if not self.is_connected: return {} try: pattern = f"{self._prefix}*" result = {} async for key in self._redis.scan_iter(match=pattern): task_id = key.replace(self._prefix, "") task = await self._load_any(key) if task: result[task_id] = task return result except Exception as e: logger.warning(f"Redis 扫描失败: {e}") return {} async def get_task_count(self) -> int: """获取任务总数;Redis 不可用时返回 0(调用方需容忍)""" if not self.is_connected: return 0 try: pattern = f"{self._prefix}*" count = 0 async for _ in self._redis.scan_iter(match=pattern): count += 1 return count except Exception as e: logger.warning(f"Redis 计数失败: {e}") return 0 # ---- 工具方法 ---- @staticmethod def _make_serializable(obj: Any) -> Any: """确保对象可 JSON 序列化""" if isinstance(obj, dict): return {k: RedisTaskManager._make_serializable(v) for k, v in obj.items()} if isinstance(obj, (list, tuple)): return [RedisTaskManager._make_serializable(v) for v in obj] if isinstance(obj, datetime): return obj.isoformat() if hasattr(obj, "value"): # Enum 类型 return obj.value if isinstance(obj, (int, float, str, bool, type(None))): return obj return str(obj) # 全局单例 redis_task_manager = RedisTaskManager.get_instance()