Files
2026-07-22 14:48:29 +08:00

544 lines
19 KiB
Python

from __future__ import annotations
import asyncio
import hashlib
import json
import logging
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any, Awaitable, Callable, Mapping
from app.config import settings
from app.enums.celery_runtime import CeleryRuntimeEvent, CeleryRuntimeState
from app.services.celery_runtime.worker_service import current_worker_identity
from app.services.operation_log_service import log_operation_event
from app.services.redis_registry_service import (
RedisExecutionLockError,
RedisExecutionLockLost,
RedisExecutionLockUnavailable,
datetime_to_epoch,
get_registry_redis,
)
try:
from redis.exceptions import RedisError
except ImportError: # pragma: no cover
RedisError = RuntimeError # type: ignore[assignment]
logger = logging.getLogger("video_gen")
DbHeartbeat = Callable[[str], Awaitable[bool]]
# KEYS:
# 1 lock, 2 active hash, 3 active zset, 4 worker-instance task set,
# 5 runtime location hash.
#
# ARGV:
# 1 token, 2 lock ttl ms, 3 runtime id, 4 payload, 5 check_at,
# 6 worker task-set ttl seconds, 7 location payload, 8 track worker flag.
_ACQUIRE_SCRIPT = """
if redis.call('exists', KEYS[1]) == 1 then
return 0
end
redis.call('psetex', KEYS[1], ARGV[2], ARGV[1])
redis.call('hset', KEYS[2], ARGV[3], ARGV[4])
redis.call('zadd', KEYS[3], ARGV[5], ARGV[3])
if ARGV[8] == '1' then
redis.call('sadd', KEYS[4], ARGV[3])
redis.call('expire', KEYS[4], ARGV[6])
end
redis.call('hset', KEYS[5], ARGV[3], ARGV[7])
return 1
"""
_HEARTBEAT_SCRIPT = """
if redis.call('get', KEYS[1]) ~= ARGV[1] then
return 0
end
redis.call('pexpire', KEYS[1], ARGV[2])
local merged_payload = ARGV[4]
local current_payload = redis.call('hget', KEYS[2], ARGV[3])
if current_payload then
local current_ok, current_obj = pcall(cjson.decode, current_payload)
local update_ok, update_obj = pcall(cjson.decode, ARGV[4])
if current_ok and update_ok then
for key, value in pairs(update_obj) do
current_obj[key] = value
end
merged_payload = cjson.encode(current_obj)
end
end
redis.call('hset', KEYS[2], ARGV[3], merged_payload)
redis.call('zadd', KEYS[3], ARGV[5], ARGV[3])
if ARGV[8] == '1' then
redis.call('sadd', KEYS[4], ARGV[3])
redis.call('expire', KEYS[4], ARGV[6])
end
redis.call('hset', KEYS[5], ARGV[3], ARGV[7])
return 1
"""
_COMPLETE_SCRIPT = """
if redis.call('get', KEYS[1]) ~= ARGV[1] then
return 0
end
redis.call('del', KEYS[1])
redis.call('hdel', KEYS[2], ARGV[2])
redis.call('zrem', KEYS[3], ARGV[2])
redis.call('srem', KEYS[4], ARGV[2])
redis.call('hdel', KEYS[5], ARGV[2])
return 1
"""
def _now_epoch() -> int:
return datetime_to_epoch(datetime.now(timezone.utc))
def _token_digest(token: str) -> str:
return hashlib.sha256(token.encode("utf-8")).hexdigest()[:16]
def build_runtime_id(domain: str, owner_type: str, owner_id: str, attempt_no: int | None = None) -> str:
attempt = max(1, int(attempt_no or 1))
return f"{domain}:{owner_type}:{owner_id}:attempt:{attempt}"
def _worker_task_set_key(worker_instance_id: str) -> str:
return f"{settings.CELERY_RUNTIME_WORKER_TASK_SET_PREFIX}:{worker_instance_id}"
@dataclass(frozen=True, slots=True)
class RuntimeIdentity:
domain: str
owner_type: str
owner_id: str
attempt_no: int
task_name: str
queue: str
registry_item_id: str | None = None
@property
def runtime_id(self) -> str:
return self.registry_item_id or build_runtime_id(
self.domain,
self.owner_type,
self.owner_id,
self.attempt_no,
)
@dataclass(slots=True)
class CeleryRuntimeLease:
identity: RuntimeIdentity
lock_key: str
hash_key: str
zset_key: str
token: str
ttl_seconds: int
heartbeat_interval_seconds: int
payload: dict[str, Any]
db_heartbeat: DbHeartbeat | None = None
db_heartbeat_grace_seconds: int = 60
_stop_event: asyncio.Event = field(default_factory=asyncio.Event, init=False, repr=False)
_heartbeat_task: asyncio.Task[Any] | None = field(default=None, init=False, repr=False)
_lost_error: RedisExecutionLockError | None = field(default=None, init=False, repr=False)
@classmethod
async def acquire(
cls,
*,
identity: RuntimeIdentity,
lock_key: str,
hash_key: str,
zset_key: str,
token: str,
ttl_seconds: int,
heartbeat_interval_seconds: int,
pipeline_stage: str | None = None,
input_hash: str | None = None,
business_version: int | str | None = None,
extra_payload: Mapping[str, Any] | None = None,
db_heartbeat: DbHeartbeat | None = None,
) -> "CeleryRuntimeLease | None":
redis = await get_registry_redis()
if redis is None:
log_operation_event(
domain="celery_runtime",
event_type=CeleryRuntimeEvent.RUNTIME_REDIS_UNAVAILABLE.value,
event_status="failed",
source="celery",
task_id=identity.owner_id,
detail={"runtime_id": identity.runtime_id, "domain": identity.domain},
)
raise RedisExecutionLockUnavailable(f"Redis runtime unavailable: {identity.runtime_id}")
worker = current_worker_identity()
now = _now_epoch()
ttl = max(1, int(ttl_seconds or 60))
heartbeat_interval = max(
1,
min(ttl - 1 if ttl > 1 else 1, int(heartbeat_interval_seconds or 30)),
)
check_at = now + ttl
task_set_ttl = max(
ttl * 2,
int(settings.CELERY_RUNTIME_WORKER_TASK_SET_TTL_SECONDS or 86400),
)
task_set_key = _worker_task_set_key(worker.worker_instance_id)
track_worker = "1" if worker.supports_targeted_recovery else "0"
payload: dict[str, Any] = {
"runtime_schema_version": int(settings.CELERY_RUNTIME_SCHEMA_VERSION or 2),
"runtime_id": identity.runtime_id,
"domain": identity.domain,
"task_name": identity.task_name,
"queue": identity.queue,
"owner_type": identity.owner_type,
"owner_id": identity.owner_id,
"attempt_no": identity.attempt_no,
"business_version": business_version,
"input_hash": input_hash,
"worker_name": worker.worker_name,
"worker_instance_id": worker.worker_instance_id,
"worker_main_pid": worker.worker_main_pid,
"execution_pid": worker.execution_pid,
"execution_thread_id": worker.execution_thread_id,
"worker_started_at": worker.started_at,
"worker_identity_quality": worker.identity_quality,
"supports_targeted_recovery": worker.supports_targeted_recovery,
"host": worker.host,
"host_boot_id": worker.host_boot_id,
"celery_task_id": worker.celery_task_id,
"lock_key": lock_key,
"lock_token_digest": _token_digest(token),
"runtime_state": CeleryRuntimeState.ACTIVE.value,
"pipeline_stage": pipeline_stage,
"started_at": now,
"heartbeat_at": now,
"lease_until": check_at,
"check_at": check_at,
"recovery_count": 0,
}
if extra_payload:
payload["extra"] = dict(extra_payload)
location_payload = json.dumps(
{
"runtime_schema_version": int(settings.CELERY_RUNTIME_SCHEMA_VERSION or 2),
"runtime_id": identity.runtime_id,
"domain": identity.domain,
"hash_key": hash_key,
"zset_key": zset_key,
"lock_key": lock_key,
"worker_name": worker.worker_name,
"worker_instance_id": worker.worker_instance_id,
"supports_targeted_recovery": worker.supports_targeted_recovery,
},
ensure_ascii=False,
default=str,
)
log_operation_event(
domain="celery_runtime",
event_type=CeleryRuntimeEvent.RUNTIME_ACQUIRE_START.value,
event_status="started",
source="celery",
task_id=identity.owner_id,
detail={
"runtime_id": identity.runtime_id,
"worker_instance_id": worker.worker_instance_id,
"execution_pid": worker.execution_pid,
},
)
try:
acquired = await redis.eval(
_ACQUIRE_SCRIPT,
5,
lock_key,
hash_key,
zset_key,
task_set_key,
settings.CELERY_RUNTIME_LOCATION_HASH_KEY,
token,
ttl * 1000,
identity.runtime_id,
json.dumps(payload, ensure_ascii=False, default=str),
check_at,
task_set_ttl,
location_payload,
track_worker,
)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.exception("Celery runtime acquire failed. runtime_id=%s", identity.runtime_id)
raise RedisExecutionLockUnavailable(
f"Redis runtime acquire failed: {identity.runtime_id}: {exc}"
) from exc
if not acquired:
log_operation_event(
domain="celery_runtime",
event_type=CeleryRuntimeEvent.RUNTIME_LOCK_HELD.value,
event_status="skipped",
source="celery",
task_id=identity.owner_id,
detail={"runtime_id": identity.runtime_id, "lock_key": lock_key},
)
return None
lease = cls(
identity=identity,
lock_key=lock_key,
hash_key=hash_key,
zset_key=zset_key,
token=token,
ttl_seconds=ttl,
heartbeat_interval_seconds=heartbeat_interval,
payload=payload,
db_heartbeat=db_heartbeat,
db_heartbeat_grace_seconds=max(60, heartbeat_interval * 2),
)
lease._heartbeat_task = asyncio.create_task(lease._heartbeat_loop())
log_operation_event(
domain="celery_runtime",
event_type=CeleryRuntimeEvent.RUNTIME_ACQUIRED.value,
event_status="success",
source="celery",
task_id=identity.owner_id,
detail={
"runtime_id": identity.runtime_id,
"domain": identity.domain,
"attempt_no": identity.attempt_no,
"worker_instance_id": worker.worker_instance_id,
"worker_main_pid": worker.worker_main_pid,
"execution_pid": worker.execution_pid,
"lock_token_suffix": token[-8:],
"lease_until": check_at,
},
)
return lease
async def _heartbeat_loop(self) -> None:
while not self._stop_event.is_set():
try:
await asyncio.wait_for(
self._stop_event.wait(),
timeout=self.heartbeat_interval_seconds,
)
return
except asyncio.TimeoutError:
pass
try:
await self._heartbeat_once()
except RedisExecutionLockError as exc:
self._lost_error = exc
log_operation_event(
domain="celery_runtime",
event_type=CeleryRuntimeEvent.RUNTIME_HEARTBEAT_LOST.value,
event_status="failed",
source="celery",
task_id=self.identity.owner_id,
detail={"runtime_id": self.identity.runtime_id, "error": str(exc)},
error=str(exc),
)
return
except Exception as exc:
self._lost_error = RedisExecutionLockUnavailable(str(exc))
logger.exception(
"Celery runtime heartbeat failed. runtime_id=%s",
self.identity.runtime_id,
)
return
async def _heartbeat_once(self) -> None:
redis = await get_registry_redis()
if redis is None:
raise RedisExecutionLockUnavailable(
f"Redis runtime heartbeat unavailable: {self.identity.runtime_id}"
)
now = _now_epoch()
check_at = now + self.ttl_seconds
payload = dict(self.payload)
payload.update(
{
"heartbeat_at": now,
"lease_until": check_at,
"check_at": check_at,
# threads 模式下线程 ID可能随下一次执行变化,但同一 lease 生命周期固定。
"execution_pid": self.payload.get("execution_pid"),
"execution_thread_id": self.payload.get("execution_thread_id"),
}
)
worker_instance_id = str(payload.get("worker_instance_id") or "")
task_set_key = _worker_task_set_key(worker_instance_id)
task_set_ttl = max(
self.ttl_seconds * 2,
int(settings.CELERY_RUNTIME_WORKER_TASK_SET_TTL_SECONDS or 86400),
)
track_worker = "1" if bool(payload.get("supports_targeted_recovery")) else "0"
location_payload = json.dumps(
{
"runtime_schema_version": int(settings.CELERY_RUNTIME_SCHEMA_VERSION or 2),
"runtime_id": self.identity.runtime_id,
"domain": self.identity.domain,
"hash_key": self.hash_key,
"zset_key": self.zset_key,
"lock_key": self.lock_key,
"worker_name": payload.get("worker_name"),
"worker_instance_id": worker_instance_id,
"supports_targeted_recovery": bool(payload.get("supports_targeted_recovery")),
},
ensure_ascii=False,
default=str,
)
try:
renewed = await redis.eval(
_HEARTBEAT_SCRIPT,
5,
self.lock_key,
self.hash_key,
self.zset_key,
task_set_key,
settings.CELERY_RUNTIME_LOCATION_HASH_KEY,
self.token,
self.ttl_seconds * 1000,
self.identity.runtime_id,
json.dumps(payload, ensure_ascii=False, default=str),
check_at,
task_set_ttl,
location_payload,
track_worker,
)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
raise RedisExecutionLockUnavailable(
f"Redis runtime heartbeat failed: {self.identity.runtime_id}: {exc}"
) from exc
if not renewed:
raise RedisExecutionLockLost(
f"Redis runtime ownership lost: {self.identity.runtime_id}"
)
self.payload = payload
if self.db_heartbeat is not None:
started_at = int(self.payload.get("started_at") or now)
if now - started_at < self.db_heartbeat_grace_seconds:
return
owned = await self.db_heartbeat(self.token)
if not owned:
log_operation_event(
domain="celery_runtime",
event_type=CeleryRuntimeEvent.RUNTIME_DB_LEASE_LOST.value,
event_status="failed",
source="celery",
task_id=self.identity.owner_id,
detail={"runtime_id": self.identity.runtime_id},
)
raise RedisExecutionLockLost(
f"Database lease ownership lost: {self.identity.runtime_id}"
)
async def ensure_owned(self) -> None:
if self._lost_error is not None:
raise self._lost_error
redis = await get_registry_redis()
if redis is None:
raise RedisExecutionLockUnavailable(
f"Redis runtime unavailable: {self.identity.runtime_id}"
)
try:
value = await redis.get(self.lock_key)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
raise RedisExecutionLockUnavailable(
f"Redis runtime check failed: {self.identity.runtime_id}: {exc}"
) from exc
if str(value or "") != self.token:
self._lost_error = RedisExecutionLockLost(
f"Redis runtime ownership lost: {self.identity.runtime_id}"
)
raise self._lost_error
async def close(self) -> None:
self._stop_event.set()
if self._heartbeat_task is not None:
try:
await self._heartbeat_task
except Exception:
logger.debug("runtime heartbeat close failed", exc_info=True)
redis = await get_registry_redis()
if redis is None:
return
worker_instance_id = str(self.payload.get("worker_instance_id") or "")
task_set_key = _worker_task_set_key(worker_instance_id)
cleaned = False
try:
cleaned = bool(
await redis.eval(
_COMPLETE_SCRIPT,
5,
self.lock_key,
self.hash_key,
self.zset_key,
task_set_key,
settings.CELERY_RUNTIME_LOCATION_HASH_KEY,
self.token,
self.identity.runtime_id,
)
)
except Exception:
logger.warning(
"Celery runtime cleanup failed. runtime_id=%s",
self.identity.runtime_id,
exc_info=True,
)
log_operation_event(
domain="celery_runtime",
event_type=CeleryRuntimeEvent.RUNTIME_COMPLETED.value,
event_status="success" if cleaned else "skipped",
source="celery",
task_id=self.identity.owner_id,
detail={
"runtime_id": self.identity.runtime_id,
"ownership_cleanup": cleaned,
},
)
async def runtime_lock_exists(lock_key: str) -> bool:
redis = await get_registry_redis()
if redis is None:
raise RedisExecutionLockUnavailable(f"Redis unavailable while checking lock: {lock_key}")
try:
return bool(await redis.exists(lock_key))
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
raise RedisExecutionLockUnavailable(f"Redis lock check failed: {lock_key}: {exc}") from exc
async def runtime_lock_values(lock_keys: list[str]) -> dict[str, str | None]:
"""批量读取执行锁;Redis 不可用时 fail-closed。"""
if not lock_keys:
return {}
redis = await get_registry_redis()
if redis is None:
raise RedisExecutionLockUnavailable("Redis unavailable while batch checking runtime locks")
try:
values = await redis.mget(lock_keys)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
raise RedisExecutionLockUnavailable(f"Redis batch lock check failed: {exc}") from exc
return {
lock_key: (str(value) if value not in (None, "") else None)
for lock_key, value in zip(lock_keys, values)
}