323 lines
12 KiB
Python
323 lines
12 KiB
Python
# 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()
|