478 lines
17 KiB
Python
478 lines
17 KiB
Python
from __future__ import annotations
|
||
|
||
import ctypes
|
||
import json
|
||
import os
|
||
import socket
|
||
import threading
|
||
import time
|
||
import uuid
|
||
from dataclasses import dataclass, field
|
||
from functools import lru_cache
|
||
from pathlib import Path
|
||
from typing import Any
|
||
|
||
from app.config import settings
|
||
from app.enums.celery_runtime import CeleryRuntimeEvent, WorkerIdentityQuality
|
||
from app.services.operation_log_service import log_operation_event
|
||
from app.services.redis_registry_service import (
|
||
RedisExecutionLockUnavailable,
|
||
get_registry_redis,
|
||
)
|
||
|
||
try:
|
||
from celery import current_task
|
||
except Exception: # pragma: no cover
|
||
current_task = None # type: ignore[assignment]
|
||
|
||
|
||
_PROCESS_STARTED_AT = int(time.time())
|
||
_PROCESS_INSTANCE_ID = uuid.uuid4().hex
|
||
_PROCESS_IDENTITY_LOCK = threading.RLock()
|
||
_WORKER_REGISTRATION_LOCK = threading.RLock()
|
||
_HEARTBEAT_THROTTLE_LOCK = threading.RLock()
|
||
|
||
_ENV_INSTANCE_TOKEN = "CELERY_WORKER_INSTANCE_TOKEN"
|
||
_ENV_MAIN_PID = "CELERY_WORKER_MAIN_PID"
|
||
_ENV_STARTED_AT = "CELERY_WORKER_STARTED_AT"
|
||
_ENV_HOST_BOOT_ID = "CELERY_HOST_BOOT_ID"
|
||
_ENV_WORKER_NAME = "CELERY_WORKER_NODE_NAME"
|
||
_ENV_BEFORE_POOL = "CELERY_WORKER_IDENTITY_BEFORE_POOL"
|
||
|
||
_registered_worker: "WorkerRegistration | None" = None
|
||
_last_heartbeat_attempt_monotonic = 0.0
|
||
_last_stale_scan_attempt_monotonic = 0.0
|
||
_heartbeat_failure_count = 0
|
||
|
||
|
||
def _parse_optional_int(value: Any) -> int | None:
|
||
raw = str(value or "").strip()
|
||
if not raw:
|
||
return None
|
||
try:
|
||
parsed = int(raw)
|
||
except (TypeError, ValueError):
|
||
return None
|
||
return parsed if parsed > 0 else None
|
||
|
||
|
||
def _parse_optional_float(value: Any) -> float | None:
|
||
raw = str(value or "").strip()
|
||
if not raw:
|
||
return None
|
||
try:
|
||
return float(raw)
|
||
except (TypeError, ValueError):
|
||
return None
|
||
|
||
|
||
@lru_cache(maxsize=1)
|
||
def _boot_id() -> str:
|
||
"""返回同一次操作系统启动期间稳定的主机启动标识。
|
||
|
||
Linux 直接读取内核 boot_id;Windows 使用 GetTickCount64 估算启动时间并
|
||
生成 UUID5。无法获取时返回 unknown-boot。该字段仅用于辅助诊断,不参与
|
||
任务锁或数据库 fencing 的最终正确性判断。
|
||
"""
|
||
linux_path = Path("/proc/sys/kernel/random/boot_id")
|
||
try:
|
||
value = linux_path.read_text(encoding="utf-8").strip()
|
||
if value:
|
||
return value
|
||
except Exception:
|
||
pass
|
||
|
||
if os.name == "nt":
|
||
try:
|
||
uptime_ms = int(ctypes.windll.kernel32.GetTickCount64()) # type: ignore[attr-defined]
|
||
# 以 10 秒为粒度消除多个进程独立计算时的亚秒抖动。
|
||
boot_epoch_bucket = int((time.time() - uptime_ms / 1000.0) // 10 * 10)
|
||
source = f"windows-boot:{socket.gethostname()}:{boot_epoch_bucket}"
|
||
return str(uuid.uuid5(uuid.NAMESPACE_OID, source))
|
||
except Exception:
|
||
pass
|
||
|
||
return "unknown-boot"
|
||
|
||
|
||
def initialize_worker_main_identity(
|
||
worker_name: str | None = None,
|
||
*,
|
||
before_pool: bool,
|
||
) -> None:
|
||
"""在 Celery Worker 主进程中初始化一次实例 token。
|
||
|
||
before_pool=True 必须由 celeryd_init/worker_init 调用,以保证 Linux prefork
|
||
子进程通过 fork 继承相同 token。worker_ready 只允许做 late fallback,且
|
||
late fallback 会关闭实例级精准恢复,避免主进程与已创建子进程身份不一致。
|
||
"""
|
||
normalized_name = str(worker_name or "").strip()
|
||
current_pid = os.getpid()
|
||
|
||
should_log = False
|
||
with _PROCESS_IDENTITY_LOCK:
|
||
existing_pid = _parse_optional_int(os.getenv(_ENV_MAIN_PID))
|
||
existing_token = str(os.getenv(_ENV_INSTANCE_TOKEN, "") or "").strip()
|
||
existing_before_pool = str(os.getenv(_ENV_BEFORE_POOL, "") or "").strip() == "1"
|
||
|
||
if not existing_token or existing_pid != current_pid:
|
||
os.environ[_ENV_INSTANCE_TOKEN] = uuid.uuid4().hex
|
||
os.environ[_ENV_MAIN_PID] = str(current_pid)
|
||
os.environ[_ENV_STARTED_AT] = str(int(time.time()))
|
||
os.environ[_ENV_HOST_BOOT_ID] = _boot_id()
|
||
os.environ[_ENV_BEFORE_POOL] = "1" if before_pool else "0"
|
||
should_log = True
|
||
elif before_pool and not existing_before_pool:
|
||
# celeryd_init 与 worker_init 都可能触发,幂等提升为 before-pool。
|
||
os.environ[_ENV_BEFORE_POOL] = "1"
|
||
should_log = True
|
||
|
||
if normalized_name:
|
||
os.environ[_ENV_WORKER_NAME] = normalized_name
|
||
|
||
_process_identity_base.cache_clear()
|
||
|
||
if should_log:
|
||
identity = current_worker_identity(worker_name_override=normalized_name or None)
|
||
log_operation_event(
|
||
domain="celery_runtime",
|
||
event_type=(
|
||
CeleryRuntimeEvent.WORKER_IDENTITY_INITIALIZED.value
|
||
if identity.supports_targeted_recovery
|
||
else CeleryRuntimeEvent.WORKER_IDENTITY_FALLBACK.value
|
||
),
|
||
event_status="success" if identity.supports_targeted_recovery else "warning",
|
||
source="worker_init",
|
||
detail={
|
||
"worker_name": identity.worker_name,
|
||
"worker_instance_id": identity.worker_instance_id,
|
||
"worker_main_pid": identity.worker_main_pid,
|
||
"execution_pid": identity.execution_pid,
|
||
"host_boot_id": identity.host_boot_id,
|
||
"identity_quality": identity.identity_quality,
|
||
"supports_targeted_recovery": identity.supports_targeted_recovery,
|
||
"before_pool": before_pool,
|
||
},
|
||
)
|
||
|
||
|
||
def reset_process_identity_cache() -> None:
|
||
"""prefork 子进程启动后清理 fork 前的 Python 对象缓存。
|
||
|
||
环境变量中的主 Worker token 保留;每个子进程重新构建自己的执行 PID 和
|
||
线程 ID,不会生成新的 Worker 实例 token。
|
||
"""
|
||
_process_identity_base.cache_clear()
|
||
|
||
|
||
@dataclass(frozen=True, slots=True)
|
||
class _ProcessIdentityBase:
|
||
worker_instance_token: str
|
||
worker_main_pid: int | None
|
||
worker_started_at: int
|
||
host_boot_id: str
|
||
identity_quality: str
|
||
supports_targeted_recovery: bool
|
||
|
||
|
||
@lru_cache(maxsize=1)
|
||
def _process_identity_base() -> _ProcessIdentityBase:
|
||
token = str(os.getenv(_ENV_INSTANCE_TOKEN, "") or "").strip()
|
||
before_pool = str(os.getenv(_ENV_BEFORE_POOL, "") or "").strip() == "1"
|
||
worker_main_pid = _parse_optional_int(os.getenv(_ENV_MAIN_PID))
|
||
started_at = _parse_optional_float(os.getenv(_ENV_STARTED_AT))
|
||
host_boot_id = str(os.getenv(_ENV_HOST_BOOT_ID, "") or "").strip() or _boot_id()
|
||
|
||
supports_targeted_recovery = bool(token and before_pool)
|
||
if not token:
|
||
token = _PROCESS_INSTANCE_ID
|
||
|
||
execution_pid = os.getpid()
|
||
if not supports_targeted_recovery:
|
||
quality = WorkerIdentityQuality.FALLBACK.value
|
||
elif worker_main_pid is None or host_boot_id == "unknown-boot":
|
||
quality = WorkerIdentityQuality.INSTANCE_TOKEN_ONLY.value
|
||
else:
|
||
quality = WorkerIdentityQuality.FULL.value
|
||
|
||
return _ProcessIdentityBase(
|
||
worker_instance_token=token,
|
||
worker_main_pid=worker_main_pid,
|
||
worker_started_at=int(started_at or _PROCESS_STARTED_AT),
|
||
host_boot_id=host_boot_id,
|
||
identity_quality=quality,
|
||
supports_targeted_recovery=supports_targeted_recovery,
|
||
)
|
||
|
||
|
||
@dataclass(frozen=True, slots=True)
|
||
class WorkerIdentity:
|
||
worker_name: str
|
||
worker_instance_id: str
|
||
worker_instance_token: str
|
||
host: str
|
||
host_boot_id: str
|
||
worker_main_pid: int | None
|
||
execution_pid: int
|
||
execution_thread_id: int
|
||
started_at: int
|
||
identity_quality: str
|
||
supports_targeted_recovery: bool
|
||
celery_task_id: str | None = None
|
||
|
||
|
||
def current_worker_identity(*, worker_name_override: str | None = None) -> WorkerIdentity:
|
||
hostname = socket.gethostname()
|
||
celery_task_id: str | None = None
|
||
request_worker_name = ""
|
||
|
||
try:
|
||
request = getattr(current_task, "request", None)
|
||
request_worker_name = str(getattr(request, "hostname", "") or "").strip()
|
||
celery_task_id = str(getattr(request, "id", "") or "").strip() or None
|
||
except Exception:
|
||
pass
|
||
|
||
worker_name = (
|
||
str(worker_name_override or "").strip()
|
||
or request_worker_name
|
||
or str(os.getenv(_ENV_WORKER_NAME, "") or "").strip()
|
||
or hostname
|
||
)
|
||
base = _process_identity_base()
|
||
instance_id = f"{worker_name}:{base.host_boot_id}:{base.worker_instance_token}"
|
||
|
||
return WorkerIdentity(
|
||
worker_name=worker_name,
|
||
worker_instance_id=instance_id,
|
||
worker_instance_token=base.worker_instance_token,
|
||
host=hostname,
|
||
host_boot_id=base.host_boot_id,
|
||
worker_main_pid=base.worker_main_pid,
|
||
execution_pid=os.getpid(),
|
||
execution_thread_id=threading.get_ident(),
|
||
started_at=base.worker_started_at,
|
||
identity_quality=base.identity_quality,
|
||
supports_targeted_recovery=base.supports_targeted_recovery,
|
||
celery_task_id=celery_task_id,
|
||
)
|
||
|
||
|
||
@dataclass(frozen=True, slots=True)
|
||
class WorkerRegistration:
|
||
identity: WorkerIdentity
|
||
queues: tuple[str, ...] = field(default_factory=tuple)
|
||
pool_type: str | None = None
|
||
configured_concurrency: int | None = None
|
||
|
||
def payload(self, *, heartbeat_at: int) -> dict[str, Any]:
|
||
return {
|
||
"runtime_schema_version": int(settings.CELERY_RUNTIME_SCHEMA_VERSION or 2),
|
||
"worker_name": self.identity.worker_name,
|
||
"worker_instance_id": self.identity.worker_instance_id,
|
||
"host": self.identity.host,
|
||
"host_boot_id": self.identity.host_boot_id,
|
||
"worker_main_pid": self.identity.worker_main_pid,
|
||
"started_at": self.identity.started_at,
|
||
"identity_quality": self.identity.identity_quality,
|
||
"supports_targeted_recovery": self.identity.supports_targeted_recovery,
|
||
"queues": list(self.queues),
|
||
"pool_type": self.pool_type,
|
||
"configured_concurrency": self.configured_concurrency,
|
||
"heartbeat_at": heartbeat_at,
|
||
}
|
||
|
||
|
||
def _worker_instance_key(worker_instance_id: str) -> str:
|
||
return f"{settings.CELERY_RUNTIME_WORKER_INSTANCE_PREFIX}:{worker_instance_id}"
|
||
|
||
|
||
def _worker_name_instances_key(worker_name: str) -> str:
|
||
return f"{settings.CELERY_RUNTIME_WORKER_NAME_INSTANCE_ZSET_PREFIX}:{worker_name}"
|
||
|
||
|
||
def _normalize_queues(values: Any) -> tuple[str, ...]:
|
||
if values is None:
|
||
return ()
|
||
if isinstance(values, str):
|
||
values = [values]
|
||
result: list[str] = []
|
||
try:
|
||
iterator = iter(values)
|
||
except TypeError:
|
||
return ()
|
||
for value in iterator:
|
||
name = str(getattr(value, "name", value) or "").strip()
|
||
if name and name not in result:
|
||
result.append(name)
|
||
return tuple(result)
|
||
|
||
|
||
async def register_worker_instance(
|
||
*,
|
||
worker_name: str,
|
||
queues: Any = None,
|
||
pool_type: str | None = None,
|
||
configured_concurrency: int | None = None,
|
||
) -> WorkerIdentity:
|
||
"""注册 Worker 主实例;Redis 不可用时不阻塞 Worker 启动。"""
|
||
global _registered_worker, _last_heartbeat_attempt_monotonic
|
||
|
||
normalized_name = str(worker_name or "").strip() or socket.gethostname()
|
||
os.environ[_ENV_WORKER_NAME] = normalized_name
|
||
identity = current_worker_identity(worker_name_override=normalized_name)
|
||
registration = WorkerRegistration(
|
||
identity=identity,
|
||
queues=_normalize_queues(queues),
|
||
pool_type=str(pool_type or "").strip() or None,
|
||
configured_concurrency=_parse_optional_int(configured_concurrency),
|
||
)
|
||
|
||
with _WORKER_REGISTRATION_LOCK:
|
||
_registered_worker = registration
|
||
|
||
now = int(time.time())
|
||
redis = await get_registry_redis()
|
||
if redis is None:
|
||
raise RedisExecutionLockUnavailable("Redis unavailable while registering Celery worker instance")
|
||
|
||
ttl = max(30, int(settings.CELERY_RUNTIME_WORKER_HEARTBEAT_TTL_SECONDS or 90))
|
||
index_ttl = max(ttl * 2, int(settings.CELERY_RUNTIME_WORKER_TASK_SET_TTL_SECONDS or 86400))
|
||
payload = json.dumps(registration.payload(heartbeat_at=now), ensure_ascii=False, default=str)
|
||
|
||
pipe = redis.pipeline(transaction=False)
|
||
pipe.set(_worker_instance_key(identity.worker_instance_id), payload, ex=ttl)
|
||
pipe.zadd(_worker_name_instances_key(identity.worker_name), {identity.worker_instance_id: now})
|
||
pipe.expire(_worker_name_instances_key(identity.worker_name), index_ttl)
|
||
await pipe.execute()
|
||
|
||
with _HEARTBEAT_THROTTLE_LOCK:
|
||
_last_heartbeat_attempt_monotonic = time.monotonic()
|
||
|
||
log_operation_event(
|
||
domain="celery_runtime",
|
||
event_type=CeleryRuntimeEvent.WORKER_REGISTERED.value,
|
||
event_status="success",
|
||
source="worker_ready",
|
||
detail=registration.payload(heartbeat_at=now),
|
||
)
|
||
return identity
|
||
|
||
|
||
def claim_worker_heartbeat_slot() -> bool:
|
||
"""对 Celery heartbeat_sent 信号做本进程节流。"""
|
||
global _last_heartbeat_attempt_monotonic
|
||
|
||
interval = max(5, int(settings.CELERY_RUNTIME_WORKER_HEARTBEAT_INTERVAL_SECONDS or 30))
|
||
now = time.monotonic()
|
||
with _HEARTBEAT_THROTTLE_LOCK:
|
||
if now - _last_heartbeat_attempt_monotonic < interval:
|
||
return False
|
||
_last_heartbeat_attempt_monotonic = now
|
||
return True
|
||
|
||
|
||
def claim_worker_stale_scan_slot() -> bool:
|
||
"""限制同一 Worker 主进程的旧实例扫描频率。"""
|
||
global _last_stale_scan_attempt_monotonic
|
||
|
||
interval = max(30, int(settings.CELERY_RUNTIME_WORKER_STALE_SCAN_INTERVAL_SECONDS or 120))
|
||
now = time.monotonic()
|
||
with _HEARTBEAT_THROTTLE_LOCK:
|
||
if now - _last_stale_scan_attempt_monotonic < interval:
|
||
return False
|
||
_last_stale_scan_attempt_monotonic = now
|
||
return True
|
||
|
||
|
||
async def heartbeat_current_worker_instance() -> bool:
|
||
"""刷新 Worker 主实例 TTL;不刷新任何业务任务锁。"""
|
||
global _heartbeat_failure_count
|
||
|
||
with _WORKER_REGISTRATION_LOCK:
|
||
registration = _registered_worker
|
||
if registration is None:
|
||
return False
|
||
|
||
redis = await get_registry_redis()
|
||
if redis is None:
|
||
_heartbeat_failure_count += 1
|
||
_log_worker_heartbeat_failure_if_needed(registration, "redis_unavailable")
|
||
return False
|
||
|
||
now = int(time.time())
|
||
ttl = max(30, int(settings.CELERY_RUNTIME_WORKER_HEARTBEAT_TTL_SECONDS or 90))
|
||
index_ttl = max(ttl * 2, int(settings.CELERY_RUNTIME_WORKER_TASK_SET_TTL_SECONDS or 86400))
|
||
payload = json.dumps(registration.payload(heartbeat_at=now), ensure_ascii=False, default=str)
|
||
|
||
try:
|
||
pipe = redis.pipeline(transaction=False)
|
||
pipe.set(_worker_instance_key(registration.identity.worker_instance_id), payload, ex=ttl)
|
||
pipe.zadd(
|
||
_worker_name_instances_key(registration.identity.worker_name),
|
||
{registration.identity.worker_instance_id: now},
|
||
)
|
||
pipe.expire(_worker_name_instances_key(registration.identity.worker_name), index_ttl)
|
||
await pipe.execute()
|
||
except Exception as exc:
|
||
_heartbeat_failure_count += 1
|
||
_log_worker_heartbeat_failure_if_needed(registration, str(exc))
|
||
return False
|
||
|
||
_heartbeat_failure_count = 0
|
||
return True
|
||
|
||
|
||
def _log_worker_heartbeat_failure_if_needed(registration: WorkerRegistration, error: str) -> None:
|
||
threshold = max(1, int(settings.CELERY_RUNTIME_WORKER_HEARTBEAT_FAILURE_LOG_THRESHOLD or 3))
|
||
if _heartbeat_failure_count != threshold and _heartbeat_failure_count % (threshold * 5) != 0:
|
||
return
|
||
log_operation_event(
|
||
domain="celery_runtime",
|
||
event_type=CeleryRuntimeEvent.WORKER_HEARTBEAT_LOST.value,
|
||
event_status="failed",
|
||
source="worker_heartbeat",
|
||
detail={
|
||
"worker_name": registration.identity.worker_name,
|
||
"worker_instance_id": registration.identity.worker_instance_id,
|
||
"failure_count": _heartbeat_failure_count,
|
||
"error": error,
|
||
},
|
||
error=error,
|
||
)
|
||
|
||
|
||
async def unregister_current_worker_instance() -> None:
|
||
"""优雅退出时删除活跃实例 key,保留名称索引供旧任务精准恢复。"""
|
||
global _registered_worker
|
||
|
||
with _WORKER_REGISTRATION_LOCK:
|
||
registration = _registered_worker
|
||
_registered_worker = None
|
||
if registration is None:
|
||
return
|
||
|
||
redis = await get_registry_redis()
|
||
if redis is not None:
|
||
try:
|
||
await redis.delete(_worker_instance_key(registration.identity.worker_instance_id))
|
||
except Exception:
|
||
pass
|
||
|
||
log_operation_event(
|
||
domain="celery_runtime",
|
||
event_type=CeleryRuntimeEvent.WORKER_SHUTDOWN.value,
|
||
event_status="success",
|
||
source="worker_shutdown",
|
||
detail={
|
||
"worker_name": registration.identity.worker_name,
|
||
"worker_instance_id": registration.identity.worker_instance_id,
|
||
"worker_main_pid": registration.identity.worker_main_pid,
|
||
},
|
||
)
|
||
|
||
|
||
def registered_worker_identity() -> WorkerIdentity | None:
|
||
with _WORKER_REGISTRATION_LOCK:
|
||
return _registered_worker.identity if _registered_worker is not None else None
|