Files
video-gen/video-gen-api/app/services/redis_registry_service.py
2026-06-11 17:54:40 +08:00

355 lines
11 KiB
Python

# app/services/redis_registry_service.py
from __future__ import annotations
import asyncio
import inspect
import json
import logging
import os
import threading
import uuid
from datetime import datetime, timezone
from typing import Any, Dict, Iterable, List, Optional, Union
from app.config import settings
try:
from redis.exceptions import RedisError
except ImportError: # pragma: no cover - redis 未安装时降级
RedisError = RuntimeError # type: ignore[assignment]
logger = logging.getLogger("video_gen")
_redis_clients: Dict[tuple[int, int, int], Any] = {}
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def ensure_aware_utc(value: Optional[datetime]) -> Optional[datetime]:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
def datetime_to_epoch(value: Optional[datetime]) -> int:
checked_value = ensure_aware_utc(value) or utc_now()
return int(checked_value.timestamp())
def normalize_registry_score(value: Optional[Union[datetime, int, float]]) -> int:
if isinstance(value, datetime):
return datetime_to_epoch(value)
if value is None:
return datetime_to_epoch(utc_now())
return int(float(value))
def registry_redis_url() -> str:
"""Celery 容灾注册表统一使用 Celery broker Redis。
不能改成只读 settings.REDIS_URL,否则线上 CELERY_BROKER_URL 使用独立
Redis DB 时,旧下载 active 注册表会被写到另一个库,导致恢复扫描失效。
"""
return settings.CELERY_BROKER_URL or settings.REDIS_URL or ""
def _is_supported_redis_url(redis_url: str) -> bool:
if not redis_url:
return False
lowered = redis_url.lower()
return lowered.startswith(("redis://", "rediss://", "unix://"))
async def get_registry_redis() -> Optional[Any]:
"""获取 Celery 容灾 Redis 连接。
重点:redis.asyncio 的连接/连接池绑定 event loop,不能跨 loop 复用。
Celery -P threads 或 worker_ready + task 线程混用时,如果使用单个全局
Redis 客户端,会触发 got Future attached to a different loop。
因此这里按 pid + thread_id + event_loop_id 缓存客户端,确保同一个客户端
只在创建它的事件循环里使用。Redis 不可用时返回 None,调用方降级为
DB fallback,不能影响生成主链路。
"""
redis_url = registry_redis_url()
if not _is_supported_redis_url(redis_url):
if redis_url:
logger.warning(
"Celery 容灾 Redis 注册表仅支持 redis/rediss/unix URL,当前 broker 不是 Redis,降级为 DB 容灾。url=%s",
redis_url,
)
return None
try:
from redis.asyncio import Redis
except ImportError as exc:
logger.warning(
"Celery 容灾 Redis 注册表不可用,redis 依赖未安装。error=%s",
exc,
)
return None
try:
loop = asyncio.get_running_loop()
except RuntimeError:
return None
client_key = (os.getpid(), threading.get_ident(), id(loop))
cached = _redis_clients.get(client_key)
if cached is not None:
return cached
try:
redis_client = Redis.from_url(redis_url, decode_responses=True)
await redis_client.ping()
_redis_clients[client_key] = redis_client
return redis_client
except (RedisError, OSError, RuntimeError) as exc:
logger.warning(
"Celery 容灾 Redis 注册表不可用,降级为仅 DB 容灾。error=%s",
exc,
)
_redis_clients.pop(client_key, None)
return None
async def close_registry_redis() -> None:
"""关闭当前进程内已缓存的 Redis 注册表连接。
关闭动作尽量只关闭当前 event loop 对应的客户端;如果调用方处于进程
退出阶段,则逐个尝试关闭,失败忽略,避免影响 worker 退出。
"""
if not _redis_clients:
return
try:
loop = asyncio.get_running_loop()
current_key = (os.getpid(), threading.get_ident(), id(loop))
items = [(current_key, _redis_clients.pop(current_key, None))]
except RuntimeError:
items = list(_redis_clients.items())
_redis_clients.clear()
for _, client in items:
if client is None:
continue
try:
close_method = getattr(client, "close", None) or getattr(client, "aclose", None)
if close_method is None:
continue
close_result = close_method()
if inspect.isawaitable(close_result):
await close_result
except (RedisError, OSError, RuntimeError) as exc:
logger.debug("关闭 Celery 容灾 Redis 注册表连接失败。error=%s", exc)
async def redis_upsert_registry_item(
*,
hash_key: str,
zset_key: str,
item_id: str,
payload: Dict[str, Any],
check_at: Optional[Union[datetime, int, float]],
log_context: str = "registry",
) -> None:
redis = await get_registry_redis()
if redis is None:
return
score = normalize_registry_score(check_at)
updated_payload = dict(payload)
updated_payload["check_at"] = score
updated_payload["updated_at"] = updated_payload.get("updated_at") or datetime_to_epoch(utc_now())
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.hset(hash_key, item_id, json.dumps(updated_payload, ensure_ascii=False, default=str))
pipe.zadd(zset_key, {item_id: score})
await pipe.execute()
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"写入 Redis 注册表失败。context=%s, item_id=%s, error=%s",
log_context,
item_id,
exc,
)
async def redis_remove_registry_item(
*,
hash_key: str,
zset_key: str,
item_id: str,
log_context: str = "registry",
) -> None:
redis = await get_registry_redis()
if redis is None:
return
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.hdel(hash_key, item_id)
pipe.zrem(zset_key, item_id)
await pipe.execute()
except (RedisError, OSError, RuntimeError) as exc:
logger.warning(
"删除 Redis 注册表失败。context=%s, item_id=%s, error=%s",
log_context,
item_id,
exc,
)
async def redis_get_due_registry_ids(
*,
zset_key: str,
limit: Optional[int] = None,
now: Optional[datetime] = None,
log_context: str = "registry",
) -> List[str]:
redis = await get_registry_redis()
if redis is None:
return []
batch_limit = int(limit or 100)
score = datetime_to_epoch(now or utc_now())
try:
result = await redis.zrangebyscore(
zset_key,
min="-inf",
max=score,
start=0,
num=batch_limit,
)
return [str(item) for item in result]
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning("扫描 Redis ZSet 失败。context=%s, error=%s", log_context, exc)
return []
async def redis_get_registry_payloads(
*,
hash_key: str,
item_ids: Iterable[str],
log_context: str = "registry",
) -> Dict[str, Dict[str, Any]]:
cleaned_item_ids = [str(item) for item in item_ids if item]
if not cleaned_item_ids:
return {}
redis = await get_registry_redis()
if redis is None:
return {}
try:
raw_values = await redis.hmget(hash_key, cleaned_item_ids)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning("读取 Redis Hash 失败。context=%s, error=%s", log_context, exc)
return {}
result: Dict[str, Dict[str, Any]] = {}
for item_id, raw in zip(cleaned_item_ids, raw_values):
if not raw:
continue
try:
value = json.loads(raw)
except (TypeError, ValueError, json.JSONDecodeError):
continue
if isinstance(value, dict):
result[item_id] = value
return result
async def redis_postpone_registry_item(
*,
hash_key: str,
zset_key: str,
item_id: str,
payload: Optional[Dict[str, Any]] = None,
check_at: Optional[Union[datetime, int, float]] = None,
log_context: str = "registry",
) -> None:
redis = await get_registry_redis()
if redis is None:
return
score = normalize_registry_score(check_at)
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.zadd(zset_key, {item_id: score})
if payload is not None:
updated_payload = dict(payload)
updated_payload["check_at"] = score
updated_payload["updated_at"] = datetime_to_epoch(utc_now())
pipe.hset(hash_key, item_id, json.dumps(updated_payload, ensure_ascii=False, default=str))
await pipe.execute()
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"刷新 Redis 注册表检查时间失败。context=%s, item_id=%s, error=%s",
log_context,
item_id,
exc,
)
async def redis_acquire_lock(
*,
lock_key: str,
ttl_seconds: int,
token: Optional[str] = None,
log_context: str = "lock",
) -> Optional[str]:
"""尝试获取 Redis 分布式锁。
返回 token 表示抢锁成功;返回 None 表示 Redis 不可用或锁已被其他 worker 持有。
"""
redis = await get_registry_redis()
if redis is None:
return None
lock_token = token or uuid.uuid4().hex
ttl = max(1, int(ttl_seconds or 60))
try:
acquired = await redis.set(lock_key, lock_token, nx=True, ex=ttl)
return lock_token if acquired else None
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning("获取 Redis 锁失败。context=%s, lock_key=%s, error=%s", log_context, lock_key, exc)
return None
async def redis_release_lock(
*,
lock_key: str,
token: str,
log_context: str = "lock",
) -> bool:
"""只释放 token 匹配的锁,避免误删其他 worker 新抢到的锁。"""
redis = await get_registry_redis()
if redis is None:
return False
script = """
if redis.call('get', KEYS[1]) == ARGV[1] then
return redis.call('del', KEYS[1])
else
return 0
end
"""
try:
released = await redis.eval(script, 1, lock_key, token)
return bool(released)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning("释放 Redis 锁失败。context=%s, lock_key=%s, error=%s", log_context, lock_key, exc)
return False