Files
geMoldInsight/src/shared/services/redis_task_manager.py
T
2026-08-31 18:01:34 +08:00

323 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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()