Files
geMoldInsight/src/services/redis_task_manager.py
T

225 lines
7.4 KiB
Python
Raw Normal View History

2026-04-13 09:43:24 +08:00
# services/redis_task_manager.py
"""Redis 任务管理器 - 替代内存字典,支持 TTL 自动清理"""
import json
import os
from typing import Dict, Any, Optional
from datetime import datetime
import redis.asyncio as aioredis
from utils.logger import get_logger
logger = get_logger(__name__)
class RedisTaskManager:
"""基于 Redis 的任务状态管理"""
_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"""
if self._connected and self._redis:
return
host = os.getenv("REDIS_HOST", "szcjw")
port = int(os.getenv("REDIS_PORT", "6379"))
password = os.getenv("REDIS_PASSWORD", "")
db = int(os.getenv("REDIS_DB", "0"))
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 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
# ---- 内存回退 ----
_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)
# ---- 公共接口 ----
async def set_task(self, task_id: str, data: Dict[str, Any], ttl: Optional[int] = None):
"""设置任务数据"""
effective_ttl = ttl or self._ttl
# 确保数据可序列化
serializable = self._make_serializable(data)
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
await self._redis.setex(key, effective_ttl, json.dumps(serializable, ensure_ascii=False))
return
except Exception as e:
logger.warning(f"Redis 写入失败,回退到内存: {e}")
self._fallback_set(task_id, serializable)
async def get_task(self, task_id: str) -> Optional[Dict[str, Any]]:
"""获取任务数据"""
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
raw = await self._redis.get(key)
if raw:
return json.loads(raw)
return None
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]):
"""更新任务的部分字段"""
current = await self.get_task(task_id)
if current is None:
logger.warning(f"任务 {task_id} 不存在,无法更新")
return
current.update(self._make_serializable(updates))
await self.set_task(task_id, current)
async def delete_task(self, task_id: str):
"""删除任务"""
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
await self._redis.delete(key)
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}*"
keys = []
async for key in self._redis.scan_iter(match=pattern):
keys.append(key)
result = {}
for key in keys:
task_id = key.replace(self._prefix, "")
raw = await self._redis.get(key)
if raw:
result[task_id] = json.loads(raw)
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()