This commit is contained in:
2026-08-31 18:01:34 +08:00
parent 3ea59551db
commit bee439cf34
46 changed files with 1884 additions and 1898 deletions
+121 -35
View File
@@ -1,20 +1,27 @@
# services/redis_task_manager.py
"""Redis 任务管理器 - 替代内存字典,支持 TTL 自动清理"""
"""Redis 任务管理器 - 替代内存字典,支持 TTL 自动清理。
存储格式:Redis Hash(field -> JSON 字符串)。
- update_task 走 HSET 字段级原子更新,消除旧 get->merge->set 三步竞态
(后台处理流程与导出端点并发写同一任务时丢更新);
- 进度 tick 只重写变化字段,不再全量重写整个任务 blob;
- 兼容读旧 string 格式(升级前写入的在途任务),新写入一律 Hash。
"""
import json
import os
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 的任务状态管理"""
"""基于 Redis Hash 的任务状态管理"""
_instance: Optional["RedisTaskManager"] = None
@@ -31,14 +38,14 @@ class RedisTaskManager:
return cls._instance
async def connect(self):
"""连接 Redis"""
"""连接 Redis(配置统一来自 shared.config.settings,不再硬编码主机名)"""
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"))
host = settings.REDIS_HOST
port = settings.REDIS_PORT
password = settings.REDIS_PASSWORD
db = settings.REDIS_DB
try:
self._redis = aioredis.Redis(
@@ -84,6 +91,16 @@ class RedisTaskManager:
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]] = {}
@@ -102,55 +119,128 @@ class RedisTaskManager:
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
# 确保数据可序列化
serializable = self._make_serializable(data)
mapping = self._dump_mapping(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))
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, serializable)
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:
key = f"{self._prefix}{task_id}"
raw = await self._redis.get(key)
if raw:
return json.loads(raw)
return None
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]):
"""更新任务的部分字段"""
current = await self.get_task(task_id)
"""字段级原子更新(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))
await self.set_task(task_id, current)
self._fallback_set(task_id, current)
async def delete_task(self, task_id: str):
"""删除任务"""
"""删除任务(DEL 对 Hash/string 均有效)"""
if self.is_connected:
try:
key = f"{self._prefix}{task_id}"
await self._redis.delete(key)
await self._redis.delete(self._key(task_id))
return
except Exception as e:
logger.warning(f"Redis 删除失败,回退到内存: {e}")
@@ -162,16 +252,12 @@ class RedisTaskManager:
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:
async for key in self._redis.scan_iter(match=pattern):
task_id = key.replace(self._prefix, "")
raw = await self._redis.get(key)
if raw:
result[task_id] = json.loads(raw)
task = await self._load_any(key)
if task:
result[task_id] = task
return result
except Exception as e:
logger.warning(f"Redis 扫描失败,回退到内存: {e}")