# services/redis_task_manager.py """Redis 任务管理器 - 替代内存字典,支持 TTL 自动清理。 存储格式:Redis Hash(field -> JSON 字符串)。 - update_task 走 HSET 字段级原子更新,消除旧 get->merge->set 三步竞态 (后台处理流程与导出端点并发写同一任务时丢更新); - 进度 tick 只重写变化字段,不再全量重写整个任务 blob; - 兼容读旧 string 格式(升级前写入的在途任务),新写入一律 Hash。 """ 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 # ---- 内存回退 ---- _fallback_tasks: Dict[str, Dict[str, Any]] = {} def _fallback_set(self, task_id: str, data: Dict[str, Any]): self._fallback_tasks[task_id] = data def _fallback_get(self, task_id: str) -> Optional[Dict[str, Any]]: return self._fallback_tasks.get(task_id) def _fallback_delete(self, task_id: str): self._fallback_tasks.pop(task_id, None) def _fallback_all(self) -> Dict[str, Dict[str, Any]]: return dict(self._fallback_tasks) def _fallback_count(self) -> int: return len(self._fallback_tasks) # ---- 内部工具 ---- 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 格式清理)""" effective_ttl = ttl or self._ttl mapping = self._dump_mapping(data) if self.is_connected: 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() return except Exception as e: logger.warning(f"Redis 写入失败,回退到内存: {e}") self._fallback_set(task_id, self._make_serializable(data)) async def get_task(self, task_id: str) -> Optional[Dict[str, Any]]: """获取任务数据(Hash / 旧 string 兼容)""" if self.is_connected: try: return await self._load_any(self._key(task_id)) except Exception as e: logger.warning(f"Redis 读取失败,回退到内存: {e}") return self._fallback_get(task_id) async def update_task(self, task_id: str, updates: Dict[str, Any]): """字段级原子更新(HSET),无读改写竞态。 兼容旧 string 格式:先迁移为 Hash 再更新。 """ mapping = self._dump_mapping(updates) if self.is_connected: 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) return except Exception as e: logger.warning(f"Redis 更新失败,回退到内存: {e}") # 内存回退保持读改写语义(单进程内存无并发竞态) current = self._fallback_get(task_id) if current is None: logger.warning(f"任务 {task_id} 不存在,无法更新") return current.update(self._make_serializable(updates)) self._fallback_set(task_id, current) async def delete_task(self, task_id: str): """删除任务(DEL 对 Hash/string 均有效)""" if self.is_connected: try: await self._redis.delete(self._key(task_id)) return except Exception as e: logger.warning(f"Redis 删除失败,回退到内存: {e}") self._fallback_delete(task_id) async def get_all_tasks(self) -> Dict[str, Dict[str, Any]]: """获取所有任务""" if self.is_connected: 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 self._fallback_all() async def get_task_count(self) -> int: """获取任务总数""" if self.is_connected: 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 self._fallback_count() async def cleanup_old_tasks(self, max_age_seconds: int = 86400 * 7): """清理过期任务(Redis 由 TTL 自动管理,内存回退需手动清理)""" now = datetime.now() to_delete = [] for task_id, task in self._fallback_tasks.items(): completed_at = task.get("completed_at") if completed_at: try: completed_dt = datetime.fromisoformat(completed_at) if (now - completed_dt).total_seconds() > max_age_seconds: to_delete.append(task_id) except (ValueError, TypeError): pass for task_id in to_delete: del self._fallback_tasks[task_id] if to_delete: logger.info(f"清理了 {len(to_delete)} 个过期内存任务") # ---- 工具方法 ---- @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()