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

478 lines
17 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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_idWindows 使用 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